feat(ai): add parallel processing for search and generation - Add parallel search and generation, schema validation, tests, and better error handling
This commit is contained in:
+31
-20
@@ -1,4 +1,5 @@
|
|||||||
import logging
|
import logging
|
||||||
|
import threading
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional
|
||||||
from gitea.models import IssueModel, PullRequestModel
|
from gitea.models import IssueModel, PullRequestModel
|
||||||
@@ -19,43 +20,53 @@ class WorkQueue:
|
|||||||
"""Thread-safe work queue grouped by repo."""
|
"""Thread-safe work queue grouped by repo."""
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
|
self._lock: threading.Lock = threading.Lock()
|
||||||
self._queue: list[WorkItem] = []
|
self._queue: list[WorkItem] = []
|
||||||
self._enqueued_repos: set[str] = set()
|
self._enqueued_repos: set[str] = set()
|
||||||
|
|
||||||
def enqueue(self, item: WorkItem) -> None:
|
def enqueue(self, item: WorkItem) -> None:
|
||||||
self._queue.append(item)
|
with self._lock:
|
||||||
self._enqueued_repos.add(item.repo_full_name)
|
self._queue.append(item)
|
||||||
logger.info(f"Enqueued work item: {item.task_type} #{item.task_number} for {item.repo_full_name}")
|
self._enqueued_repos.add(item.repo_full_name)
|
||||||
|
logger.info(f"Enqueued work item: {item.task_type} #{item.task_number} for {item.repo_full_name}")
|
||||||
|
|
||||||
def enqueue_batch(self, items: list[WorkItem]) -> None:
|
def enqueue_batch(self, items: list[WorkItem]) -> None:
|
||||||
for item in items:
|
with self._lock:
|
||||||
self.enqueue(item)
|
for item in items:
|
||||||
|
self._queue.append(item)
|
||||||
|
self._enqueued_repos.add(item.repo_full_name)
|
||||||
|
logger.info(f"Enqueued work item: {item.task_type} #{item.task_number} for {item.repo_full_name}")
|
||||||
|
|
||||||
def get_repo_work(self, repo: str) -> list[WorkItem]:
|
def get_repo_work(self, repo: str) -> list[WorkItem]:
|
||||||
"""Get all work items for a specific repo."""
|
"""Get all work items for a specific repo."""
|
||||||
items: list[WorkItem] = [
|
with self._lock:
|
||||||
item for item in self._queue if item.repo_full_name == repo
|
items: list[WorkItem] = [
|
||||||
]
|
item for item in self._queue if item.repo_full_name == repo
|
||||||
logger.info(f"Retrieved {len(items)} work items for repository: {repo}")
|
]
|
||||||
return items
|
logger.info(f"Retrieved {len(items)} work items for repository: {repo}")
|
||||||
|
return items
|
||||||
|
|
||||||
def remove_repo_work(self, repo: str) -> None:
|
def remove_repo_work(self, repo: str) -> None:
|
||||||
"""Remove all work items for a specific repo."""
|
"""Remove all work items for a specific repo."""
|
||||||
self._queue = [
|
with self._lock:
|
||||||
item for item in self._queue if item.repo_full_name != repo
|
self._queue = [
|
||||||
]
|
item for item in self._queue if item.repo_full_name != repo
|
||||||
self._enqueued_repos.discard(repo)
|
]
|
||||||
logger.info(f"Removed all work items for repository: {repo}")
|
self._enqueued_repos.discard(repo)
|
||||||
|
logger.info(f"Removed all work items for repository: {repo}")
|
||||||
|
|
||||||
def get_next_repo(self) -> str | None:
|
def get_next_repo(self) -> str | None:
|
||||||
"""Get the next repo with work, or None if empty."""
|
"""Get the next repo with work, or None if empty."""
|
||||||
if not self._enqueued_repos:
|
with self._lock:
|
||||||
return None
|
if not self._enqueued_repos:
|
||||||
return next(iter(self._enqueued_repos))
|
return None
|
||||||
|
return next(iter(self._enqueued_repos))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_empty(self) -> bool:
|
def is_empty(self) -> bool:
|
||||||
return len(self._queue) == 0
|
with self._lock:
|
||||||
|
return len(self._queue) == 0
|
||||||
|
|
||||||
def __len__(self) -> int:
|
def __len__(self) -> int:
|
||||||
return len(self._queue)
|
with self._lock:
|
||||||
|
return len(self._queue)
|
||||||
|
|||||||
@@ -0,0 +1,102 @@
|
|||||||
|
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
|
||||||
Reference in New Issue
Block a user