feat: implement Gitea notification polling and state persistence #1
+112
-37
@@ -1,13 +1,16 @@
|
|||||||
"""Top-level coordinator: polls Gitea, queues work, dispatches to agent."""
|
"""Top-level coordinator: polls Gitea unread notifications, queues work, dispatches to agent."""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
|
import datetime
|
||||||
|
import json
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any, Optional
|
||||||
|
|
||||||
from core.queue import WorkQueue, WorkItem
|
from core.queue import WorkQueue, WorkItem
|
||||||
from core.dispatcher import AgentDispatcher
|
from core.dispatcher import AgentDispatcher
|
||||||
from gitea.client import GiteaClient
|
from gitea.client import GiteaClient
|
||||||
from gitea.models import IssueModel, PullRequestModel
|
from gitea.models import IssueModel, PullRequestModel, RepositoryModel
|
||||||
from gitea.tools.gitea_tools import GiteaTools
|
from gitea.tools.gitea_tools import GiteaTools
|
||||||
from gitea.config import AGENT_MODEL_ID, AGENT_MAX_RETRIES
|
from gitea.config import AGENT_MODEL_ID, AGENT_MAX_RETRIES
|
||||||
from gitea.workspace import WorkspaceManager
|
from gitea.workspace import WorkspaceManager
|
||||||
@@ -16,7 +19,7 @@ logger: logging.Logger = logging.getLogger("agent-orchestrator")
|
|||||||
|
|
||||||
|
|
||||||
class AgentOrchestrator:
|
class AgentOrchestrator:
|
||||||
"""Top-level coordinator: polls Gitea, queues work, dispatches to agent."""
|
"""Top-level coordinator: polls Gitea notifications, queues work, dispatches to agent."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -31,51 +34,120 @@ class AgentOrchestrator:
|
|||||||
self._work_queue = WorkQueue()
|
self._work_queue = WorkQueue()
|
||||||
self._dispatcher = AgentDispatcher(client, tools, model_name, max_retries)
|
self._dispatcher = AgentDispatcher(client, tools, model_name, max_retries)
|
||||||
|
|
||||||
|
def _get_state_file_path(self) -> Path:
|
||||||
|
"""Get the path to the persistent state file."""
|
||||||
|
return Path(__file__).parent.parent / "agent_state.json"
|
||||||
|
|
||||||
|
def _read_last_checked(self) -> Optional[str]:
|
||||||
|
"""Read the last checked timestamp from persistent storage."""
|
||||||
|
state_file = self._get_state_file_path()
|
||||||
|
if state_file.exists():
|
||||||
|
try:
|
||||||
|
with open(state_file, "r") as f:
|
||||||
|
data = json.load(f)
|
||||||
|
return data.get("last_checked")
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Error reading agent_state.json: {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _write_last_checked(self, timestamp: str) -> None:
|
||||||
|
"""Write the last checked timestamp to persistent storage."""
|
||||||
|
state_file = self._get_state_file_path()
|
||||||
|
try:
|
||||||
|
with open(state_file, "w") as f:
|
||||||
|
json.dump({"last_checked": timestamp}, f)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error writing agent_state.json: {e}")
|
||||||
|
|
||||||
async def poll_and_dispatch(self) -> None:
|
async def poll_and_dispatch(self) -> None:
|
||||||
"""Poll Gitea for tasks, enqueue them, and dispatch to agent."""
|
"""Poll Gitea unread notifications, enqueue them, and dispatch to agent."""
|
||||||
issues: list[IssueModel] = self._client.list_assigned_issues()
|
last_checked = self._read_last_checked()
|
||||||
prs: list[PullRequestModel] = self._client.list_assigned_pull_requests()
|
logger.info(f"Polling unread notifications since: {last_checked or 'beginning'}")
|
||||||
|
|
||||||
if issues:
|
notifications = self._client.list_unread_notifications(since=last_checked)
|
||||||
logger.info(f"Found {len(issues)} assigned issues")
|
|
||||||
self._enqueue_tasks("issue", issues)
|
|
||||||
else:
|
|
||||||
logger.info("No assigned issues found.")
|
|
||||||
|
|
||||||
if prs:
|
if not notifications:
|
||||||
logger.info(f"Found {len(prs)} assigned PRs")
|
logger.info("No new notifications found.")
|
||||||
self._enqueue_tasks("pr", prs)
|
return
|
||||||
else:
|
|
||||||
logger.info("No assigned PRs found.")
|
|
||||||
|
|
||||||
|
logger.info(f"Retrieved {len(notifications)} unread notifications.")
|
||||||
|
|
||||||
|
# Filter and enqueue tasks from notifications
|
||||||
|
latest_timestamp = last_checked
|
||||||
|
for n in notifications:
|
||||||
|
notification_id = n.get("id")
|
||||||
|
subject = n.get("subject") or {}
|
||||||
|
subj_type = subject.get("type", "").lower()
|
||||||
|
subj_url = subject.get("url", "")
|
||||||
|
updated_at = n.get("updated_at")
|
||||||
|
|
||||||
|
# Track latest updated_at to advance checkpoint
|
||||||
|
if updated_at and (not latest_timestamp or updated_at > latest_timestamp):
|
||||||
|
latest_timestamp = updated_at
|
||||||
|
|
||||||
|
repo_info = n.get("repository") or {}
|
||||||
|
repo_full_name = repo_info.get("full_name", "")
|
||||||
|
if not repo_full_name or not subj_url:
|
||||||
|
continue
|
||||||
|
|
||||||
|
owner, repo_name = repo_full_name.split("/")
|
||||||
|
try:
|
||||||
|
task_number = int(subj_url.rstrip("/").split("/")[-1])
|
||||||
|
except (ValueError, IndexError):
|
||||||
|
logger.warning(f"Could not parse task number from subject URL: {subj_url}")
|
||||||
|
continue
|
||||||
|
|
||||||
|
if subj_type == "issue":
|
||||||
|
try:
|
||||||
|
issue = self._client.get_issue(owner, repo_name, task_number)
|
||||||
|
if issue.repository is None:
|
||||||
|
issue = issue.model_copy(update={"repository": RepositoryModel(**repo_info)})
|
||||||
|
|
||||||
|
item = WorkItem(
|
||||||
|
repo_full_name=repo_full_name,
|
||||||
|
task_type="issue",
|
||||||
|
task_number=task_number,
|
||||||
|
task_info=issue,
|
||||||
|
notification_id=notification_id,
|
||||||
|
priority=0
|
||||||
|
)
|
||||||
|
self._work_queue.enqueue(item)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to fetch issue #{task_number} for notification: {e}")
|
||||||
|
elif subj_type in ("pull", "pullrequest"):
|
||||||
|
try:
|
||||||
|
pr = self._client.get_pull_request(owner, repo_name, task_number)
|
||||||
|
if pr.repository is None:
|
||||||
|
pr = pr.model_copy(update={"repository": RepositoryModel(**repo_info)})
|
||||||
|
|
||||||
|
item = WorkItem(
|
||||||
|
repo_full_name=repo_full_name,
|
||||||
|
task_type="pr",
|
||||||
|
task_number=task_number,
|
||||||
|
task_info=pr,
|
||||||
|
notification_id=notification_id,
|
||||||
|
priority=0
|
||||||
|
)
|
||||||
|
self._work_queue.enqueue(item)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to fetch PR #{task_number} for notification: {e}")
|
||||||
|
|
||||||
|
# Process enqueued work
|
||||||
if not self._work_queue.is_empty:
|
if not self._work_queue.is_empty:
|
||||||
await self._process_work()
|
await self._process_work()
|
||||||
|
|
||||||
def _enqueue_tasks(self, task_type: str, tasks: list[IssueModel] | list[PullRequestModel]) -> None:
|
# Update last checked timestamp checkpoint
|
||||||
for task in tasks:
|
if latest_timestamp:
|
||||||
repo_full_name: str | None = task.repository.full_name if task.repository else None
|
self._write_last_checked(latest_timestamp)
|
||||||
task_number: int = task.number
|
|
||||||
if not repo_full_name or not task_number:
|
|
||||||
continue
|
|
||||||
|
|
||||||
item = WorkItem(
|
|
||||||
repo_full_name=repo_full_name,
|
|
||||||
task_type=task_type,
|
|
||||||
task_number=task_number,
|
|
||||||
task_info=task,
|
|
||||||
priority=0,
|
|
||||||
)
|
|
||||||
self._work_queue.enqueue(item)
|
|
||||||
logger.info(f"Enqueued {task_type} #{task_number} from {repo_full_name}")
|
|
||||||
|
|
||||||
async def _process_work(self) -> None:
|
async def _process_work(self) -> None:
|
||||||
"""Process all queued work, repo by repo."""
|
"""Process all queued work, repo by repo, then mark notifications as read."""
|
||||||
while not self._work_queue.is_empty:
|
while not self._work_queue.is_empty:
|
||||||
repo: str | None = self._work_queue.get_next_repo()
|
repo = self._work_queue.get_next_repo()
|
||||||
if not repo:
|
if not repo:
|
||||||
break
|
break
|
||||||
|
|
||||||
work_items: list[WorkItem] = self._work_queue.get_repo_work(repo)
|
work_items = self._work_queue.get_repo_work(repo)
|
||||||
self._work_queue.remove_repo_work(repo)
|
self._work_queue.remove_repo_work(repo)
|
||||||
|
|
||||||
# Ensure the workspace repository is cloned and sanitized
|
# Ensure the workspace repository is cloned and sanitized
|
||||||
@@ -89,8 +161,11 @@ class AgentOrchestrator:
|
|||||||
logger.info(f"Sanitized existing repo at {repo_path}")
|
logger.info(f"Sanitized existing repo at {repo_path}")
|
||||||
|
|
||||||
logger.info(f"Dispatching {len(work_items)} tasks for {repo}")
|
logger.info(f"Dispatching {len(work_items)} tasks for {repo}")
|
||||||
results: list[str] = await self._dispatcher.dispatch(repo, work_items)
|
results = await self._dispatcher.dispatch(repo, work_items)
|
||||||
|
|
||||||
for i, result in enumerate(results):
|
for i, result in enumerate(results):
|
||||||
item = work_items[i]
|
item = work_items[i]
|
||||||
logger.info(f"Completed {item.task_type} #{item.task_number}: {result[:200]}")
|
logger.info(f"Completed {item.task_type} #{item.task_number}: {result[:200]}")
|
||||||
|
if item.notification_id is not None:
|
||||||
|
self._client.mark_notification_as_read(item.notification_id)
|
||||||
|
logger.info(f"Marked Gitea notification thread {item.notification_id} as read.")
|
||||||
|
|||||||
+2
-1
@@ -1,6 +1,6 @@
|
|||||||
import logging
|
import logging
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
from typing import Any
|
from typing import Any, Optional
|
||||||
from gitea.models import IssueModel, PullRequestModel
|
from gitea.models import IssueModel, PullRequestModel
|
||||||
|
|
||||||
logger: logging.Logger = logging.getLogger("work-queue")
|
logger: logging.Logger = logging.getLogger("work-queue")
|
||||||
@@ -11,6 +11,7 @@ class WorkItem(BaseModel):
|
|||||||
task_type: str # 'issue' or 'pr'
|
task_type: str # 'issue' or 'pr'
|
||||||
task_number: int
|
task_number: int
|
||||||
task_info: IssueModel | PullRequestModel
|
task_info: IssueModel | PullRequestModel
|
||||||
|
notification_id: Optional[int] = None
|
||||||
priority: int = 0
|
priority: int = 0
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+34
-1
@@ -1,7 +1,7 @@
|
|||||||
import httpx
|
import httpx
|
||||||
import json
|
import json
|
||||||
import base64
|
import base64
|
||||||
from typing import Any
|
from typing import Any, Optional
|
||||||
from .config import GITEA_URL, GITEA_TOKEN
|
from .config import GITEA_URL, GITEA_TOKEN
|
||||||
from .models import (
|
from .models import (
|
||||||
UserModel,
|
UserModel,
|
||||||
@@ -396,3 +396,36 @@ class GiteaClient(IssuesClient, PullRequestsClient, FilesClient, RefsClient, Rep
|
|||||||
if isinstance(data, list):
|
if isinstance(data, list):
|
||||||
return [item.get("content", "") for item in data if item.get("type") == "file"]
|
return [item.get("content", "") for item in data if item.get("type") == "file"]
|
||||||
return base64.b64decode(data.get("content", "")).decode() if data.get("content") else ""
|
return base64.b64decode(data.get("content", "")).decode() if data.get("content") else ""
|
||||||
|
|
||||||
|
def list_unread_notifications(self, since: Optional[str] = None) -> list[dict[str, Any]]:
|
||||||
|
try:
|
||||||
|
with httpx.Client() as client:
|
||||||
|
url = f"{self.base_url}/api/v1/notifications?all=false"
|
||||||
|
if since:
|
||||||
|
url += f"&since={since}"
|
||||||
|
response = client.get(url, headers=self.headers)
|
||||||
|
response.raise_for_status()
|
||||||
|
notifications: list[dict[str, Any]] = response.json()
|
||||||
|
|
||||||
|
result: list[dict[str, Any]] = []
|
||||||
|
for n in notifications:
|
||||||
|
repo_info = n.get("repository") or {}
|
||||||
|
owner_info = repo_info.get("owner") or {}
|
||||||
|
owner_login = owner_info.get("login", "")
|
||||||
|
if owner_login == "meeks":
|
||||||
|
result.append(n)
|
||||||
|
return result
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Error listing unread notifications: {e}")
|
||||||
|
return []
|
||||||
|
|
||||||
|
def mark_notification_as_read(self, thread_id: int) -> bool:
|
||||||
|
try:
|
||||||
|
with httpx.Client() as client:
|
||||||
|
url = f"{self.base_url}/api/v1/notifications/threads/{thread_id}"
|
||||||
|
response = client.patch(url, headers=self.headers)
|
||||||
|
response.raise_for_status()
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Error marking notification thread {thread_id} as read: {e}")
|
||||||
|
return False
|
||||||
|
|||||||
@@ -0,0 +1,133 @@
|
|||||||
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
import pytest
|
||||||
|
from unittest.mock import MagicMock, AsyncMock, patch
|
||||||
|
|
||||||
|
from core.orchestrator import AgentOrchestrator
|
||||||
|
from gitea.client import GiteaClient
|
||||||
|
from gitea.tools.gitea_tools import GiteaTools
|
||||||
|
from gitea.models import IssueModel, PullRequestModel, RepositoryModel
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.anyio
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def temp_state_file(tmp_path: Path) -> Path:
|
||||||
|
"""Fixture to mock state file path."""
|
||||||
|
state_file = tmp_path / "agent_state.json"
|
||||||
|
return state_file
|
||||||
|
|
||||||
|
|
||||||
|
@patch("core.orchestrator.AgentOrchestrator._get_state_file_path")
|
||||||
|
@patch("core.orchestrator.AgentDispatcher")
|
||||||
|
@patch("core.orchestrator.WorkspaceManager")
|
||||||
|
async def test_poll_and_dispatch_no_notifications(
|
||||||
|
mock_workspace_class: MagicMock,
|
||||||
|
mock_dispatcher_class: MagicMock,
|
||||||
|
mock_get_path: MagicMock,
|
||||||
|
temp_state_file: Path
|
||||||
|
) -> None:
|
||||||
|
mock_get_path.return_value = temp_state_file
|
||||||
|
mock_client = MagicMock(spec=GiteaClient)
|
||||||
|
mock_tools = MagicMock(spec=GiteaTools)
|
||||||
|
|
||||||
|
# Return no notifications
|
||||||
|
mock_client.list_unread_notifications.return_value = []
|
||||||
|
|
||||||
|
orchestrator = AgentOrchestrator(mock_client, mock_tools)
|
||||||
|
await orchestrator.poll_and_dispatch()
|
||||||
|
|
||||||
|
mock_client.list_unread_notifications.assert_called_once_with(since=None)
|
||||||
|
assert not temp_state_file.exists()
|
||||||
|
|
||||||
|
|
||||||
|
@patch("core.orchestrator.AgentOrchestrator._get_state_file_path")
|
||||||
|
@patch("core.orchestrator.AgentDispatcher")
|
||||||
|
@patch("core.orchestrator.WorkspaceManager")
|
||||||
|
async def test_poll_and_dispatch_with_notifications(
|
||||||
|
mock_workspace_class: MagicMock,
|
||||||
|
mock_dispatcher_class: MagicMock,
|
||||||
|
mock_get_path: MagicMock,
|
||||||
|
temp_state_file: Path
|
||||||
|
) -> None:
|
||||||
|
mock_get_path.return_value = temp_state_file
|
||||||
|
mock_client = MagicMock(spec=GiteaClient)
|
||||||
|
mock_tools = MagicMock(spec=GiteaTools)
|
||||||
|
|
||||||
|
# Set up mock Gitea notifications
|
||||||
|
notifications = [
|
||||||
|
{
|
||||||
|
"id": 101,
|
||||||
|
"updated_at": "2026-06-30T10:00:00Z",
|
||||||
|
"subject": {
|
||||||
|
"type": "issue",
|
||||||
|
"url": "https://gitea.meeks.freeddns.org/api/v1/repos/meeks/repo1/issues/42"
|
||||||
|
},
|
||||||
|
"repository": {
|
||||||
|
"name": "repo1",
|
||||||
|
"full_name": "meeks/repo1",
|
||||||
|
"owner": {"login": "meeks"}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 102,
|
||||||
|
"updated_at": "2026-06-30T11:00:00Z",
|
||||||
|
"subject": {
|
||||||
|
"type": "pull",
|
||||||
|
"url": "https://gitea.meeks.freeddns.org/api/v1/repos/meeks/repo1/pulls/10"
|
||||||
|
},
|
||||||
|
"repository": {
|
||||||
|
"name": "repo1",
|
||||||
|
"full_name": "meeks/repo1",
|
||||||
|
"owner": {"login": "meeks"}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
]
|
||||||
|
mock_client.list_unread_notifications.return_value = notifications
|
||||||
|
|
||||||
|
# Mock issue and PR get methods
|
||||||
|
issue_model = IssueModel(number=42, title="Bug issue", repository=RepositoryModel(name="repo1", full_name="meeks/repo1"))
|
||||||
|
pr_model = PullRequestModel(number=10, title="Fix PR", repository=RepositoryModel(name="repo1", full_name="meeks/repo1"))
|
||||||
|
mock_client.get_issue.return_value = issue_model
|
||||||
|
mock_client.get_pull_request.return_value = pr_model
|
||||||
|
|
||||||
|
# Mock dispatcher and workspace path
|
||||||
|
mock_dispatcher_instance = MagicMock()
|
||||||
|
mock_dispatcher_instance.dispatch = AsyncMock(return_value=["Issue comment posted", "PR verified"])
|
||||||
|
mock_dispatcher_class.return_value = mock_dispatcher_instance
|
||||||
|
|
||||||
|
mock_workspace_instance = MagicMock()
|
||||||
|
mock_workspace_instance.get_repo_path.return_value.exists.return_value = True
|
||||||
|
mock_workspace_class.return_value = mock_workspace_instance
|
||||||
|
|
||||||
|
# Create orchestrator and poll
|
||||||
|
orchestrator = AgentOrchestrator(mock_client, mock_tools)
|
||||||
|
await orchestrator.poll_and_dispatch()
|
||||||
|
|
||||||
|
# Assert notifications were checked with None (first execution)
|
||||||
|
mock_client.list_unread_notifications.assert_called_once_with(since=None)
|
||||||
|
|
||||||
|
# Assert issue and PR details were fetched
|
||||||
|
mock_client.get_issue.assert_called_once_with("meeks", "repo1", 42)
|
||||||
|
mock_client.get_pull_request.assert_called_once_with("meeks", "repo1", 10)
|
||||||
|
|
||||||
|
# Assert work was processed by dispatcher
|
||||||
|
mock_dispatcher_instance.dispatch.assert_called_once()
|
||||||
|
work_items = mock_dispatcher_instance.dispatch.call_args[0][1]
|
||||||
|
assert len(work_items) == 2
|
||||||
|
assert work_items[0].task_number == 42
|
||||||
|
assert work_items[0].notification_id == 101
|
||||||
|
assert work_items[1].task_number == 10
|
||||||
|
assert work_items[1].notification_id == 102
|
||||||
|
|
||||||
|
# Assert notifications were marked as read
|
||||||
|
mock_client.mark_notification_as_read.assert_any_call(101)
|
||||||
|
mock_client.mark_notification_as_read.assert_any_call(102)
|
||||||
|
assert mock_client.mark_notification_as_read.call_count == 2
|
||||||
|
|
||||||
|
# Assert checkpoint date was persisted
|
||||||
|
assert temp_state_file.exists()
|
||||||
|
with open(temp_state_file, "r") as f:
|
||||||
|
state = json.load(f)
|
||||||
|
# Checkpoint should match latest updated_at
|
||||||
|
assert state["last_checked"] == "2026-06-30T11:00:00Z"
|
||||||
Reference in New Issue
Block a user