103 lines
2.7 KiB
Python
103 lines
2.7 KiB
Python
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
|