import asyncio import logging from contextlib import asynccontextmanager from dataclasses import dataclass, field from enum import StrEnum from collections.abc import Callable from typing import Any @dataclass class ChannelContext: """单个通道请求/会话的上下文对象,承载一次交互所需的全部运行时信息。""" channel_type: str # 通道类型标识,例如 "feishu"、"wechat" account_id: str # 该通道下的账户唯一标识 cancel_event: asyncio.Event = field(default_factory=asyncio.Event) # 用于取消当前任务的异步事件 config: dict = field(default_factory=dict) # 通道级配置字典 queue: asyncio.Queue[Any] | None = None # 可选的消息队列,用于异步流式回传 logger: logging.Logger | None = None # 专用日志器,__post_init__ 中自动初始化 request_id: str | None = None # 可选的链路追踪请求 ID def __post_init__(self): # 若调用方未传入 logger,则按 "yuxi.channel.{channel_type}.{account_id}" 的命名规范自动创建 if self.logger is None: self.logger = logging.getLogger(f"yuxi.channel.{self.channel_type}.{self.account_id}") @dataclass(slots=True) class ConversationBinding: """会话绑定关系,描述一条消息所属的对话层级结构。""" channel: str # 通道标识 account_id: str # 账户标识 conversation_id: str # 当前会话/对话 ID parent_conversation_id: str | None = None # 父会话 ID,用于支持嵌套会话场景 thread_id: str | None = None # 线程 ID,用于支持话题(thread)级别的细分 class ContextVisibilityMode(StrEnum): """补充上下文可见性控制模式,决定哪些附加信息可以进入大模型上下文。""" ALL = "all" # 全部放行,不做过滤 ALLOWLIST = "allowlist" # 仅允许发送者在白名单中的内容 ALLOWLIST_QUOTE = "allowlist_quote" # 白名单逻辑,但引用(quote)类型始终放行 NONE = "none" # 全部拒绝(通常仅保留系统级提示) def _evaluate_supplemental_visibility( mode: ContextVisibilityMode, kind: str, sender_allowed: bool, ) -> bool: """根据可见性模式判断某类补充上下文是否应被保留。 参数: mode: 当前采用的可见性模式。 kind: 上下文类型,如 "quote"、"forwarded"、"thread"。 sender_allowed: 发送者是否已在白名单中。 """ if mode == ContextVisibilityMode.ALL: return True if sender_allowed: return True if mode == ContextVisibilityMode.ALLOWLIST_QUOTE and kind == "quote": return True return False @dataclass(slots=True) class SupplementalContext: """补充上下文,包含引用、转发、话题线程等可能来自外部的不受信信息。 在注入大模型前,应通过 filter_by_visibility() 按业务策略进行过滤, 防止不可信内容污染系统提示或诱导模型产生错误回复。 """ quote: dict | None = None # 引用消息上下文 forwarded: dict | None = None # 转发消息上下文 thread: dict | None = None # 话题/线程上下文 untrusted_context: list[dict] | None = None # 其他未受信上下文列表 group_system_prompt: str | None = None # 群组级别的系统提示(通常不受可见性过滤影响) def filter_by_visibility(self, mode: ContextVisibilityMode) -> "SupplementalContext": """按指定可见性模式过滤当前实例,返回过滤后的新 SupplementalContext。 untrusted_context 与 group_system_prompt 不参与过滤,始终保留。 """ if mode == ContextVisibilityMode.ALL: return self quote = self.quote forwarded = self.forwarded thread = self.thread if not _evaluate_supplemental_visibility(mode, "quote", bool(quote and quote.get("sender_allowed"))): quote = None if not _evaluate_supplemental_visibility( mode, "forwarded", bool(forwarded and forwarded.get("sender_allowed")) ): forwarded = None if not _evaluate_supplemental_visibility(mode, "thread", bool(thread and thread.get("sender_allowed"))): thread = None return SupplementalContext( quote=quote, forwarded=forwarded, thread=thread, untrusted_context=self.untrusted_context, group_system_prompt=self.group_system_prompt, ) def is_empty(self) -> bool: """判断当前补充上下文是否不包含任何有效信息。""" return not any((self.quote, self.forwarded, self.thread, self.untrusted_context, self.group_system_prompt)) class ChannelRuntimeContexts: """跨生命周期状态注册表, keyed by (channel_id, account_id, capability). 借鉴 openclaw 的 ChannelRuntimeContextRegistry 模式, 允许插件和模块在不产生紧耦合的情况下动态共享运行时状态。 Usage:: registry = ChannelRuntimeContexts() dispose = registry.register("feishu", "account1", "gateway_state", {"connected": True}) state = registry.get("feishu", "account1", "gateway_state") dispose() """ def __init__(self): # 内部存储结构: key -> (token, context) # token 用于在并发 dispose 时做安全校验,防止误删后续重新注册的同名条目 self._store: dict[str, tuple[int, Any]] = {} self._counter = 0 @staticmethod def _make_key(channel_id: str, account_id: str, capability: str) -> str: """使用 \0 拼接三元组,避免不同字段组合产生 key 冲突。""" return f"{channel_id}\0{account_id}\0{capability}" def register( self, channel_id: str, account_id: str, capability: str, context: Any, ) -> Callable[[], None]: """注册一个运行时上下文,返回 dispose 闭包用于手动注销。 注意: - 同一 (channel_id, account_id, capability) 重复注册将覆盖旧值。 - dispose 是幂等的,多次调用无副作用。 """ key = self._make_key(channel_id, account_id, capability) self._counter += 1 token = self._counter self._store[key] = (token, context) disposed = False def dispose(): nonlocal disposed if disposed: return disposed = True entry = self._store.get(key) # 仅当条目存在且 token 匹配时才删除,防止并发覆盖后误删新条目 if entry is None or entry[0] != token: return del self._store[key] return dispose def get(self, channel_id: str, account_id: str, capability: str) -> Any | None: """获取指定三元组对应的运行时上下文,若不存在则返回 None。""" entry = self._store.get(self._make_key(channel_id, account_id, capability)) return entry[1] if entry is not None else None def dispose_all(self): """清空全部注册状态并重置计数器,通常用于服务关闭或测试清理。""" self._store.clear() self._counter = 0 @asynccontextmanager async def task_scoped_contexts(registry: ChannelRuntimeContexts): """异步上下文管理器, 在退出时自动 dispose 作用域内所有注册。 Usage:: registry = ChannelRuntimeContexts() async with task_scoped_contexts(registry) as scoped: scoped.register("feishu", "main", "temp", data) # 退出时自动清理 """ leases: list[Callable[[], None]] = [] class _Scoped: """作用域内暴露的注册代理,自动收集所有 dispose 句柄以便退出时统一清理。""" def register(self, channel_id: str, account_id: str, capability: str, context: Any): dispose = registry.register(channel_id, account_id, capability, context) leases.append(dispose) return dispose try: yield _Scoped() finally: # 按后进先出顺序清理,尽量保证依赖关系正确释放 for dispose in reversed(leases): try: dispose() except Exception: # 清理异常不应阻断后续释放流程 pass