from __future__ import annotations import logging import time from yuxi.channel.domain.exception.agent_crash_error import AgentCrashError from yuxi.channel.domain.exception.recoverable_error import RecoverableError from yuxi.channel.domain.model.message.dispatch_result import DispatchResult from yuxi.channel.domain.port.bot_loop_guard_port import BotLoopGuardPort from yuxi.channel.domain.port.circuit_breaker_port import CircuitBreakerPort from yuxi.channel.domain.port.content_filter_port import ContentFilterPort from yuxi.channel.domain.port.metrics_port import MetricsPort from yuxi.channel.domain.repository.message_log_repository import MessageLogRepositoryPort from yuxi.channel.application.service.delivery_service import DeliveryService from yuxi.channel.application.service.session_resolver import SessionResolver logger = logging.getLogger(__name__) class DispatchService: def __init__( self, content_filter: ContentFilterPort, bot_loop_guard: BotLoopGuardPort, circuit_breaker: CircuitBreakerPort, session_resolver: SessionResolver, delivery_service: DeliveryService, *, message_log_repo: MessageLogRepositoryPort | None = None, metrics: MetricsPort | None = None, ) -> None: self._content_filter = content_filter self._bot_loop_guard = bot_loop_guard self._circuit_breaker = circuit_breaker self._session_resolver = session_resolver self._delivery = delivery_service self._message_log_repo = message_log_repo self._metrics = metrics @property def delivery_service(self) -> DeliveryService: return self._delivery @property def session_resolver(self) -> SessionResolver: return self._session_resolver async def dispatch(self, payload: dict) -> DispatchResult: message_id = payload["message_id"] channel_type = payload["channel_type"] trace_id = payload.get("trace_id", "") sender_id = payload.get("sender_id", "") await self._create_log( trace_id=trace_id, message_id=message_id, channel_type=channel_type, sender_id=sender_id, content_summary=payload.get("content", "")[:500], agent_config_id=payload.get("agent_config_id"), ) start = time.monotonic() try: result = await self._do_dispatch(payload) if self._metrics: await self._metrics.record_worker_dispatch_duration(channel_type, time.monotonic() - start) await self._metrics.record_worker_dispatch_total( channel_type, "success" if result.success else "failure" ) return result except RecoverableError: logger.warning("recoverable error, re-raise for retry: %s [%s]", message_id, trace_id) raise except AgentCrashError: logger.error("agent crash: %s [%s]", message_id, trace_id) if self._metrics: await self._metrics.record_worker_dispatch_total(channel_type, "agent_crash") return DispatchResult(success=False, message_id=message_id, error="agent_crash") except Exception as exc: logger.exception("dispatch error: %s [%s]", message_id, trace_id) if self._metrics: await self._metrics.record_worker_dispatch_total(channel_type, "error") await self._update_log( trace_id, message_id, worker_result="error", status="failed", error_message=str(exc)[:500], ) return DispatchResult(success=False, message_id=message_id, error=str(exc)) async def _do_dispatch(self, payload: dict) -> DispatchResult: message_id = payload["message_id"] channel_type = payload["channel_type"] session_id = payload.get("session_id", "") trace_id = payload.get("trace_id", "") sender_id = payload.get("sender_id", "") is_group = payload.get("metadata", {}).get("is_group", False) content = payload["content"] filter_result = await self._content_filter.check(content, channel_type=channel_type) if not filter_result.passed: logger.warning("content blocked in worker: %s [%s]", filter_result.violations, trace_id) return DispatchResult(success=True, message_id=message_id) if filter_result.masked_fields: payload = {**payload, "content": filter_result.filtered_content} if not await self._bot_loop_guard.check(session_id, sender_id=sender_id, is_group=is_group): logger.warning( "bot loop detected: session=%s sender=%s is_group=%s", session_id, sender_id, is_group, ) return DispatchResult(success=True, message_id=message_id) session = await self._session_resolver.resolve(payload) if session is None: return DispatchResult(success=False, message_id=message_id, error="session_error") agent_config_id = payload.get("agent_config_id") if not agent_config_id: try: agent_config_id = int(session.agent_id) except (ValueError, TypeError): agent_config_id = 1 if not await self._circuit_breaker.is_available(agent_config_id): logger.warning("circuit open for agent_config_id=%s, skip [%s]", agent_config_id, trace_id) return DispatchResult(success=False, message_id=message_id, error="circuit_open") try: result = await self._delivery.deliver(payload, session) if result.success: await self._circuit_breaker.record_success(agent_config_id) return result except AgentCrashError: await self._circuit_breaker.record_failure(agent_config_id) raise except RecoverableError: raise except Exception: await self._circuit_breaker.record_failure(agent_config_id) raise async def _create_log( self, *, trace_id: str, message_id: str, channel_type: str, sender_id: str | None = None, content_summary: str | None = None, agent_config_id: int | None = None, ) -> None: if not self._message_log_repo: return try: await self._message_log_repo.create_log( trace_id=trace_id, message_id=message_id, channel_type=channel_type, direction="inbound", sender_id=sender_id, content_summary=content_summary, agent_config_id=agent_config_id, ) except Exception: logger.debug("message log create failed for %s", trace_id) async def _update_log( self, trace_id: str, message_id: str, *, worker_result: str, status: str, error_message: str | None = None, ) -> None: if not self._message_log_repo: return try: await self._message_log_repo.update_worker_result( trace_id=trace_id, message_id=message_id, worker_result=worker_result, status=status, error_message=error_message, ) except Exception: logger.debug("message log update failed for %s", trace_id)