from __future__ import annotations import asyncio from datetime import datetime, timezone from yuxi.channels.models import ChannelMessage, ChannelResponse from yuxi.channels.policy.context_policy import ContextCommand, ContextPolicy from yuxi.channels.policy.dedup_policy import DedupPolicy from yuxi.channels.policy.group_chat_policy import GroupChatMode, GroupChatPolicy from yuxi.channels.policy.media_policy import MediaPolicy from yuxi.channels.policy.schedule_policy import SchedulePolicy from yuxi.channels.policy.security_policy import BaseSecurityPolicy from yuxi.channels.policy.voice_policy import VoicePolicy from yuxi.channels.policy.welcome_policy import WelcomePolicy from yuxi.channels.protocols.outbound import ChannelOutboundProtocol from yuxi.channels.services.context import ChatAbortEntry, ChatRunBuffer from yuxi.channels.session_mapper import VIRTUAL_DEPARTMENT_ID, SessionMapper from yuxi.utils.logging_config import logger class _ChannelUser: __slots__ = ("id", "department_id", "username", "user_id") def __init__(self, uid: str, dept_id: int): self.id = uid self.department_id = dept_id self.username = f"channel_user_{uid}" self.user_id = uid class MessageRouter: def __init__( self, channel_manager=None, dedup_policy: DedupPolicy | None = None, context_policy: ContextPolicy | None = None, group_chat_policy: GroupChatPolicy | None = None, welcome_policy: WelcomePolicy | None = None, schedule_policy: SchedulePolicy | None = None, media_policy: MediaPolicy | None = None, voice_policy: VoicePolicy | None = None, ): self._channel_manager = channel_manager self.dedup_policy = dedup_policy or DedupPolicy() self.context_policy = context_policy or ContextPolicy() self.media_policy = media_policy or MediaPolicy() self.voice_policy = voice_policy or VoicePolicy() self._schedule_policies: dict[str, SchedulePolicy] = {} self._group_chat_policies: dict[str, GroupChatPolicy] = {} self._welcome_policies: dict[str, WelcomePolicy] = {} self._security_policies: dict[str, BaseSecurityPolicy] = {} self._default_schedule_policy = schedule_policy or SchedulePolicy() self._default_group_chat_policy = group_chat_policy or GroupChatPolicy() self._default_welcome_policy = welcome_policy or WelcomePolicy() self.chat_abort_controllers: dict[str, ChatAbortEntry] = {} self.chat_run_buffers: dict[str, ChatRunBuffer] = {} def _get_schedule_policy(self, channel_id: str) -> SchedulePolicy: return self._schedule_policies.get(channel_id, self._default_schedule_policy) def _get_group_chat_policy(self, channel_id: str) -> GroupChatPolicy: return self._group_chat_policies.get(channel_id, self._default_group_chat_policy) def _get_welcome_policy(self, channel_id: str) -> WelcomePolicy: return self._welcome_policies.get(channel_id, self._default_welcome_policy) def _get_security_policy(self, channel_id: str, policy_data: dict) -> BaseSecurityPolicy: if channel_id not in self._security_policies: self._security_policies[channel_id] = BaseSecurityPolicy(policy_data) return self._security_policies[channel_id] async def _load_channel_policy(self, channel_id: str) -> dict | None: from sqlalchemy import select from yuxi.storage.postgres.manager import pg_manager from yuxi.storage.postgres.models_channels import ChannelPolicyConfig try: async with pg_manager.get_async_session_context() as db: result = await db.execute( select(ChannelPolicyConfig).where(ChannelPolicyConfig.channel_id == channel_id) ) policy = result.scalar_one_or_none() if policy: return policy.to_dict() except Exception: logger.warning(f"Failed to load policy for channel {channel_id}", exc_info=True) return None def _apply_policy_to_schedule(self, channel_id: str, policy_data: dict) -> SchedulePolicy: from datetime import time as dt_time from yuxi.channels.policy.schedule_policy import ScheduleConfig, TimeWindow schedule_config = ScheduleConfig( work_hours=TimeWindow( dt_time.fromisoformat(policy_data.get("work_hours_start", "09:00")), dt_time.fromisoformat(policy_data.get("work_hours_end", "18:00")), ), off_hours_reply=policy_data.get("off_hours_reply"), timezone_offset_hours=policy_data.get("timezone_offset", 8), ) policy = SchedulePolicy() policy.configure(schedule_config) self._schedule_policies[channel_id] = policy return policy def _apply_policy_to_group_chat(self, channel_id: str, policy_data: dict) -> GroupChatPolicy: mode_str = policy_data.get("group_chat_mode", "mention_only") try: mode = GroupChatMode(mode_str) except ValueError: mode = GroupChatMode.MENTION_ONLY policy = GroupChatPolicy() policy.configure(mode, whitelist=policy_data.get("whitelist_ids", [])) self._group_chat_policies[channel_id] = policy return policy def _apply_policy_to_welcome(self, channel_id: str, policy_data: dict) -> WelcomePolicy: welcome_msg = policy_data.get("welcome_message") policy = WelcomePolicy() policy.configure(message_template=welcome_msg if welcome_msg else None) self._welcome_policies[channel_id] = policy return policy async def route_inbound(self, message: ChannelMessage) -> None: identity = message.identity if await self.dedup_policy.check_and_remember(message): logger.debug(f"Dropping duplicate message from {identity.channel_id}") return adapter = self._channel_manager._adapters.get(identity.channel_id) if self._channel_manager else None if adapter is not None: assert isinstance(adapter, ChannelOutboundProtocol), ( f"Adapter {identity.channel_id} ({type(adapter).__name__}) must implement ChannelOutboundProtocol" ) context_result = self.context_policy.parse(message) if context_result.handled: await self._handle_context_command(message, context_result.command, context_result.args) return policy_data = await self._load_channel_policy(identity.channel_id) if policy_data and isinstance(policy_data, dict): schedule_policy = self._apply_policy_to_schedule(identity.channel_id, policy_data) group_chat_policy = self._apply_policy_to_group_chat(identity.channel_id, policy_data) welcome_policy = self._apply_policy_to_welcome(identity.channel_id, policy_data) security_policy = self._get_security_policy(identity.channel_id, policy_data) if not security_policy.check_dm_access(identity.channel_user_id).allowed: logger.info( f"Security policy blocked DM from {identity.channel_user_id} on channel {identity.channel_id}" ) return else: schedule_policy = self._get_schedule_policy(identity.channel_id) group_chat_policy = self._get_group_chat_policy(identity.channel_id) welcome_policy = self._get_welcome_policy(identity.channel_id) if not schedule_policy.is_working_hours(): reply = schedule_policy.get_off_hours_reply() if reply: response = ChannelResponse(identity=identity, content=reply) await self._send_response(identity.channel_id, response) return is_at_bot = bool(message.mentions and message.mentions.is_bot_mentioned) if not group_chat_policy.should_respond(message, is_at_bot): return from yuxi.repositories.channel_message_record_repository import ChannelMessageRecordRepository from yuxi.storage.postgres.manager import pg_manager async with pg_manager.get_async_session_context() as db: session_mapper = SessionMapper(db) internal_user_id = await session_mapper.resolve_user(message) thread_id = await session_mapper.resolve_thread(message, internal_user_id) msg_record_repo = ChannelMessageRecordRepository(db) agent_config_id = await self._resolve_agent_config_id(message, db) record = await msg_record_repo.create_record(message, agent_config_id=agent_config_id) if welcome_policy.mark_welcomed(internal_user_id): welcome_response = ChannelResponse( identity=identity, content=welcome_policy.get_welcome_message(), ) await self._send_response(identity.channel_id, welcome_response) run_id = f"{message.identity.channel_id}:{message.identity.channel_message_id}" try: task = asyncio.ensure_future( self._invoke_agent( db=db, query=message.content, thread_id=thread_id, internal_user_id=internal_user_id, agent_config_id=agent_config_id, message=message, ) ) self.chat_abort_controllers[run_id] = ChatAbortEntry(task=task) t_start = datetime.now(datetime.UTC) response_content = await task elapsed_ms = int((datetime.now(datetime.UTC) - t_start).total_seconds() * 1000) self.chat_abort_controllers.pop(run_id, None) response = ChannelResponse(identity=identity, content=response_content) await self._send_response(identity.channel_id, response) await msg_record_repo.mark_success(record.id, response, response_time_ms=elapsed_ms) self._record_stats_success(elapsed_ms) except asyncio.CancelledError: logger.info(f"Chat aborted for run {run_id}") self.chat_abort_controllers.pop(run_id, None) error_response = ChannelResponse( identity=identity, content="对话已被中断。", ) await self._send_response(identity.channel_id, error_response) except Exception as e: logger.error(f"Agent invocation failed: {e}") error_response = ChannelResponse( identity=identity, content="\u62b1\u6b49\uff0c\u5904\u7406\u4f60\u7684\u6d88\u606f\u65f6\u51fa\u9519\u4e86\uff0c\u8bf7\u7a0d\u540e\u518d\u8bd5\u3002", ) await self._send_response(identity.channel_id, error_response) await msg_record_repo.mark_error(record.id, str(e)) self._record_stats_error() async def route_outbound(self, agent_result, channel_id: str, identity) -> None: response = ChannelResponse( identity=identity, content=agent_result.response_text, attachments=getattr(agent_result, "attachments", []), ) await self._send_response(channel_id, response) def abort_chat(self, run_id: str) -> bool: entry = self.chat_abort_controllers.get(run_id) if entry is None: return False entry.abort() return True async def _handle_context_command(self, message: ChannelMessage, command: ContextCommand, args: str) -> None: identity = message.identity if command == ContextCommand.RESET: from yuxi.storage.postgres.manager import pg_manager async with pg_manager.get_async_session_context() as db: session_mapper = SessionMapper(db) internal_user_id = await session_mapper.resolve_user(message) await session_mapper.reset_thread(message, internal_user_id) response = ChannelResponse( identity=identity, content="\u5bf9\u8bdd\u4e0a\u4e0b\u6587\u5df2\u91cd\u7f6e\uff0c\u65b0\u7684\u4f1a\u8bdd\u5df2\u521b\u5efa", ) await self._send_response(identity.channel_id, response) elif command == ContextCommand.HISTORY: await self._cmd_history(message) elif command == ContextCommand.CONTEXT: await self._cmd_context(message) elif command == ContextCommand.SUMMARY: await self._cmd_summary(message) async def _cmd_history(self, message: ChannelMessage) -> None: identity = message.identity try: from yuxi.repositories.channel_message_record_repository import ChannelMessageRecordRepository from yuxi.storage.postgres.manager import pg_manager async with pg_manager.get_async_session_context() as db: repo = ChannelMessageRecordRepository(db) records = await repo.get_recent_records( identity.channel_id, identity.channel_chat_id or "", limit=10, ) if not records: response = ChannelResponse( identity=identity, content="\u6682\u65e0\u5bf9\u8bdd\u5386\u53f2\u8bb0\u5f55\u3002", ) else: lines = ["\u260e \u6700\u8fd1\u5bf9\u8bdd\u5386\u53f2\uff1a", ""] for r in reversed(records): created = r.created_at.strftime("%H:%M") if r.created_at else "" q_text = r.content_preview[:60] + ("..." if len(r.content_preview) > 60 else "") a_text = (r.reply_content_preview or "")[:60] if a_text: a_text = a_text + ("..." if len(r.reply_content_preview or "") > 60 else "") status_icon = "\u2705" if r.status == "success" else "\u274c" lines.append(f"[{created}] Q: {q_text}") if a_text: lines.append(f" A: {a_text} {status_icon}") else: lines.append(f" [{r.status}] {status_icon}") lines.append("") response = ChannelResponse(identity=identity, content="\n".join(lines)) self._record_stats_success(0) except Exception as e: logger.error(f"/history failed: {e}") response = ChannelResponse( identity=identity, content=f"\u83b7\u53d6\u5386\u53f2\u8bb0\u5f55\u5931\u8d25\uff1a{str(e)[:100]}", ) self._record_stats_error() await self._send_response(identity.channel_id, response) async def _cmd_context(self, message: ChannelMessage) -> None: identity = message.identity try: from yuxi.repositories.channel_message_record_repository import ChannelMessageRecordRepository from yuxi.storage.postgres.manager import pg_manager async with pg_manager.get_async_session_context() as db: session_mapper = SessionMapper(db) internal_user_id = await session_mapper.resolve_user(message) thread_id = await session_mapper.resolve_thread(message, internal_user_id) repo = ChannelMessageRecordRepository(db) msg_count_24h = await repo.get_chat_message_count(identity.channel_id, identity.channel_chat_id or "") lines = [ "\ud83d\udcca \u5f53\u524d\u5bf9\u8bdd\u4e0a\u4e0b\u6587\uff1a", "", f"\u6e20\u9053\uff1a{identity.channel_id} ({identity.channel_type.value})", f"\u804a\u5929 ID\uff1a{identity.channel_chat_id or 'N/A'}", f"\u4f1a\u8bdd ID\uff1a{thread_id[:8]}...", f"\u7528\u6237 ID\uff1a{internal_user_id[:12]}...", f"24h \u6d88\u606f\u6570\uff1a{msg_count_24h}", ] response = ChannelResponse(identity=identity, content="\n".join(lines)) self._record_stats_success(0) except Exception as e: logger.error(f"/context failed: {e}") response = ChannelResponse( identity=identity, content=f"\u83b7\u53d6\u4e0a\u4e0b\u6587\u4fe1\u606f\u5931\u8d25\uff1a{str(e)[:100]}", ) self._record_stats_error() await self._send_response(identity.channel_id, response) async def _cmd_summary(self, message: ChannelMessage) -> None: identity = message.identity try: from yuxi.storage.postgres.manager import pg_manager async with pg_manager.get_async_session_context() as db: session_mapper = SessionMapper(db) internal_user_id = await session_mapper.resolve_user(message) thread_id = await session_mapper.resolve_thread(message, internal_user_id) agent_config_id = await self._resolve_agent_config_id(message, db) summary_prompt = ( "\u8bf7\u7528\u4e00\u53e5\u8bdd\u6458\u8981\u603b\u7ed3\u4e0a\u8ff0\u5bf9\u8bdd\u7684\u6838\u5fc3\u5185\u5bb9\u3002" "\u53ea\u8f93\u51fa\u6458\u8981\u5185\u5bb9\uff0c\u4e0d\u8981\u8f93\u51fa\u5176\u4ed6\u4efb\u4f55\u5185\u5bb9\u3002" ) t_start = datetime.now(timezone.utc) # noqa: UP017 async with pg_manager.get_async_session_context() as db: summary_text = await self._invoke_agent( db=db, query=summary_prompt, thread_id=thread_id, internal_user_id=internal_user_id, agent_config_id=agent_config_id, message=message, ) elapsed_ms = int((datetime.now(timezone.utc) - t_start).total_seconds() * 1000) # noqa: UP017 response = ChannelResponse( identity=identity, content=f"\ud83d\udcdd \u5bf9\u8bdd\u6458\u8981\uff1a\n\n{summary_text}", ) self._record_stats_success(elapsed_ms) except Exception as e: logger.error(f"/summary failed: {e}") response = ChannelResponse( identity=identity, content=f"\u751f\u6210\u6458\u8981\u5931\u8d25\uff1a{str(e)[:100]}", ) self._record_stats_error() await self._send_response(identity.channel_id, response) async def _invoke_agent( self, db, query: str, thread_id: str, internal_user_id: str, agent_config_id: int, message: ChannelMessage, ) -> str: from yuxi.services.chat_service import stream_agent_chat channel_config = self._get_channel_config(message.identity.channel_id) department_id = channel_config.get("department_id", VIRTUAL_DEPARTMENT_ID) channel_user = _ChannelUser(internal_user_id, department_id) run_id = f"{message.identity.channel_id}:{message.identity.channel_message_id}" buffer = ChatRunBuffer(run_id=run_id) self.chat_run_buffers[run_id] = buffer adapter = None if self._channel_manager: adapter = self._channel_manager._adapters.get(message.identity.channel_id) supports_streaming = getattr(adapter, "supports_streaming", False) try: async for chunk in stream_agent_chat( query=query, agent_config_id=agent_config_id, thread_id=thread_id, meta={ "source": "channel", "channel_id": message.identity.channel_id, "channel_type": message.identity.channel_type.value, }, image_content=None, current_user=channel_user, db=db, ): import json try: data = json.loads(chunk.decode("utf-8").strip()) if data.get("status") == "loading" and data.get("response"): buffer.append_chunk(data["response"]) if supports_streaming and adapter: await adapter.send_stream_chunk( chat_id=message.identity.channel_chat_id or "", message_id=message.identity.channel_message_id or "", chunk_text=data["response"], finished=False, ) except (json.JSONDecodeError, UnicodeDecodeError): logger.debug(f"Non-JSON stream chunk: {chunk[:100]!r}") except Exception: logger.debug(f"Unexpected stream parse error for chunk: {chunk[:100]!r}") buffer.mark_finished() if supports_streaming and adapter: await adapter.send_stream_chunk( chat_id=message.identity.channel_chat_id or "", message_id=message.identity.channel_message_id or "", chunk_text="", finished=True, ) return buffer.get_full_text() finally: self.chat_run_buffers.pop(run_id, None) async def _resolve_agent_config_id(self, message: ChannelMessage, db=None) -> int: channel_id = message.identity.channel_id content = message.content.strip() agent_config_id = message.metadata.get("agent_config_id") if agent_config_id is not None: return int(agent_config_id) channel_config = self._get_channel_config(channel_id) if content.startswith("/"): cmd = content.split()[0].lower() cmd_routing = channel_config.get("command_routing", {}) if cmd in cmd_routing: return int(cmd_routing[cmd]) if db is not None: from sqlalchemy import select as sa_select from yuxi.storage.postgres.models_channels import ChannelRoutingRule result = await db.execute( sa_select(ChannelRoutingRule.agent_config_id) .where( ChannelRoutingRule.channel_id == channel_id, ChannelRoutingRule.command == cmd, ) .limit(1) ) row = result.scalar_one_or_none() if row is not None: return await self._resolve_agent_id_to_config_id(row, db) channel_default = channel_config.get("agent_config_id") if channel_default is not None: try: return int(channel_default) except (ValueError, TypeError): return await self._resolve_agent_id_to_config_id(str(channel_default), db) global_default = self._get_global_default_agent_id() if global_default is not None: return global_default return 1 async def _resolve_agent_id_to_config_id(self, agent_id: str, db) -> int: from yuxi.repositories.agent_config_repository import AgentConfigRepository repo = AgentConfigRepository(db) config = await repo.get_or_create_default(department_id=-1, agent_id=agent_id) if config is not None: return config.id return 1 def _get_channel_config(self, channel_id: str) -> dict: if self._channel_manager and hasattr(self._channel_manager, "_channels_config"): return self._channel_manager._channels_config.get(channel_id, {}) return {} def _get_global_default_agent_id(self) -> int | None: from yuxi import config as conf return getattr(conf, "default_agent_id", None) async def _send_response(self, channel_id: str, response: ChannelResponse) -> None: if self._channel_manager and hasattr(self._channel_manager, "send_outbound"): await self._channel_manager.send_outbound(channel_id, response) def _record_stats_success(self, elapsed_ms: int) -> None: collector = getattr(self._channel_manager, "_stats_collector", None) if self._channel_manager else None if collector: collector.record_request() collector.record_response_time(float(elapsed_ms)) def _record_stats_error(self) -> None: collector = getattr(self._channel_manager, "_stats_collector", None) if self._channel_manager else None if collector: collector.record_request() collector.record_error()