import asyncio import logging import time from collections.abc import AsyncIterator, Callable from sqlalchemy import select from yuxi.channel.hooks.lifecycle import HookEvent, LifecycleHookRegistry from yuxi.channel.message.block_reply_pipeline import create_block_reply_pipeline from yuxi.channel.message.bridge import AgentBridge from yuxi.channel.message.conversation_fence import ConversationFence from yuxi.channel.message.dispatch import MessageDispatcher from yuxi.channel.message.idempotency import claim, release from yuxi.channel.message.langfuse_trace import ChannelTrace, create_channel_trace from yuxi.channel.message.metrics import ( record_dispatch_duration_ms, record_message, record_rate_limit_reject, set_inflight, ) from yuxi.channel.message.models import ( PeerKind, StreamCancelledError, StreamStatus, StreamTimeoutError, UnifiedMessage, ) from yuxi.channel.message.rate_limiter import rate_limit_manager from yuxi.channel.protocols import ( InboundMessageHook, MessageSendingHook, OutboundProtocol, SendLifecycleHook, ) from yuxi.channel.plugins.registry import ChannelPluginRegistry from yuxi.channel.routing.matcher import RouteMatcher from yuxi.channel.routing.models import RouteBinding from yuxi.channel.routing.session_key import SessionKeyBuilder from yuxi.channel.security.allowlist import AllowlistChecker from yuxi.channel.security.identity_link import IdentityLinkResolver from yuxi.channel.security.pairing import InMemoryPairingStore, PairingManager from yuxi.channel.message.circuit_breaker import CircuitBreaker from yuxi.repositories.channel_binding_repo import ChannelBindingRepository, to_route_binding from yuxi.repositories.channel_msg_record_repo import ChannelMsgRecordRepository from yuxi.repositories.channel_thread_mapping_repo import ChannelThreadMappingRepository from yuxi.repositories.channel_user_mapping_repo import ChannelUserMappingRepository from yuxi.storage.postgres.manager import pg_manager from yuxi.storage.postgres.models_business import User logger = logging.getLogger(__name__) _BINDINGS_CACHE_TTL = 60.0 _IDEMPOTENCY_TTL = 300 _FALLBACK_MESSAGES: dict[str, str] = { "no_route": "暂不支持该渠道的消息处理。", "route_error": "消息路由异常,请稍后重试。", "dispatch_error": "消息处理异常,请稍后重试。", "approval_denied": "操作未通过审批。", "approval_check_error": "审批校验异常,请稍后重试。", "rate_limited": "当前消息量较大,请稍后重试。", } class MessageProcessor: def __init__( self, stream_fn: Callable[..., AsyncIterator[bytes]], channel_manager=None, *, msg_repo=None, binding_repo=None, user_mapping_repo=None, thread_mapping_repo=None, allowlist=None, pairing=None, identity_link_resolver=None, response_prefix: str = "", human_delay: tuple[float, float] | None = None, ): self._bridge = AgentBridge(stream_fn, channel_manager) self._msg_repo = msg_repo or ChannelMsgRecordRepository() self._binding_repo = binding_repo or ChannelBindingRepository() self._user_mapping_repo = user_mapping_repo or ChannelUserMappingRepository() self._thread_mapping_repo = thread_mapping_repo or ChannelThreadMappingRepository() self._hooks = LifecycleHookRegistry() self._allowlist = allowlist or AllowlistChecker() if pairing: self._pairing = pairing else: self._pairing = PairingManager(InMemoryPairingStore(), allowlist=self._allowlist) if allowlist is None: self._allowlist.load_dm_allowlist(["*"]) self._allowlist.load_group_allowlist(["*"]) self._response_prefix = response_prefix self._human_delay = human_delay self._identity_links = identity_link_resolver or IdentityLinkResolver() self._matcher = RouteMatcher() self._session_key_builder = SessionKeyBuilder() self._active_message_count = 0 self._active_count_lock = asyncio.Lock() self._cached_bindings: list[RouteBinding] = [] self._bindings_ts: float = 0.0 self._fence = ConversationFence() self._self_user_ids: dict[str, set[str]] = {} def set_self_user_ids(self, channel_type: str, account_id: str, user_ids: set[str]) -> None: key = f"{channel_type}:{account_id}" self._self_user_ids[key] = user_ids @property def hooks(self) -> LifecycleHookRegistry: return self._hooks async def _get_bindings(self) -> list[RouteBinding]: now = time.monotonic() if now - self._bindings_ts < _BINDINGS_CACHE_TTL and self._cached_bindings: return self._cached_bindings db_bindings = await self._binding_repo.list_all() self._cached_bindings = [to_route_binding(b) for b in db_bindings] self._bindings_ts = now await self._matcher.set_bindings(self._cached_bindings) return self._cached_bindings async def _get_user_for_channel(self, internal_user_id: str | None) -> User | None: if not internal_user_id: return None try: uid = int(internal_user_id) except (ValueError, TypeError): return None async with pg_manager.get_async_session_context() as db: return (await db.execute(select(User).where(User.id == uid))).scalar_one_or_none() def _is_echo(self, msg: UnifiedMessage) -> bool: key = f"{msg.channel_type}:{msg.account_id}" self_ids = self._self_user_ids.get(key) if not self_ids: return False return msg.sender.id in self_ids async def _ensure_security_loaded(self) -> None: await self._identity_links.ensure_loaded() self._session_key_builder = SessionKeyBuilder( identity_links=self._identity_links.list_links() ) async def process(self, msg: UnifiedMessage) -> None: if self._is_echo(msg): logger.debug("Skipping echo message: %s/%s", msg.channel_type, msg.msg_id) return limiter = await rate_limit_manager.get_limiter(msg.channel_type, msg.account_id) acquire_result = await limiter.try_acquire() if not acquire_result: record_rate_limit_reject(msg.channel_type) logger.warning( "Rate limited: channel=%s account=%s retry_after=%.2fs", msg.channel_type, msg.account_id, acquire_result.retry_after_sec, ) await self._send_fallback(msg, _FALLBACK_MESSAGES["rate_limited"]) return if not await claim(msg.msg_id, ttl=_IDEMPOTENCY_TTL): logger.debug("Duplicate/inflight message ignored: %s", msg.msg_id) return try: async with self._active_count_lock: self._active_message_count += 1 set_inflight(self._active_message_count) await self._ensure_security_loaded() await self._process_impl(msg, limiter) finally: await release(msg.msg_id) async with self._active_count_lock: self._active_message_count = max(0, self._active_message_count - 1) set_inflight(self._active_message_count) async def _process_impl(self, msg: UnifiedMessage, limiter) -> None: trace = await create_channel_trace(msg) await trace.add_span("idempotency_check", metadata={"claimed": True}) conversation_key = self._fence.key_for(msg) version, stop_event = self._fence.enter(conversation_key) lock = self._fence.lock_for(conversation_key) logger.debug( "processor: phase=enter msg_id=%s channel=%s conv_key=%s version=%d", msg.msg_id, msg.channel_type, conversation_key, version, ) async with lock: if version != self._fence.current_version(conversation_key): logger.debug("Foreground fence: message superseded for %s", conversation_key) record_message(msg.channel_type, "superseded") await trace.finish(error="superseded") return record_message(msg.channel_type, "received") plugin = ChannelPluginRegistry.get(msg.channel_type) if isinstance(plugin, InboundMessageHook): try: result = await plugin.on_message_received(msg) if result is None: logger.info("Message dropped by inbound hook: %s", msg.msg_id) await trace.finish(error="dropped_by_hook") return except Exception: logger.exception("Inbound hook failed, proceeding: %s", msg.msg_id) try: await self._hooks.fire(HookEvent.MESSAGE_RECEIVED, msg) except Exception: logger.exception("MESSAGE_RECEIVED hook failed: %s", msg.msg_id) chat_id = msg.group.id if msg.group and msg.group.id else msg.sender.id chat_type = "group" if msg.sender.kind != PeerKind.DIRECT else "direct" record_data = { "channel_id": msg.account_id, "channel_type": msg.channel_type, "message_id": msg.msg_id, "chat_id": chat_id, "chat_type": chat_type, "content_type": msg.message_type.value if msg.message_type else "text", "sender_user_id": msg.sender.id, "content_preview": (msg.content or "")[:500], "reply_to_message_id": msg.metadata.get("reply_to_message_id") if msg.metadata else None, "extra_data": msg.raw_payload, } record = None start_time = time.monotonic() try: record = await self._msg_repo.create_record(record_data) if stop_event.is_set(): logger.debug("Aborted before dispatch for %s", conversation_key) elapsed_ms = int((time.monotonic() - start_time) * 1000) await self._msg_repo.mark_error(record.id, "aborted_by_fence", elapsed_ms) await trace.finish(error="aborted_by_fence") return bindings = await self._get_bindings() dispatcher = MessageDispatcher( bindings, self._allowlist, self._pairing, matcher=self._matcher, session_key_builder=self._session_key_builder, user_mapping_repo=self._user_mapping_repo, thread_mapping_repo=self._thread_mapping_repo, ) try: await self._hooks.fire(HookEvent.BEFORE_AGENT_RUN, msg) except Exception: logger.exception("BEFORE_AGENT_RUN hook failed: %s", msg.msg_id) result = await dispatcher.dispatch(msg, trace=trace) record_dispatch_duration_ms( msg.channel_type, (time.monotonic() - start_time) * 1000, ) logger.debug( "processor: phase=dispatched msg_id=%s handled_by=%s success=%s agent_config_id=%s", msg.msg_id, result.handled_by, result.success, result.agent_config_id, ) if result.handled_by == "command": if result.command_response: await self._send_reply(msg, result.command_response) try: await self._hooks.fire(HookEvent.AFTER_AGENT_RUN, msg, result, result.command_response) except Exception: logger.exception("AFTER_AGENT_RUN hook failed: %s", msg.msg_id) elapsed_ms = int((time.monotonic() - start_time) * 1000) await self._msg_repo.mark_success( record.id, "", (result.command_response or "")[:500], elapsed_ms, ) await trace.finish() return if not result.success: try: await self._hooks.fire(HookEvent.AFTER_AGENT_RUN, msg, result, None) except Exception: logger.exception("AFTER_AGENT_RUN hook failed: %s", msg.msg_id) error_msg = result.error or "dispatch_failed" fallback = _FALLBACK_MESSAGES.get(error_msg) if fallback is None and result.error and result.error.startswith("approval_denied:"): fallback = _FALLBACK_MESSAGES["approval_denied"] if fallback: await self._send_fallback(msg, fallback) elapsed_ms = int((time.monotonic() - start_time) * 1000) await self._msg_repo.mark_error(record.id, error_msg, elapsed_ms) await trace.finish(error=error_msg) return if stop_event.is_set(): logger.debug("Aborted after dispatch for %s", conversation_key) try: await self._hooks.fire(HookEvent.AFTER_AGENT_RUN, msg, result, None) except Exception: logger.exception("AFTER_AGENT_RUN hook failed: %s", msg.msg_id) elapsed_ms = int((time.monotonic() - start_time) * 1000) await self._msg_repo.mark_error(record.id, "aborted_by_fence", elapsed_ms) await trace.finish(error="aborted_by_fence") return logger.debug( "processor: phase=reply_start msg_id=%s channel=%s agent_config_id=%s", msg.msg_id, msg.channel_type, result.agent_config_id, ) await self._handle_agent_reply(msg, result, plugin, stop_event, record, start_time, conversation_key, trace, limiter) except Exception as e: logger.exception("Message processing failed: %s", msg.msg_id) await trace.finish(error=f"{type(e).__name__}: {e}") try: await self._hooks.fire(HookEvent.AFTER_AGENT_RUN, msg, None, None) except Exception: logger.exception("AFTER_AGENT_RUN hook failed: %s", msg.msg_id) elapsed_ms = int((time.monotonic() - start_time) * 1000) error_text = f"{type(e).__name__}: {e}" if record is not None: await self._msg_repo.mark_error(record.id, error_text[:500], elapsed_ms) record_message(msg.channel_type, "error") async def _handle_agent_reply( self, msg: UnifiedMessage, result, plugin, stop_event: asyncio.Event | None, record, start_time: float, conversation_key: str, trace: ChannelTrace, limiter, ) -> None: user = await self._get_user_for_channel(result.internal_user_id) if user is None: logger.error( "channel_turn skip: msg_id=%s channel=%s reason=no_user_for_internal_id:%s", msg.msg_id, msg.channel_type, result.internal_user_id, ) elapsed_ms = int((time.monotonic() - start_time) * 1000) await self._msg_repo.mark_error(record.id, "no_user_for_internal_id", elapsed_ms) await trace.finish(error="no_user_for_internal_id") return agent_config_id = result.agent_config_id chat_id = msg.group.id if msg.group and msg.group.id else msg.sender.id try: await self._hooks.fire(HookEvent.AGENT_BOOTSTRAP, msg, result) except Exception: logger.exception("AGENT_BOOTSTRAP hook failed: %s", msg.msg_id) stream_status: StreamStatus | None = None if plugin is not None and isinstance(plugin, OutboundProtocol): try: await self._hooks.fire(HookEvent.MESSAGE_SENDING, msg, "") except Exception: logger.exception("MESSAGE_SENDING hook failed: %s", msg.msg_id) reply_text, stream_status = await self._process_with_pipeline( msg, result, agent_config_id, user, plugin, stop_event, limiter=limiter, trace=trace, ) else: async with pg_manager.get_async_session_context() as db: stream_result = await self._bridge.invoke_with_result( msg, result, agent_config_id, user, db, ) reply_text = stream_result.accumulated_text or None stream_status = stream_result.status if reply_text: if plugin is not None: reply_text = await self._apply_sending_hooks(msg, plugin, chat_id, reply_text) if reply_text is None: elapsed_ms = int((time.monotonic() - start_time) * 1000) await self._msg_repo.mark_error(record.id, "blocked_by_sending_hook", elapsed_ms) await trace.finish(error="blocked_by_sending_hook") return success = await self._send_with_retry(plugin, chat_id, reply_text, reply_to_id=msg.msg_id) if not success: logger.exception("Failed to send reply via %s after retries", msg.channel_type) else: logger.warning("No outbound plugin for channel: %s", msg.channel_type) elapsed_ms = int((time.monotonic() - start_time) * 1000) await self._msg_repo.mark_error(record.id, "no_outbound_plugin", elapsed_ms) await trace.finish(error="no_outbound_plugin") return if not reply_text and stream_status in (StreamStatus.TIMEOUT, StreamStatus.ERROR, StreamStatus.CANCELLED): fallback_map = { StreamStatus.TIMEOUT: "处理超时,请稍后重试。", StreamStatus.ERROR: "处理异常,已通知管理员。", StreamStatus.CANCELLED: "处理已取消。", } fallback = fallback_map.get(stream_status) if fallback: await self._send_fallback(msg, fallback) elapsed_ms = int((time.monotonic() - start_time) * 1000) await self._msg_repo.mark_error(record.id, stream_status.value, elapsed_ms) await trace.finish(error=stream_status.value) return if isinstance(plugin, SendLifecycleHook): try: await plugin.after_send_success( chat_id, {"content": reply_text, "msg_id": msg.msg_id}, ) except Exception: logger.exception("Send lifecycle hook failed: %s", msg.msg_id) try: await self._hooks.fire(HookEvent.MESSAGE_SENT, msg, reply_text) except Exception: logger.exception("MESSAGE_SENT hook failed: %s", msg.msg_id) try: await self._hooks.fire(HookEvent.AFTER_AGENT_RUN, msg, result, reply_text) except Exception: logger.exception("AFTER_AGENT_RUN hook failed: %s", msg.msg_id) elapsed_ms = int((time.monotonic() - start_time) * 1000) await self._msg_repo.mark_success( record.id, "", (reply_text or "")[:500], elapsed_ms, ) record_message(msg.channel_type, "success") await trace.finish() async def _apply_sending_hooks( self, msg: UnifiedMessage, plugin, chat_id: str, reply_text: str, ) -> str | None: if isinstance(plugin, MessageSendingHook): try: modified = await plugin.on_message_sending(msg, reply_text) if modified is None: logger.info("Message blocked by MessageSendingHook: %s", msg.msg_id) return None reply_text = modified except Exception: logger.exception("MessageSendingHook failed: %s", msg.msg_id) if isinstance(plugin, SendLifecycleHook): try: await plugin.before_send_attempt(chat_id, reply_text) except Exception: logger.exception("before_send_attempt failed: %s", msg.msg_id) try: result = await self._hooks.fire_sequential( HookEvent.MESSAGE_SENDING, initial=reply_text, msg=msg, ) reply_text = result if result is not None else reply_text except Exception: logger.exception("MESSAGE_SENDING hook failed: %s", msg.msg_id) return reply_text async def _process_with_pipeline( self, msg: UnifiedMessage, result, agent_config_id: int, user, plugin, stop_event: asyncio.Event | None = None, *, limiter=None, trace: "ChannelTrace | None" = None, ) -> tuple[str | None, StreamStatus | None]: target_id = msg.group.id if msg.group and msg.group.id else msg.sender.id async def send_fn(content: str) -> str | None: if stop_event and stop_event.is_set(): return None if not content: return None send_content = content if isinstance(plugin, MessageSendingHook): try: modified = await plugin.on_message_sending(msg, send_content) if modified is None: logger.info("Pipeline send blocked by MessageSendingHook: %s", msg.msg_id) return None send_content = modified except Exception: logger.exception("Pipeline MessageSendingHook failed: %s", msg.msg_id) if isinstance(plugin, SendLifecycleHook): try: await plugin.before_send_attempt(target_id, send_content) except Exception: logger.exception("Pipeline before_send_attempt failed: %s", msg.msg_id) success = await self._send_with_retry(plugin, target_id, send_content, reply_to_id=msg.msg_id) if not success: logger.exception("Pipeline send failed via %s after retries", msg.channel_type) return None pipeline = create_block_reply_pipeline( send_fn=send_fn, msg=msg, response_prefix=self._response_prefix, human_delay=self._human_delay, ) accumulated: list[str] = [] stream_error: StreamStatus | None = None exec_span_id = None if trace: exec_span_id = await trace.add_span( "agent_execution", metadata={ "agent_config_id": agent_config_id, "thread_id": result.thread_id, }, ) async with pg_manager.get_async_session_context() as db: try: stream_iter = self._bridge.invoke_stream( msg, result, agent_config_id, user, db, ) if limiter is not None: stream_iter = limiter.run_with_limit(stream_iter) async for chunk in stream_iter: if stop_event and stop_event.is_set(): await pipeline.abort() logger.debug("Pipeline aborted by foreground fence for %s", self._fence.key_for(msg)) break await pipeline.enqueue(chunk) if chunk.content: accumulated.append(chunk.content) except StreamTimeoutError: stream_error = StreamStatus.TIMEOUT logger.warning("Agent stream timeout for channel=%s peer=%s", msg.channel_type, msg.sender.id) except StreamCancelledError: stream_error = StreamStatus.CANCELLED logger.debug("Agent stream cancelled for channel=%s peer=%s", msg.channel_type, msg.sender.id) await pipeline.wait_idle() if exec_span_id: status_msg = stream_error.value if stream_error else "ok" await trace.end_span( exec_span_id, level="ERROR" if stream_error else "DEFAULT", status_message=status_msg, metadata={"stream_status": status_msg, "chunks": len(accumulated)}, ) if stop_event and stop_event.is_set(): return None, None return ("".join(accumulated) if accumulated else None, stream_error) async def _send_fallback(self, msg: UnifiedMessage, fallback_text: str) -> None: plugin = ChannelPluginRegistry.get(msg.channel_type) if plugin is None or not isinstance(plugin, OutboundProtocol): return try: target_id = msg.group.id if msg.group and msg.group.id else msg.sender.id await plugin.send_text(target_id, fallback_text, reply_to_id=msg.msg_id) except Exception: logger.exception("Failed to send fallback reply via %s", msg.channel_type) async def _send_with_retry( self, plugin, target_id: str, content: str, *, reply_to_id: str | None = None, max_retries: int = 3, circuit_breaker: "CircuitBreaker | None" = None, ) -> bool: last_error = None total_attempts = max_retries + 1 for attempt in range(total_attempts): if circuit_breaker is not None and not circuit_breaker.allow_request(): logger.warning( "Circuit breaker OPEN for %s, skipping send", plugin.id if hasattr(plugin, "id") else "unknown", ) return False try: await plugin.send_text(target_id, content, reply_to_id=reply_to_id) if circuit_breaker is not None: circuit_breaker.record_success() return True except Exception as e: last_error = e if circuit_breaker is not None: circuit_breaker.record_failure() logger.warning( "Send attempt %d/%d failed (target=%s): %s", attempt + 1, total_attempts, target_id, e, ) if attempt < max_retries: delay = min(0.3 * (2 ** attempt), 30.0) await asyncio.sleep(delay) logger.error("All %d send attempts failed (target=%s): %s", total_attempts, target_id, last_error) if isinstance(plugin, SendLifecycleHook): try: await plugin.after_send_failure(target_id, str(last_error)) except Exception: pass return False async def _send_reply(self, msg: UnifiedMessage, content: str) -> str | None: plugin = ChannelPluginRegistry.get(msg.channel_type) if plugin is None or not isinstance(plugin, OutboundProtocol): logger.warning("No outbound plugin for channel: %s", msg.channel_type) return None target_id = msg.group.id if msg.group and msg.group.id else msg.sender.id reply_to_id = msg.msg_id send_content = await self._apply_sending_hooks(msg, plugin, target_id, content) if send_content is None: return None success = await self._send_with_retry(plugin, target_id, send_content, reply_to_id=reply_to_id) if not success: logger.exception("Failed to send reply via %s after retries", msg.channel_type) return None