import threading from core.queue import WorkQueue, WorkItem from gitea.models import IssueModel def test_work_queue_basic_operations() -> None: queue: WorkQueue = WorkQueue() assert queue.is_empty assert len(queue) == 0 assert queue.get_next_repo() is None item1 = WorkItem( repo_full_name="meeks/repo1", task_type="issue", task_number=1, task_info=IssueModel(number=1), ) item2 = WorkItem( repo_full_name="meeks/repo1", task_type="issue", task_number=2, task_info=IssueModel(number=2), ) item3 = WorkItem( repo_full_name="meeks/repo2", task_type="issue", task_number=3, task_info=IssueModel(number=3), ) queue.enqueue(item1) assert not queue.is_empty assert len(queue) == 1 assert queue.get_next_repo() == "meeks/repo1" queue.enqueue_batch([item2, item3]) assert len(queue) == 3 # get_repo_work repo1_work = queue.get_repo_work("meeks/repo1") assert len(repo1_work) == 2 assert repo1_work[0].task_number == 1 assert repo1_work[1].task_number == 2 # remove_repo_work queue.remove_repo_work("meeks/repo1") assert len(queue) == 1 assert queue.get_next_repo() == "meeks/repo2" queue.remove_repo_work("meeks/repo2") assert queue.is_empty assert len(queue) == 0 assert queue.get_next_repo() is None def test_work_queue_thread_safety() -> None: queue: WorkQueue = WorkQueue() num_threads: int = 10 items_per_thread: int = 100 barrier = threading.Barrier(num_threads) def worker(thread_idx: int) -> None: barrier.wait() # synchronize start for i in range(items_per_thread): item = WorkItem( repo_full_name=f"meeks/repo_{thread_idx}", task_type="issue", task_number=i, task_info=IssueModel(number=i), ) queue.enqueue(item) threads: list[threading.Thread] = [] for idx in range(num_threads): t = threading.Thread(target=worker, args=(idx,)) threads.append(t) t.start() for t in threads: t.join() # Verify that all items are enqueued assert len(queue) == num_threads * items_per_thread # Concurrently remove repo work barrier_remove = threading.Barrier(num_threads) def remover(thread_idx: int) -> None: barrier_remove.wait() queue.remove_repo_work(f"meeks/repo_{thread_idx}") remove_threads: list[threading.Thread] = [] for idx in range(num_threads): t = threading.Thread(target=remover, args=(idx,)) remove_threads.append(t) t.start() for t in remove_threads: t.join() assert queue.is_empty assert len(queue) == 0