import asyncio import contextlib import logging from collections.abc import Awaitable, Callable from datetime import UTC, datetime, timedelta from typing import Any from uuid import uuid4 from sqlalchemy import select, update from app.core.config import Settings, get_settings from app.core.contracts import AgentRequest, AgentRequestMetadata, AgentResult, RequestContext from app.core.errors import AgentError, RecoverableAgentError, RunLeaseLostError from app.infrastructure.db import SessionFactory from app.infrastructure.neo4j_graph_driver import Neo4jGraphDriver from app.model.audit import InteractionAudit from app.model.conversation import ConversationMessage from app.model.platform import AgentRun, DomainEventOutbox, RequestIdempotency from app.repository.agent_run_repository import AgentRunRepository from app.service.agent.bootstrap import get_agent_factory from app.service.agent.executor import AgentExecutor from app.service.agent.factory import AgentFactory from app.service.agent_persistence_service import AgentPersistenceService from app.service.identity_service import IdentityService from app.service.memory_service import MemoryService from app.service.portfolio_projection_scheduler import PortfolioProjectionScheduler from app.service.relationship_service import RelationshipService from app.worker.outbox_worker import OutboxWorker from app.worker.portfolio_graph_projection_worker import PortfolioGraphProjectionWorker from app.worker.product_governance_monitor_worker import ProductGovernanceMonitorWorker from app.worker.product_history_sync_worker import ProductHistorySyncWorker from app.worker.product_market_quote_sync_worker import ProductMarketQuoteSyncWorker from app.worker.product_metric_refresh_worker import ProductMetricRefreshWorker logger = logging.getLogger(__name__) class WorkerRuntime: def __init__( self, factory: AgentFactory | None = None, settings: Settings | None = None, resolve_identity: Callable[[RequestContext], Awaitable[RequestContext]] | None = None, ) -> None: self.factory = factory if factory is not None else get_agent_factory() self.settings = settings or get_settings() self.resolve_identity = resolve_identity or IdentityService().resolve self.portfolio_projection_scheduler = PortfolioProjectionScheduler() self.portfolio_projection_worker = PortfolioGraphProjectionWorker( RelationshipService(Neo4jGraphDriver(self.settings)) ) self.product_history_sync_worker = ProductHistorySyncWorker() self.product_governance_monitor_worker = ProductGovernanceMonitorWorker() self.product_market_quote_sync_worker = ProductMarketQuoteSyncWorker() self.product_metric_refresh_worker = ProductMetricRefreshWorker() async def dispatch_one(self, *, run_id: str | None = None) -> bool: # Outbox acknowledges a durable SQL queue entry, not an in-memory task. async with SessionFactory() as session: async def dispatch(payload: dict[str, Any]) -> None: run = await AgentRunRepository(session).get(str(payload["run_id"])) if run is None: raise ValueError("run not found") async def project_portfolio(payload: dict[str, Any]) -> None: customer_id = payload.get("customer_id") if type(customer_id) is not int or customer_id <= 0: raise ValueError("invalid portfolio projection customer") await self.portfolio_projection_worker.project_customer(customer_id) return await OutboxWorker(session, { "agent.run_requested": dispatch, "portfolio.projection_requested": project_portfolio, }).publish_one(aggregate_id=run_id) async def run_once(self) -> bool: governance_changes = await self.product_governance_monitor_worker.refresh_once() quotes_refreshed = await self.product_market_quote_sync_worker.refresh_once() history_refreshed = await self.product_history_sync_worker.refresh_once() metrics_refreshed = await self.product_metric_refresh_worker.refresh_once() projection_enqueued = await self.portfolio_projection_scheduler.enqueue_changes() dispatched = await self.dispatch_one() now = datetime.now(UTC).replace(tzinfo=None) async with SessionFactory() as session: run_id = await session.scalar(select(AgentRun.run_id).where( AgentRun.status.in_(("queued", "running", "cancel_requested")), (AgentRun.locked_until.is_(None) | (AgentRun.locked_until < now)), ).order_by(AgentRun.created_at).limit(1)) if run_id is None: return any(( dispatched, bool(projection_enqueued), bool(metrics_refreshed), bool(history_refreshed), bool(quotes_refreshed), bool(governance_changes), )) return ( await self.execute(run_id) or dispatched or bool(projection_enqueued) or bool(metrics_refreshed) or bool(history_refreshed) ) async def execute(self, run_id: str) -> bool: # A new fencing token for each claim also fences restarts of the same process. worker_id = str(uuid4()) async with SessionFactory() as session, session.begin(): run = await session.scalar(select(AgentRun).where( AgentRun.run_id == run_id).with_for_update()) if run is None: return False if run.status == "cancel_requested": run.status = "cancelled" run.completed_at = datetime.now(UTC).replace(tzinfo=None) run.locked_until = None return True claimed = await AgentRunRepository(session).claim( run_id, worker_id, self.settings.worker_lease_seconds) if not claimed: return False task = asyncio.create_task(self._execute_claimed(run_id, worker_id)) heartbeat = asyncio.create_task(self._heartbeat(run_id, worker_id, task)) try: await task except asyncio.CancelledError: await self._failure(run_id, worker_id, "RUN_INTERRUPTED", retryable=True) # Runtime cancellation is shutdown; a lost lease only cancels the child. current = asyncio.current_task() if current is not None and current.cancelling(): raise except RunLeaseLostError: await self._failure(run_id, worker_id, "RUN_LEASE_LOST", retryable=True) except Exception as exc: code = exc.code if isinstance(exc, AgentError) else "AGENT_INTERNAL_ERROR" await self._failure(run_id, worker_id, code, retryable=isinstance(exc, RecoverableAgentError)) logger.warning("run failed run_id=%s error_type=%s", run_id, type(exc).__name__) finally: heartbeat.cancel() with contextlib.suppress(asyncio.CancelledError): await heartbeat return True async def _heartbeat( self, run_id: str, worker_id: str, task: asyncio.Task[None] ) -> None: try: while True: await asyncio.sleep(self.settings.worker_lease_seconds / 3) async with SessionFactory() as session, session.begin(): renewed = await AgentRunRepository(session).renew( run_id, worker_id, self.settings.worker_lease_seconds) if not renewed: task.cancel() return except Exception: task.cancel() logger.warning("lease renewal failed run_id=%s", run_id) async def _execute_claimed(self, run_id: str, worker_id: str) -> None: async with SessionFactory() as session: run = await AgentRunRepository(session).get(run_id) if run is None: raise ValueError("run not found") message = await session.get(ConversationMessage, run.request_message_id) idem = await session.get(RequestIdempotency, run.idempotency_id) event = await session.scalar(select(DomainEventOutbox).where( DomainEventOutbox.aggregate_id == run_id, DomainEventOutbox.event_type == "agent.run_requested").limit(1)) if message is None or idem is None: raise ValueError("run input missing") metadata = event.payload.get("metadata", {}) if event else {} request = AgentRequest( agent_type=run.agent_type, message=message.content, session_id=run.session_id, idempotency_key=idem.idempotency_key, metadata=AgentRequestMetadata.model_validate(metadata), ) identity = RequestContext(user_id=str(run.user_id), trace_id=run.trace_id) # Re-check account and permissions at execution time, including delayed jobs. context = await self.resolve_identity(identity) result: AgentResult | None = None async for event_data in AgentExecutor(self.factory).execute( request.agent_type, request, context, run_id ): if event_data.event_type == "done": result = AgentResult.model_validate(event_data.payload["result"]) if result is None: raise ValueError("Agent produced no terminal result") async with SessionFactory() as session: await AgentPersistenceService(session).complete_run( run_id, result, worker_id=worker_id, memory_extraction_requested=MemoryService.should_extract_memory( conversation_content=request.message, role="user"), ) async def _failure( self, run_id: str, worker_id: str, error_code: str, *, retryable: bool ) -> None: now = datetime.now(UTC).replace(tzinfo=None) async with SessionFactory() as session, session.begin(): run = await session.scalar(select(AgentRun).where( AgentRun.run_id == run_id).with_for_update()) if run is None or run.worker_id != worker_id: return if run.status not in {"running", "cancel_requested"}: return if run.status == "cancel_requested": run.status = "cancelled" elif retryable and run.attempt_count < self.settings.worker_retry_limit: run.status = "queued" else: run.status = "failed" run.error_code, run.updated_at = error_code, now run.locked_until = ( now + timedelta(seconds=min(60, 2**run.attempt_count)) if run.status == "queued" else None ) if run.status != "queued": run.completed_at = now await session.execute(update(RequestIdempotency).where( RequestIdempotency.id == run.idempotency_id ).values(status="failed", error_code=error_code, updated_at=now)) session.add(InteractionAudit( actor_type="agent", actor_id=run.user_id, session_id=run.session_id, portal="api", action_type=f"agent.run_{run.status}", detail={"run_id": run_id, "error_code": error_code}, created_at=now, ))