Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@
@Profile("!ci | thread")
public class LeetcodeQuestionProcessService {

private static final ReentrantLock LOCK = new ReentrantLock();
private final ReentrantLock lock = new ReentrantLock();

private static final int MAX_JOBS_PER_RUN = 10;
private static final long REQUESTS_OVER_TIME = 1L;
Expand Down Expand Up @@ -89,7 +89,7 @@ private List<Job> claimBatch(final int maxSize) {
@Scheduled(initialDelay = 0, fixedDelay = 30, timeUnit = TimeUnit.MINUTES)
@Async
public CompletableFuture<Empty> drainQueue() {
if (!LOCK.tryLock()) {
if (!lock.tryLock()) {
log.info("thread attempted to drain queue, but queue is already being drained.");
return CompletableFuture.completedFuture(Empty.of());
}
Expand Down Expand Up @@ -119,7 +119,7 @@ public CompletableFuture<Empty> drainQueue() {
}
}
} finally {
LOCK.unlock();
lock.unlock();
}
return CompletableFuture.completedFuture(Empty.of());
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,11 @@
import static org.mockito.Mockito.*;

import io.github.bucket4j.BlockingBucket;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicInteger;
import java.util.concurrent.atomic.AtomicReference;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.BeforeEach;
Expand Down Expand Up @@ -75,59 +73,67 @@ void acquireShouldAppendToEndOfQueue() throws InterruptedException {

@Test
@Timeout(value = 5, unit = TimeUnit.SECONDS)
void acquireFastShouldAppendToStartOfQueue() throws InterruptedException {
CountDownLatch tickerCallLatch = new CountDownLatch(1);
CountDownLatch releaseTicker = new CountDownLatch(1);
void acquireFastShouldAppendToStartOfQueue() throws Exception {
var tickerStarted = new CountDownLatch(1);
var releaseTicker = new CountDownLatch(1);
var releaseNormal = new CountDownLatch(1);
var normalSelected = new CountDownLatch(1);
var calls = new AtomicInteger();

doAnswer(invocation -> {
tickerCallLatch.countDown();
releaseTicker.await();
int call = calls.incrementAndGet();
if (call == 1) {
tickerStarted.countDown();
releaseTicker.await();
} else if (call == 3) {
normalSelected.countDown();
releaseNormal.await();
}
return null;
})
.when(bucket)
.consume(1);

executor.submit(() -> {
try {
try {
var initial = executor.submit(() -> {
queueLock.acquire();
} catch (InterruptedException e) {
Thread.currentThread().interrupt();
}
});
assertTrue(tickerCallLatch.await(2, TimeUnit.SECONDS));

List<String> orderedCalls = Collections.synchronizedList(new ArrayList<>());
CountDownLatch doneLatch = new CountDownLatch(2);
return null;
});
assertTrue(tickerStarted.await(2, TimeUnit.SECONDS));

executor.submit(() -> {
try {
var normal = executor.submit(() -> {
queueLock.acquire();
orderedCalls.add("T2");
doneLatch.countDown();
} catch (InterruptedException e) {
Thread.currentThread().interrupt();
}
});
return null;
});
awaitQueueSize(1);

Thread.sleep(100);

executor.submit(() -> {
try {
var fast = executor.submit(() -> {
queueLock.acquireFast();
orderedCalls.add("T3");
doneLatch.countDown();
} catch (InterruptedException e) {
Thread.currentThread().interrupt();
}
});

Thread.sleep(100);

releaseTicker.countDown();
return null;
});
awaitQueueSize(2);

releaseTicker.countDown();
fast.get(2, TimeUnit.SECONDS);
assertTrue(normalSelected.await(2, TimeUnit.SECONDS));
assertFalse(normal.isDone(), "Normal request must wait until the fast request is released");

releaseNormal.countDown();
normal.get(2, TimeUnit.SECONDS);
initial.get(2, TimeUnit.SECONDS);
verify(bucket, times(3)).consume(1);
} finally {
releaseTicker.countDown();
releaseNormal.countDown();
}
}

assertTrue(doneLatch.await(2, TimeUnit.SECONDS));
assertEquals("T3", orderedCalls.get(0));
assertEquals("T2", orderedCalls.get(1));
private void awaitQueueSize(int expected) throws InterruptedException {
long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(2);
while (queueLock.queue.size() != expected && System.nanoTime() < deadline) {
Thread.sleep(1);
}
assertEquals(expected, queueLock.queue.size(), "Requests did not enter the queue in time");
}

@Test
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,19 +17,24 @@
import org.patinanetwork.codebloom.common.db.models.question.QuestionDifficulty;
import org.patinanetwork.codebloom.common.db.repos.job.JobRepository;
import org.patinanetwork.codebloom.common.db.repos.question.QuestionRepository;
import org.patinanetwork.codebloom.common.leetcode.throttled.ThrottledLeetcodeClient;
import org.patinanetwork.codebloom.common.time.StandardizedOffsetDateTime;
import org.patinanetwork.codebloom.config.NoJdaRequired;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.test.annotation.DirtiesContext;
import org.springframework.test.context.ActiveProfiles;
import org.springframework.test.context.bean.override.mockito.MockitoBean;

@SpringBootTest
@SpringBootTest(properties = {"codebloom.scheduling.enabled=false", "codebloom.notify.enabled=false"})
@ActiveProfiles({"ci", "thread"})
@TestInstance(TestInstance.Lifecycle.PER_CLASS)
@DirtiesContext(classMode = DirtiesContext.ClassMode.AFTER_CLASS)
public class LeetcodeQuestionProcessServiceTest extends NoJdaRequired {

@MockitoBean
private ThrottledLeetcodeClient leetcodeClient;

private final JobRepository jobRepository;
private final LeetcodeQuestionProcessService service;
private final QuestionRepository questionRepository;
Expand Down Expand Up @@ -177,7 +182,9 @@ void jobStatusTransitionValid() {

@Test
void drainQueueValid() {
service.drainQueue();
service.drainQueue().join();
assertEquals(
JobStatus.COMPLETE, jobRepository.findJobById(testJob.getId()).getStatus());
}

// TODO: (TAN-32) re-enable
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@

import java.util.List;
import java.util.Optional;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.TimeUnit;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.params.ParameterizedTest;
Expand Down Expand Up @@ -65,6 +67,48 @@ private void runQueue() {
.join();
}

@Test
void independentServiceCanDrainWhileAnotherInstanceIsRunning() {
var otherJobs = mock(JobRepository.class);
when(otherJobs.findIncompleteJobs(10)).thenReturn(List.of());
var otherService = new LeetcodeQuestionProcessService(otherJobs, client, questions, bank);
when(jobs.findIncompleteJobs(10)).thenAnswer(invocation -> {
CompletableFuture.runAsync(() -> otherService.drainQueue().join()).get(5, TimeUnit.SECONDS);
return List.of();
});

runQueue();

verify(otherJobs).findIncompleteJobs(10);
}

@Test
void sameServiceSkipsConcurrentDrain() {
var service = new LeetcodeQuestionProcessService(jobs, client, questions, bank);
when(jobs.findIncompleteJobs(10)).thenAnswer(invocation -> {
CompletableFuture.runAsync(() -> service.drainQueue().join()).get(5, TimeUnit.SECONDS);
return List.of();
});

service.drainQueue().join();

verify(jobs).findIncompleteJobs(10);
}

@Test
void failedDrainReleasesInstanceLock() {
var service = new LeetcodeQuestionProcessService(jobs, client, questions, bank);
when(jobs.findIncompleteJobs(10))
.thenThrow(new IllegalStateException("Database unavailable"))
.thenReturn(List.of());

assertThrows(IllegalStateException.class, () -> service.drainQueue().join());
assertDoesNotThrow(() ->
CompletableFuture.runAsync(() -> service.drainQueue().join()).get(5, TimeUnit.SECONDS));

verify(jobs, times(2)).findIncompleteJobs(10);
}

@ParameterizedTest
@NullAndEmptySource
@ValueSource(strings = {" ", "\t\n"})
Expand Down
Loading