refactor(channel): 完成频道模块大重构与功能扩展

本次提交对频道模块进行了全面重构并新增多项核心功能:
1.  优化适配器状态获取逻辑,修复状态返回空值问题
2.  新增4种频道异常类型,完善错误处理体系
3.  大幅精简Mixin类,移除冗余的抽象方法定义
4.  重构适配器注册系统,统一注册入口并新增内置适配器加载方法
5.  扩展插件系统,新增更多元数据配置项支持
6.  新增线程类型、会话范围等模型定义,扩展事件类型枚举
7.  优化用户映射逻辑,使用PostgreSQL upsert避免重复创建
8.  新增历史消息注入模块,支持多格式历史格式化与缓存管理
9.  新增线程能力配置与各平台预置适配配置
10. 新增线程绑定管理器,支持多类型线程绑定生命周期管理
11. 重构__init__.py,整理导出模块与类型
12. 扩展基础适配器类,新增凭证解析、状态存储等核心方法
13. 重写消息路由器,支持按频道加载策略、安全校验与多命令处理
14. 新增/history、/context、/summary等交互命令实现
15. 优化消息记录与统计逻辑,完善路由调度链路
This commit is contained in:
Kris 2026-05-13 16:41:11 +08:00
parent 4a5051e7d1
commit ede29b1809
14 changed files with 1243 additions and 157 deletions

View File

@ -1,3 +1,60 @@
from yuxi.channels.auth import (
AuthHealth,
AuthHealthMonitor,
BackoffConfig,
BaseTokenProvider,
CertificateTokenProvider,
CompatTokenProvider,
DmPolicy,
ExponentialBackoff,
GroupPolicy,
NetworkGuardError,
OAuth2TokenProvider,
QRTokenProvider,
SecretManager,
SecretSource,
SecurityPolicy,
SecurityPolicyEngine,
SensitiveDataFilter,
StaticTokenProvider,
TokenProviderType,
TokenState,
UnifiedTokenManager,
apply_ssrf_guard_defaults,
build_ssrf_safe_headers,
check_response_size,
check_response_text,
fetch_with_ssrf_guard,
get_health_monitor,
get_secret_manager,
get_security_policy_engine,
get_token_manager,
init_secret_manager,
is_hostname_allowed,
is_private_url,
redact_sensitive,
safe_webhook_url,
sanitize_transport_kwargs,
validate_url,
validate_url_with_whitelist,
)
from yuxi.channels.base import BaseChannelAdapter
from yuxi.channels.bridge import BridgeAdapter
from yuxi.channels.capabilities import ChannelCapabilities
from yuxi.channels.exceptions import (
ChannelAuthenticationError,
ChannelConnectionError,
ChannelException,
ChannelNotConnectedError,
ChannelRateLimitError,
DeliveryFailedError,
MessageFormatError,
)
from yuxi.channels.infra.broadcast import EventBroadcaster
from yuxi.channels.infra.circuit_breaker import CircuitBreaker, CircuitBreakerOpenError, CircuitState
from yuxi.channels.infra.config_watcher import ConfigWatcher
from yuxi.channels.manager import ChannelManager
from yuxi.channels.meta import ChannelMeta
from yuxi.channels.models import (
AgentRequest,
AgentResult,
@ -17,26 +74,8 @@ from yuxi.channels.models import (
TokenStatus,
build_snapshot_from_adapter,
)
from yuxi.channels.base import BaseChannelAdapter
from yuxi.channels.exceptions import (
ChannelAuthenticationError,
ChannelConnectionError,
ChannelException,
ChannelNotConnectedError,
ChannelRateLimitError,
DeliveryFailedError,
MessageFormatError,
)
from yuxi.channels.infra.circuit_breaker import CircuitBreaker, CircuitBreakerOpenError, CircuitState
from yuxi.channels.registry import ChannelRegistry, register_builtin_adapter
from yuxi.channels.manager import ChannelManager
from yuxi.channels.router import MessageRouter
from yuxi.channels.session_mapper import SessionMapper
from yuxi.channels.services.maintenance import MaintenanceRunner
from yuxi.channels.bridge import BridgeAdapter
from yuxi.channels.plugin import ChannelPlugin, channel_plugin
from yuxi.channels.meta import ChannelMeta
from yuxi.channels.capabilities import ChannelCapabilities
from yuxi.channels.policy.heartbeat import BaseHeartbeatAdapter
from yuxi.channels.protocols import (
ChannelConfigProtocol,
ChannelGatewayProtocol,
@ -48,13 +87,12 @@ from yuxi.channels.protocols import (
ChannelStatusProtocol,
ChannelThreadingProtocol,
)
from yuxi.channels.registry import ChannelRegistry, register_builtin_adapter
from yuxi.channels.router import MessageRouter
from yuxi.channels.services.context import ChatAbortEntry, ChatRunBuffer, GatewayRequestContext
from yuxi.channels.infra.broadcast import EventBroadcaster
from yuxi.channels.infra.config_watcher import ConfigWatcher
from yuxi.channels.services.doctor import ConfigDoctor, DiagnosisIssue
from yuxi.channels.policy.heartbeat import BaseHeartbeatAdapter
from yuxi.channels.services.maintenance import MaintenanceRunner
from yuxi.channels.session_mapper import SessionMapper
__all__ = [
"AgentRequest",
@ -113,4 +151,42 @@ __all__ = [
"ConfigWatcher",
"ConfigDoctor",
"DiagnosisIssue",
"SecretManager",
"SecretSource",
"get_secret_manager",
"init_secret_manager",
"BaseTokenProvider",
"TokenProviderType",
"TokenState",
"UnifiedTokenManager",
"get_token_manager",
"CertificateTokenProvider",
"CompatTokenProvider",
"OAuth2TokenProvider",
"QRTokenProvider",
"StaticTokenProvider",
"AuthHealth",
"AuthHealthMonitor",
"get_health_monitor",
"BackoffConfig",
"ExponentialBackoff",
"DmPolicy",
"GroupPolicy",
"SecurityPolicy",
"SecurityPolicyEngine",
"get_security_policy_engine",
"SensitiveDataFilter",
"redact_sensitive",
"NetworkGuardError",
"apply_ssrf_guard_defaults",
"build_ssrf_safe_headers",
"check_response_size",
"check_response_text",
"fetch_with_ssrf_guard",
"is_hostname_allowed",
"is_private_url",
"safe_webhook_url",
"sanitize_transport_kwargs",
"validate_url",
"validate_url_with_whitelist",
]

View File

@ -2,7 +2,7 @@ from __future__ import annotations
from abc import ABC, abstractmethod
from collections.abc import AsyncIterator, Awaitable, Callable
from typing import Any, ClassVar
from typing import TYPE_CHECKING, Any, ClassVar
from yuxi.channels.capabilities import CAPS_SIMPLE_TEXT, ChannelCapabilities
from yuxi.channels.meta import ChannelMeta
@ -13,9 +13,16 @@ from yuxi.channels.models import (
DeliveryResult,
HealthStatus,
)
from yuxi.channels.protocols.gateway import ChannelGatewayProtocol
from yuxi.channels.protocols.lifecycle import ChannelLifecycleProtocol
if TYPE_CHECKING:
from yuxi.channels.auth.backoff import ExponentialBackoff
from yuxi.channels.auth.secret_manager import SecretManager
from yuxi.channels.services.plugin_state_store import PluginStateStore
class BaseChannelAdapter(ABC):
class BaseChannelAdapter(ChannelLifecycleProtocol, ChannelGatewayProtocol, ABC):
channel_id: ClassVar[str]
channel_type: ClassVar[ChannelType]
@ -30,10 +37,77 @@ class BaseChannelAdapter(ABC):
capabilities: ClassVar[ChannelCapabilities] = CAPS_SIMPLE_TEXT
meta: ClassVar[ChannelMeta] = ChannelMeta(id="", label="")
_state_store: PluginStateStore | None = None
def __init__(self, config: dict[str, Any] | None = None):
self.config = config or {}
self._status: Any = None
self._message_handler: Callable[[ChannelMessage], Awaitable[None]] | None = None
self._stream_state: dict[str, int] = {}
self._credential_backoff = self._create_backoff()
@property
def secret_manager(self) -> SecretManager:
from yuxi.channels.auth.secret_manager import get_secret_manager
return get_secret_manager()
@staticmethod
def _create_backoff() -> ExponentialBackoff:
from yuxi.channels.auth.backoff import BackoffConfig, ExponentialBackoff
return ExponentialBackoff(BackoffConfig(base_seconds=5.0, max_seconds=300.0, jitter_pct=0.2))
async def resolve_credential(self, key: str, fallback_keys: list[str] | None = None) -> str | None:
from yuxi.channels.auth.secret_manager import SecretSource
sm = self.secret_manager
sources = [
SecretSource.CONFIG,
SecretSource.ENV,
SecretSource.FILE,
SecretSource.SECRET_REF,
SecretSource.EXEC,
]
resolved = await sm.resolve_secret(key, sources=sources, config=self.config)
if resolved:
return resolved
if fallback_keys:
for fk in fallback_keys:
resolved = await sm.resolve_secret(fk, sources=sources, config=self.config)
if resolved:
return resolved
return None
def _log_config_safely(self) -> None:
from yuxi.channels.auth.secret_manager import SecretManager
from yuxi.utils.logging_config import logger
safe_config = SecretManager.redact_config(self.config)
logger.debug(f"[{self.channel_id}] Config (redacted): {safe_config}")
async def state_get(self, key: str, namespace: str = "default") -> Any | None:
if self._state_store is None:
return None
return await self._state_store.get(self.channel_id, key, namespace)
async def state_set(
self,
key: str,
value: Any,
namespace: str = "default",
ttl_seconds: int | None = None,
) -> None:
if self._state_store is None:
return
await self._state_store.set(self.channel_id, key, value, namespace, ttl_seconds)
async def state_delete(self, key: str, namespace: str = "default") -> None:
if self._state_store is None:
return
await self._state_store.delete(self.channel_id, key, namespace)
def _get_stream_state(self, chat_id: str, msg_id: str) -> int:
return self._stream_state.get(f"{chat_id}:{msg_id}", 0)
@ -54,7 +128,7 @@ class BaseChannelAdapter(ABC):
async def send(self, response: ChannelResponse) -> DeliveryResult: ...
async def receive(self) -> AsyncIterator[ChannelMessage]:
return
raise NotImplementedError
yield # type: ignore[misc]
@abstractmethod
@ -105,7 +179,66 @@ class BaseChannelAdapter(ABC):
return True
async def _refresh_token_if_needed(self) -> bool:
return True
from yuxi.utils.logging_config import logger
try:
self._credential_backoff.mark_attempt()
return True
except Exception as e:
delay = self._credential_backoff.next_delay
logger.warning(
f"[{self.channel_id}] Token refresh failed "
f"(attempt {self._credential_backoff.attempt}), "
f"next retry in {delay:.1f}s: {e}"
)
await self._credential_backoff.wait()
return False
async def pre_connect(self) -> dict:
return {}
def is_enabled(self) -> bool:
return bool(self.config.get("enabled", False))
def is_configured(self) -> bool:
return bool(self.config)
@property
def status(self) -> str:
return getattr(self, "_status", "unknown") or "unknown"
def snapshot(self) -> dict[str, Any]:
from yuxi.channels.models import build_snapshot_from_adapter
return build_snapshot_from_adapter(self).model_dump()
def resolve_account_state(self, configured: bool, enabled: bool) -> str:
if not configured:
return "not_configured"
if not enabled:
return "disabled"
return "active"
def collect_status_issues(self, accounts: list) -> list[str]:
issues = []
if not self.is_configured():
issues.append("not_configured")
if not self.is_enabled():
issues.append("disabled")
return issues
async def check_ready(self) -> bool:
return self.is_enabled() and self.is_configured()
def on_config_changed(self, prev_cfg: dict, next_cfg: dict) -> None:
pass
async def run_startup_maintenance(self) -> None:
pass
async def logout_account(self, ctx) -> None:
raise NotImplementedError
async def login_with_qr_start(self, force: bool, timeout_ms: int) -> str:
raise NotImplementedError

View File

@ -92,7 +92,7 @@ class BridgeAdapter:
@property
def status(self) -> str:
return getattr(self._legacy, "_status", "unknown")
return getattr(self._legacy, "_status", None) or "unknown"
def snapshot(self) -> dict[str, Any]:
from yuxi.channels.models import build_snapshot_from_adapter

View File

@ -3,6 +3,24 @@ from __future__ import annotations
from pydantic import BaseModel, Field
class ThreadCapabilities(BaseModel):
supports_native_threads: bool = False
supports_topics: bool = False
supports_reply_chains: bool = False
supports_simulated_threads: bool = False
max_thread_depth: int = 1
supports_history_fetch: bool = False
supports_parent_injection: bool = False
supports_thread_binding: bool = False
history_max_messages: int = 50
history_max_chars: int = 4000
history_ttl_seconds: int = 300
@property
def requires_simulation(self) -> bool:
return not self.supports_native_threads and self.supports_simulated_threads
class TTSVoiceCapabilities(BaseModel):
synthesis_target: str = "voice-note"
transcodes_audio: bool = False
@ -36,6 +54,8 @@ class ChannelCapabilities(BaseModel):
threads: bool = False
thread_caps: ThreadCapabilities = Field(default_factory=ThreadCapabilities)
media: bool = False
max_media_size_mb: int = 100
@ -53,6 +73,10 @@ class ChannelCapabilities(BaseModel):
send_ephemeral: bool = False
forward: bool = False
narrowcast: bool = False
vision: bool = False
approval: bool = False
typing: bool = False
@ -114,3 +138,77 @@ CAPS_GROUP_CHAT_ONLY = ChannelCapabilities(
supports_streaming=True,
streaming_modes=["block"],
)
THREAD_CAPABILITIES: dict[str, ThreadCapabilities] = {
"matrix": ThreadCapabilities(
supports_native_threads=True,
supports_history_fetch=True,
supports_parent_injection=True,
supports_thread_binding=True,
history_max_messages=100,
),
"discord": ThreadCapabilities(
supports_native_threads=True,
supports_history_fetch=True,
supports_thread_binding=True,
history_max_messages=100,
),
"slack": ThreadCapabilities(
supports_native_threads=True,
supports_history_fetch=True,
supports_thread_binding=True,
history_max_messages=50,
),
"feishu": ThreadCapabilities(
supports_native_threads=True,
supports_history_fetch=True,
supports_parent_injection=True,
supports_thread_binding=True,
history_max_messages=100,
),
"telegram": ThreadCapabilities(
supports_topics=True,
supports_history_fetch=True,
supports_thread_binding=True,
),
"whatsapp": ThreadCapabilities(
supports_topics=True,
supports_history_fetch=True,
),
"mattermost": ThreadCapabilities(
supports_reply_chains=True,
supports_history_fetch=True,
supports_thread_binding=True,
),
"irc": ThreadCapabilities(
supports_simulated_threads=True,
supports_history_fetch=False,
history_max_messages=0,
),
"nostr": ThreadCapabilities(
supports_simulated_threads=True,
supports_history_fetch=False,
history_max_messages=0,
),
"qqbot": ThreadCapabilities(
supports_simulated_threads=True,
supports_history_fetch=False,
history_max_messages=0,
),
"twitch": ThreadCapabilities(
supports_simulated_threads=True,
supports_history_fetch=False,
history_max_messages=0,
),
"msteams": ThreadCapabilities(
supports_native_threads=True,
supports_history_fetch=False,
supports_parent_injection=False,
supports_thread_binding=True,
),
"wechat": ThreadCapabilities(
supports_simulated_threads=True,
supports_history_fetch=False,
supports_parent_injection=False,
),
}

View File

@ -39,3 +39,18 @@ class DeliveryFailedError(ChannelException):
def __init__(self, message: str = ""):
detail = f"Message delivery failed: {message}" if message else "Message delivery failed"
super().__init__(detail, retryable=True, retry_after_ms=3000)
class ChannelTimeoutError(ChannelException):
def __init__(self, message: str = "Channel operation timed out"):
super().__init__(message, retryable=True, retry_after_ms=2000)
class ChannelQuotaExceededError(ChannelException):
def __init__(self, message: str = "Channel quota exceeded"):
super().__init__(message, retryable=False)
class MessageTooLargeError(ChannelException):
def __init__(self, message: str = "Message body exceeds size limit"):
super().__init__(message, retryable=False)

View File

@ -0,0 +1,100 @@
from __future__ import annotations
import hashlib
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING
from cachetools import TTLCache
if TYPE_CHECKING:
from yuxi.channels.models import FetchOptions, HistoricalMessage
class HistoryFetcher(ABC):
@abstractmethod
async def fetch_thread_history(
self,
thread_id: str,
options: FetchOptions,
) -> list[HistoricalMessage]:
"""获取线程历史消息,按时间升序"""
@abstractmethod
async def fetch_parent_message(
self,
thread_id: str,
) -> HistoricalMessage | None:
"""获取线程根消息(父消息)"""
class HistoryFormatter:
FORMAT_TEMPLATES: dict[str, dict[str, str]] = {
"xml": {
"header": "[Thread history]\n",
"message": " [{sender}] {content}\n",
"footer": "[/Thread history]\n",
},
"markdown": {
"header": "**Thread History**\n\n",
"message": "> **{sender}**: {content}\n\n",
"footer": "",
},
"compact": {
"header": "",
"message": "{sender}: {content}\n",
"footer": "",
},
}
def __init__(self, format_type: str = "xml"):
self.template = self.FORMAT_TEMPLATES.get(format_type, self.FORMAT_TEMPLATES["xml"])
def format_history(
self,
messages: list[HistoricalMessage],
max_chars: int = 4000,
) -> str:
if not messages:
return ""
result: list[str] = [self.template["header"]]
current_chars = len(result[0])
for msg in reversed(messages):
content = msg.content[:200] if len(msg.content) > 200 else msg.content
line = self.template["message"].format(sender=msg.sender_name, content=content)
if current_chars + len(line) > max_chars:
break
result.insert(1, line)
current_chars += len(line)
result.append(self.template["footer"])
return "".join(result)
class HistoryCache:
def __init__(self, max_size: int = 100, ttl_seconds: int = 300):
self._cache: TTLCache[str, list[HistoricalMessage]] = TTLCache(maxsize=max_size, ttl=ttl_seconds)
self._thread_keys: dict[str, set[str]] = {}
def _make_key(self, thread_id: str, options: FetchOptions) -> str:
key_data = f"{thread_id}:{options.max_messages}:{options.before_message_id}"
return hashlib.md5(key_data.encode()).hexdigest()
def get(self, thread_id: str, options: FetchOptions) -> list[HistoricalMessage] | None:
key = self._make_key(thread_id, options)
return self._cache.get(key)
def set(self, thread_id: str, options: FetchOptions, messages: list[HistoricalMessage]) -> None:
key = self._make_key(thread_id, options)
self._cache[key] = messages
if thread_id not in self._thread_keys:
self._thread_keys[thread_id] = set()
self._thread_keys[thread_id].add(key)
def invalidate(self, thread_id: str) -> None:
keys = self._thread_keys.pop(thread_id, set())
for key in keys:
self._cache.pop(key, None)

View File

@ -8,16 +8,17 @@ from typing import Any
from sqlalchemy import func, select
from yuxi.channels.base import BaseChannelAdapter
from yuxi.channels.exceptions import ChannelException
from yuxi.channels.infra.broadcast import EventBroadcaster
from yuxi.channels.infra.circuit_breaker import CircuitBreaker
from yuxi.channels.infra.circuit_breaker import CircuitBreaker, CircuitBreakerOpenError
from yuxi.channels.infra.config_watcher import ConfigWatcher
from yuxi.channels.models import ChannelStatus
from yuxi.channels.registry import _BUILTIN_ADAPTERS, ChannelRegistry, _load_builtin_adapters
from yuxi.channels.router import MessageRouter
from yuxi.channels.services.context import GatewayRequestContext
from yuxi.channels.services.doctor import ConfigDoctor, DiagnosisIssue
from yuxi.channels.exceptions import ChannelException
from yuxi.channels.services.maintenance import MaintenanceRunner
from yuxi.channels.models import ChannelStatus
from yuxi.channels.registry import ChannelRegistry, _BUILTIN_ADAPTERS, _load_builtin_adapters
from yuxi.channels.router import MessageRouter
from yuxi.channels.services.plugin_state_store import PostgresPluginStateStore
from yuxi.channels.services.runtime_state import RuntimeState
from yuxi.channels.services.stats_collector import StatsCollector
from yuxi.channels.services.webhook_registry import WebhookRegistry
@ -47,13 +48,15 @@ class ChannelManager:
self._broadcaster: EventBroadcaster | None = None
self._ctx: GatewayRequestContext | None = None
self.runtime_state = RuntimeState()
self._auth_limiter: Any = None
self._ws_handlers: dict[str, Any] = {}
self._scheduled_tasks: list[asyncio.Task] = []
self._ws_logger: WsLogger | None = None
self._maintenance_runner: MaintenanceRunner | None = None
self._stats_collector: StatsCollector | None = None
self._webhook_registry: WebhookRegistry | None = None
self._state_store: PostgresPluginStateStore | None = None
self._ws_broadcast: Any = None
self._prev_statuses: dict[str, str] = {}
@property
def phase(self) -> str:
@ -75,13 +78,19 @@ class ChannelManager:
def doctor(self) -> ConfigDoctor | None:
return self._doctor
async def load_config(self) -> None:
if self._phase not in ("not_started",):
return
def _register_all_adapters(self) -> None:
from yuxi.channels.message_actions import ActionRegistry
_load_builtin_adapters()
for channel_id, adapter_cls in _BUILTIN_ADAPTERS.items():
self._registry.register(channel_id, adapter_cls)
ActionRegistry.register_adapter(adapter_cls)
async def load_config(self) -> None:
if self._phase not in ("not_started",):
return
self._register_all_adapters()
from yuxi import config as conf
@ -126,9 +135,7 @@ class ChannelManager:
if self._initialized:
return
_load_builtin_adapters()
for channel_id, adapter_cls in _BUILTIN_ADAPTERS.items():
self._registry.register(channel_id, adapter_cls)
self._register_all_adapters()
from yuxi import config as conf
@ -177,7 +184,6 @@ class ChannelManager:
await self._stage_load_config()
await self._stage_prepare_bootstrap()
await self._stage_start_early_runtime()
await self._stage_create_auth_limiter()
await self._stage_init_channels()
await self._stage_create_runtime_state()
await self._stage_start_runtime_services()
@ -186,7 +192,7 @@ class ChannelManager:
await self._stage_start_event_subscriptions()
self._initialized = True
self._phase = "fully_running"
logger.info("ChannelManager: 10-stage startup complete")
logger.info("ChannelManager: 9-stage startup complete")
except Exception:
logger.exception(f"ChannelManager startup failed at phase: {self._phase}")
raise
@ -195,9 +201,7 @@ class ChannelManager:
if self._phase not in ("not_started",):
return
_load_builtin_adapters()
for channel_id, adapter_cls in _BUILTIN_ADAPTERS.items():
self._registry.register(channel_id, adapter_cls)
self._register_all_adapters()
from yuxi import config as conf
@ -232,17 +236,10 @@ class ChannelManager:
)
self._phase = "early_runtime"
logger.info("ChannelManager: [3/10] early runtime started")
async def _stage_create_auth_limiter(self) -> None:
if self._phase not in ("early_runtime",):
return
self._phase = "auth_limiter_created"
logger.info("ChannelManager: [4/10] auth rate limiter created")
logger.info("ChannelManager: [3/9] early runtime started")
async def _stage_init_channels(self) -> None:
if self._phase not in ("early_runtime", "auth_limiter_created"):
if self._phase not in ("early_runtime",):
return
for channel_id in self._registry.list_channels():
@ -254,7 +251,7 @@ class ChannelManager:
logger.exception(f"Failed to start channel {channel_id}")
self._phase = "channels_started"
logger.info(f"ChannelManager: [5/10] channels started: {list(self._adapters.keys())}")
logger.info(f"ChannelManager: [4/9] channels started: {list(self._adapters.keys())}")
async def _stage_create_runtime_state(self) -> None:
if self._phase not in ("channels_started",):
@ -268,10 +265,11 @@ class ChannelManager:
main_node=True,
services=["doctor", "maintenance", "stats", "webhooks"],
)
self._state_store = PostgresPluginStateStore()
self._ws_logger = WsLogger(max_entries=1000)
self._phase = "runtime_state_created"
logger.info("ChannelManager: [6/10] runtime state created")
logger.info("ChannelManager: [5/9] runtime state created")
async def _stage_start_runtime_services(self) -> None:
if self._phase not in ("runtime_state_created",):
@ -283,7 +281,7 @@ class ChannelManager:
self._webhook_registry = WebhookRegistry(self)
self._phase = "runtime_services_started"
logger.info("ChannelManager: [7/10] runtime services started (doctor+maintenance+stats+webhooks)")
logger.info("ChannelManager: [6/9] runtime services started (doctor+maintenance+stats+webhooks)")
async def _stage_activate_scheduled_services(self) -> None:
if self._phase not in ("runtime_services_started",):
@ -297,7 +295,7 @@ class ChannelManager:
self._scheduled_tasks.append(asyncio.create_task(self._webhook_registry.run()))
self._phase = "scheduled_services_active"
logger.info("ChannelManager: [8/10] scheduled services activated (watcher+maintenance+stats+webhooks)")
logger.info("ChannelManager: [7/9] scheduled services activated (watcher+maintenance+stats+webhooks)")
async def _stage_attach_ws_handlers(self) -> None:
if self._phase not in ("scheduled_services_active",):
@ -311,7 +309,7 @@ class ChannelManager:
self._broadcaster.subscribe_callback("config.reload", self._on_config_reload)
self._phase = "ws_handlers_attached"
logger.info("ChannelManager: [9/10] websocket handlers attached")
logger.info("ChannelManager: [8/9] websocket handlers attached")
async def _stage_start_event_subscriptions(self) -> None:
if self._phase not in ("ws_handlers_attached",):
@ -320,13 +318,39 @@ class ChannelManager:
self.runtime_state.services = ["doctor", "maintenance", "stats", "webhooks", "watcher", "ws"]
self._phase = "subscriptions_started"
logger.info("ChannelManager: [10/10] event subscriptions started")
logger.info("ChannelManager: [9/9] event subscriptions started")
def set_ws_broadcast(self, cb) -> None:
self._ws_broadcast = cb
async def _push_channel_status_to_ws(self, channel_id: str, status: str, health: dict | None = None) -> None:
if not self._ws_broadcast:
return
payload: dict[str, Any] = {"channel_id": channel_id, "status": status}
if health:
payload["health"] = health
try:
await self._ws_broadcast(
{
"type": "channel_status",
"payload": payload,
"timestamp": _now_iso(),
}
)
except Exception:
pass
async def _on_channel_status_change(self, event: str, payload: Any) -> None:
logger.debug(f"WS event: {event}")
logger.debug(f"channel.status_change: {event}")
if isinstance(payload, dict):
channel_id = payload.get("channel_id")
status = payload.get("status")
health = payload.get("health")
if channel_id and status:
await self._push_channel_status_to_ws(channel_id, status, health)
async def _on_tick(self, event: str, payload: Any) -> None:
pass
pass # TODO: implement periodic tick logic (health checks, stats flush, etc.)
async def _on_chat_event(self, event: str, payload: Any) -> None:
logger.debug(f"WS chat event: {event}")
@ -374,6 +398,7 @@ class ChannelManager:
config = config or self._channels_config.get(channel_id, {})
adapter = adapter_cls(config=config)
adapter._state_store = self._state_store
adapter.on_message(self._handle_inbound_message)
pre_connect_result = await adapter.pre_connect()
@ -388,6 +413,18 @@ class ChannelManager:
if channel_id not in self._health_tasks:
self._health_tasks[channel_id] = asyncio.create_task(self._health_check_loop(channel_id))
current_status = self._adapter_status(adapter)
self._prev_statuses[channel_id] = current_status
if self._broadcaster:
await self._broadcaster.broadcast(
"channel.status_change",
{
"channel_id": channel_id,
"status": current_status,
"health": None,
},
)
logger.info(f"Channel {channel_id} started")
async def stop_channel(self, channel_id: str) -> None:
@ -408,6 +445,17 @@ class ChannelManager:
except Exception:
logger.exception(f"Error disconnecting channel {channel_id}")
self._prev_statuses.pop(channel_id, None)
if self._broadcaster:
await self._broadcaster.broadcast(
"channel.status_change",
{
"channel_id": channel_id,
"status": ChannelStatus.DISCONNECTED.value,
"health": None,
},
)
logger.info(f"Channel {channel_id} stopped")
async def restart_channel(self, channel_id: str) -> None:
@ -415,21 +463,53 @@ class ChannelManager:
await self.stop_channel(channel_id)
await self.start_channel(channel_id, config)
async def register_channel(self, channel_id: str, config: dict[str, Any] | None = None) -> None:
adapter_cls = self._registry.get(channel_id)
if not adapter_cls:
raise ValueError(f"Unknown channel type: '{channel_id}'")
if channel_id in self._channels_config:
raise ValueError(f"Channel '{channel_id}' is already registered")
self._channels_config[channel_id] = {**(config or {}), "enabled": True}
logger.info(f"Channel '{channel_id}' registered")
async def unregister_channel(self, channel_id: str) -> None:
await self.stop_channel(channel_id)
self._registry.unregister(channel_id)
self._channels_config.pop(channel_id, None)
self._circuit_breakers.pop(channel_id, None)
self._prev_statuses.pop(channel_id, None)
logger.info(f"Channel {channel_id} unregistered")
async def send_outbound(self, channel_id: str, response) -> None:
adapter = self._adapters.get(channel_id)
if not adapter:
raise ChannelException(f"Channel {channel_id} not found", retryable=False)
cb = self._circuit_breakers[channel_id]
await cb.call(lambda: adapter.send(response))
try:
await cb.call(lambda: adapter.send(response))
except CircuitBreakerOpenError:
raise ChannelException(
f"Channel {channel_id} temporarily unavailable",
retryable=True,
retry_after_ms=int(cb.recovery_timeout * 1000),
)
async def get_channel_status(self, channel_id: str | None = None) -> dict:
if channel_id:
return await self._get_single_channel_status(channel_id)
channel_ids = self._registry.list_channels()
if not channel_ids:
return {"channels": {}}
batch_stats = await self._get_batch_channel_stats(channel_ids)
all_channels = {}
for cid in self._registry.list_channels():
info = await self._get_single_channel_status(cid)
for cid in channel_ids:
info = await self._get_single_channel_status(cid, stats=batch_stats.get(cid))
stats = info.get("stats") or {}
all_channels[cid] = {
"channel_id": cid,
@ -446,20 +526,22 @@ class ChannelManager:
return {"channels": all_channels}
async def update_channel_config(self, channel_id: str, config_updates: dict[str, Any]) -> dict:
adapter = self._adapters.get(channel_id)
if not adapter:
raise ChannelException(f"Channel {channel_id} not found", retryable=False)
if hasattr(adapter, "reload_config"):
await adapter.reload_config(config_updates)
else:
for key, value in config_updates.items():
adapter.config[key] = value
registered_ids = set(self._registry.list_channels())
if channel_id not in registered_ids:
raise ChannelException(f"Channel {channel_id} is not registered", retryable=False)
if channel_id not in self._channels_config:
self._channels_config[channel_id] = {}
self._channels_config[channel_id].update(config_updates)
adapter = self._adapters.get(channel_id)
if adapter:
if hasattr(adapter, "reload_config"):
await adapter.reload_config(config_updates)
else:
for key, value in config_updates.items():
adapter.config[key] = value
return {"channel_id": channel_id, "config_updated": True}
async def test_channel(self, channel_id: str) -> dict:
@ -472,15 +554,27 @@ class ChannelManager:
health = await adapter.health_check()
latency_ms = (time.monotonic() - start) * 1000
return {
status_map = {"healthy": "success", "degraded": "degraded", "unhealthy": "failure"}
test_result = status_map.get(health.status, "failure")
result = {
"channel_id": channel_id,
"test_result": "success" if health.status == "healthy" else "degraded",
"test_result": test_result,
"latency_ms": round(latency_ms, 1),
"health": health.model_dump(),
}
if test_result == "failure" and health.status == "unhealthy":
result["error"] = f"Channel unhealthy: {health.last_error or 'unknown'}"
return result
except Exception as e:
return {"channel_id": channel_id, "test_result": "failure", "error": str(e)}
def is_registered(self, channel_id: str) -> bool:
return channel_id in set(self._registry.list_channels())
def is_running(self, channel_id: str) -> bool:
return channel_id in self._adapters
async def check_rate_limit(self, key: str, max_req: int, window_seconds: int) -> bool:
lock = self._rate_limit_locks.setdefault(key, asyncio.Lock())
async with lock:
@ -506,7 +600,11 @@ class ChannelManager:
if not adapter:
break
prev_status = self._prev_statuses.get(channel_id)
current_status = self._adapter_status(adapter)
cb = self._circuit_breakers.get(channel_id)
health = None
try:
health = await adapter.health_check()
if cb and health.status == "healthy":
@ -529,7 +627,99 @@ class ChannelManager:
except Exception:
pass
async def _get_single_channel_status(self, channel_id: str) -> dict:
if current_status != prev_status and self._broadcaster:
self._prev_statuses[channel_id] = current_status
health_dict = health.model_dump() if health else None
await self._broadcaster.broadcast(
"channel.status_change",
{
"channel_id": channel_id,
"status": current_status,
"health": health_dict,
},
)
async def _get_batch_channel_stats(self, channel_ids: list[str]) -> dict[str, dict]:
if not channel_ids:
return {}
try:
from yuxi.utils.datetime_utils import utc_now_naive
now = utc_now_naive()
today_start = now.replace(hour=0, minute=0, second=0, microsecond=0)
async with pg_manager.get_async_session_context() as session:
totals_result = await session.execute(
select(
ChannelMsgRecord.channel_id,
func.count().label("total"),
func.sum(
func.cast(
(ChannelMsgRecord.status == "success").cast(func.Integer),
func.Integer,
)
).label("success_count"),
func.sum(
func.cast(
(ChannelMsgRecord.status == "error").cast(func.Integer),
func.Integer,
)
).label("error_count"),
)
.where(ChannelMsgRecord.channel_id.in_(channel_ids))
.group_by(ChannelMsgRecord.channel_id)
)
totals = {
r.channel_id: (r.total or 0, r.success_count or 0, r.error_count or 0) for r in totals_result.all()
}
today_result = await session.execute(
select(
ChannelMsgRecord.channel_id,
func.count().label("today_count"),
)
.where(
ChannelMsgRecord.channel_id.in_(channel_ids),
ChannelMsgRecord.created_at >= today_start,
)
.group_by(ChannelMsgRecord.channel_id)
)
today_counts = {r.channel_id: r.today_count for r in today_result.all()}
stats_map = {}
for cid in channel_ids:
if cid in totals:
total, success, error = totals[cid]
stats_map[cid] = {
"total_messages": total,
"today_messages": today_counts.get(cid, 0),
"success_count": int(success),
"error_count": int(error),
"success_rate": round(success / total, 3) if total > 0 else 0,
}
else:
stats_map[cid] = {
"total_messages": 0,
"today_messages": 0,
"success_count": 0,
"error_count": 0,
"success_rate": 0,
}
return stats_map
except Exception:
return {
cid: {
"total_messages": 0,
"today_messages": 0,
"success_count": 0,
"error_count": 0,
"success_rate": 0,
}
for cid in channel_ids
}
async def _get_single_channel_status(self, channel_id: str, stats: dict | None = None) -> dict:
adapter = self._adapters.get(channel_id)
if not adapter:
adapter_cls = self._registry.get(channel_id)
@ -545,16 +735,17 @@ class ChannelManager:
}
)
channel_type = adapter_cls.channel_type.value
saved_config = self._channels_config.get(channel_id, {})
return {
"channel_id": channel_id,
"channel_type": channel_type,
"display_name": None,
"enabled": False,
"display_name": saved_config.get("display_name"),
"enabled": saved_config.get("enabled", False),
"status": ChannelStatus.DISABLED.value,
"config": {"enabled": False},
"config": saved_config if saved_config else {"enabled": False},
"capabilities": caps,
"health": None,
"stats": None,
"stats": stats if stats is not None else None,
}
return {"channel_id": channel_id, "status": "not_found"}
@ -592,7 +783,7 @@ class ChannelManager:
"capabilities": caps,
"health": health,
"circuit_state": cb_state,
"stats": await self._get_channel_stats(channel_id),
"stats": stats if stats is not None else await self._get_channel_stats(channel_id),
}
async def _get_channel_stats(self, channel_id: str) -> dict:
@ -650,7 +841,10 @@ class ChannelManager:
}
def _adapter_status(self, adapter: BaseChannelAdapter) -> str:
return getattr(adapter, "status", "unknown")
_status = getattr(adapter, "_status", None)
if _status is None:
return "unknown"
return _status.value if hasattr(_status, "value") else str(_status)
async def _ensure_virtual_department(self, db) -> None:
from yuxi.storage.postgres.models_business import Department
@ -678,4 +872,17 @@ class ChannelManager:
logger.info("Ensured default agent config for ChatbotAgent")
channel_manager = ChannelManager()
_channel_manager: ChannelManager | None = None
def get_channel_manager() -> ChannelManager:
global _channel_manager
if _channel_manager is None:
_channel_manager = ChannelManager()
return _channel_manager
def _now_iso() -> str:
from yuxi.utils.datetime_utils import format_utc_datetime, utc_now_naive
return format_utc_datetime(utc_now_naive())

View File

@ -40,33 +40,13 @@ class OutboundMixin:
return chunks
class MediaMixin:
"""媒体发送能力 Mixin"""
async def send_media(self, chat_id: str, media_type: str, data: Any) -> DeliveryResult:
raise NotImplementedError
class ReactionMixin:
"""表情回应能力 Mixin"""
async def send_reaction(self, chat_id: str, msg_id: str, emoji: str) -> DeliveryResult:
raise NotImplementedError
async def remove_reaction(self, chat_id: str, msg_id: str, emoji: str) -> DeliveryResult:
raise NotImplementedError
class EditDeleteMixin:
"""消息编辑/删除能力 Mixin"""
async def edit_message(self, chat_id: str, msg_id: str, content: str) -> DeliveryResult:
raise NotImplementedError
async def delete_message(self, chat_id: str, msg_id: str) -> DeliveryResult:
raise NotImplementedError
class PinMixin:
"""消息置顶/取消置顶能力 Mixin"""
@ -100,13 +80,6 @@ class PollMixin:
class StreamingMixin:
"""流式输出能力 Mixin"""
async def send_stream_chunk(self, chat_id: str, msg_id: str, chunk: str, finished: bool) -> DeliveryResult:
raise NotImplementedError
@property
def streaming_modes(self) -> list[str]:
return ["off"]
@property
def block_streaming(self) -> bool:
return False

View File

@ -9,6 +9,53 @@ from pydantic import BaseModel, ConfigDict, Field
from yuxi.utils.datetime_utils import utc_now_naive
class ThreadType(StrEnum):
NATIVE = "native"
TOPIC = "topic"
REPLY_CHAIN = "reply_chain"
SIMULATED = "simulated"
DIRECT = "direct"
GROUP = "group"
CHANNEL = "channel"
class SessionScope(StrEnum):
DIRECT = "dm"
GROUP = "group"
GROUP_SENDER = "group_sender"
TOPIC = "topic"
TOPIC_SENDER = "topic_sender"
THREAD = "thread"
class ThreadContext(BaseModel):
thread_id: str
thread_type: ThreadType = ThreadType.DIRECT
parent_id: str | None = None
root_message_id: str | None = None
participants: list[str] = []
created_at: datetime | None = None
metadata: dict[str, Any] = {}
class HistoricalMessage(BaseModel):
message_id: str
sender_id: str
sender_name: str
content: str
timestamp: datetime
is_from_bot: bool = False
reply_to_id: str | None = None
class FetchOptions(BaseModel):
max_messages: int = 50
max_chars: int = 4000
include_bot_messages: bool = True
before_message_id: str | None = None
after_message_id: str | None = None
class MessageType(StrEnum):
TEXT = "text"
IMAGE = "image"
@ -78,10 +125,14 @@ class EventType(StrEnum):
MESSAGE_RECEIVED = "message.received"
MESSAGE_UPDATED = "message.updated"
MESSAGE_DELETED = "message.deleted"
MESSAGES_DELETED = "messages.deleted"
BOT_ADDED = "bot.added"
BOT_REMOVED = "bot.removed"
MEMBER_JOINED = "member.joined"
MEMBER_LEFT = "member.left"
MEMBER_ADDED = "member.added"
MEMBER_REMOVED = "member.removed"
MEMBER_UPDATED = "member.updated"
CARD_ACTION = "card.action"
REACTION_ADDED = "reaction.added"
REACTION_REMOVED = "reaction.removed"
@ -89,6 +140,13 @@ class EventType(StrEnum):
TYPING = "typing"
READ_RECEIPT = "read_receipt"
SYSTEM_EVENT = "system.event"
ROLE_CREATED = "role.created"
ROLE_DELETED = "role.deleted"
ROLE_UPDATED = "role.updated"
CHANNEL_CREATED = "channel.created"
CHANNEL_UPDATED = "channel.updated"
CHANNEL_DELETED = "channel.deleted"
INTERACTION = "interaction"
class RejectReason(StrEnum):

View File

@ -5,7 +5,7 @@ from typing import Any
from yuxi.channels.base import BaseChannelAdapter
from yuxi.channels.capabilities import ChannelCapabilities
from yuxi.channels.meta import ChannelMeta
from yuxi.channels.registry import _BUILTIN_ADAPTERS
from yuxi.channels.registry import _register_builtin
def channel_plugin(
@ -47,7 +47,7 @@ def channel_plugin(
if meta is not None:
_cls.meta = meta # type: ignore[attr-defined]
_BUILTIN_ADAPTERS[cid] = _cls
_register_builtin(cid, _cls)
return _cls
if cls is not None:
@ -71,6 +71,9 @@ class ChannelPlugin:
streaming_modes=["off", "partial", "block", "progress"],
),
meta=ChannelMeta(id="telegram", label="Telegram"),
pairing={"auto_pair": True},
conversation_bindings={"max_bindings": 5},
agent_prompt="You are a Telegram bot",
)
@plugin.register
@ -108,6 +111,18 @@ class ChannelPlugin:
cls.capabilities = self.capabilities # type: ignore[attr-defined]
if self.meta is not None:
cls.meta = self.meta # type: ignore[attr-defined]
if self.channel_type is not None:
cls.channel_type = self.channel_type # type: ignore[attr-defined]
if self.pairing:
cls.pairing = self.pairing # type: ignore[attr-defined]
if self.conversation_bindings:
cls.conversation_bindings = self.conversation_bindings # type: ignore[attr-defined]
if self.agent_prompt is not None:
cls.agent_prompt = self.agent_prompt # type: ignore[attr-defined]
if self.messaging:
cls.messaging = self.messaging # type: ignore[attr-defined]
if self.directory is not None:
cls.directory = self.directory # type: ignore[attr-defined]
_BUILTIN_ADAPTERS[self.channel_id] = cls
_register_builtin(self.channel_id, cls)
return cls

View File

@ -6,12 +6,17 @@ _BUILTIN_ADAPTERS: dict[str, type[BaseChannelAdapter]] = {}
_BUILTIN_ADAPTER_ALIASES: dict[str, str] = {}
def _register_builtin(channel_id: str, cls: type[BaseChannelAdapter], aliases: list[str] | None = None) -> None:
"""统一的内置适配器注册入口 — 所有注册路径最终都调用此函数"""
_BUILTIN_ADAPTERS[channel_id] = cls
if aliases:
for alias in aliases:
_BUILTIN_ADAPTER_ALIASES[alias] = channel_id
def register_builtin_adapter(cls: type[BaseChannelAdapter] | None = None, *, aliases: list[str] | None = None):
def _decorator(cls_inner: type[BaseChannelAdapter]) -> type[BaseChannelAdapter]:
_BUILTIN_ADAPTERS[cls_inner.channel_id] = cls_inner
if aliases:
for alias in aliases:
_BUILTIN_ADAPTER_ALIASES[alias] = cls_inner.channel_id
_register_builtin(cls_inner.channel_id, cls_inner, aliases)
return cls_inner
if cls is None:
@ -209,5 +214,11 @@ class ChannelRegistry:
def unregister(self, channel_id: str) -> None:
self._adapters.pop(channel_id, None)
def load_builtins(self) -> None:
"""将 _BUILTIN_ADAPTERS 中的所有内置适配器加载到 self._adapters 中"""
_load_builtin_adapters()
for channel_id, adapter_cls in _BUILTIN_ADAPTERS.items():
self._adapters.setdefault(channel_id, adapter_cls)
BUILTIN_ADAPTERS = _BUILTIN_ADAPTERS

View File

@ -1,17 +1,20 @@
from __future__ import annotations
import asyncio
from datetime import datetime, timezone
from yuxi.channels.services.context import ChatAbortEntry, ChatRunBuffer
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 GroupChatPolicy
from yuxi.channels.policy.schedule_policy import SchedulePolicy
from yuxi.channels.policy.welcome_policy import WelcomePolicy
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.session_mapper import SessionMapper, VIRTUAL_DEPARTMENT_ID
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
@ -40,14 +43,89 @@ class MessageRouter:
self._channel_manager = channel_manager
self.dedup_policy = dedup_policy or DedupPolicy()
self.context_policy = context_policy or ContextPolicy()
self.group_chat_policy = group_chat_policy or GroupChatPolicy()
self.welcome_policy = welcome_policy or WelcomePolicy()
self.schedule_policy = schedule_policy or SchedulePolicy()
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
@ -55,24 +133,47 @@ class MessageRouter:
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
if not self.schedule_policy.is_working_hours():
reply = self.schedule_policy.get_off_hours_reply()
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 self.group_chat_policy.should_respond(message, is_at_bot):
if not group_chat_policy.should_respond(message, is_at_bot):
return
from yuxi.storage.postgres.manager import pg_manager
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)
@ -80,13 +181,13 @@ class MessageRouter:
thread_id = await session_mapper.resolve_thread(message, internal_user_id)
msg_record_repo = ChannelMessageRecordRepository(db)
agent_config_id = self._resolve_agent_config_id(message)
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 self.welcome_policy.mark_welcomed(internal_user_id):
if welcome_policy.mark_welcomed(internal_user_id):
welcome_response = ChannelResponse(
identity=identity,
content=self.welcome_policy.get_welcome_message(),
content=welcome_policy.get_welcome_message(),
)
await self._send_response(identity.channel_id, welcome_response)
@ -105,13 +206,17 @@ class MessageRouter:
)
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)
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}")
@ -130,6 +235,7 @@ class MessageRouter:
)
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(
@ -164,22 +270,142 @@ class MessageRouter:
await self._send_response(identity.channel_id, response)
elif command == ContextCommand.HISTORY:
response = ChannelResponse(
identity=identity, content="\u5386\u53f2\u8bb0\u5f55\u529f\u80fd\u6682\u672a\u5b9e\u73b0"
)
await self._send_response(identity.channel_id, response)
await self._cmd_history(message)
elif command == ContextCommand.CONTEXT:
response = ChannelResponse(
identity=identity, content="\u4e0a\u4e0b\u6587\u4fe1\u606f\u529f\u80fd\u6682\u672a\u5b9e\u73b0"
)
await self._send_response(identity.channel_id, response)
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="\u5bf9\u8bdd\u6458\u8981\u529f\u80fd\u6682\u672a\u5b9e\u73b0"
identity=identity,
content=f"\u83b7\u53d6\u5386\u53f2\u8bb0\u5f55\u5931\u8d25\uff1a{str(e)[:100]}",
)
await self._send_response(identity.channel_id, response)
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,
@ -252,7 +478,7 @@ class MessageRouter:
finally:
self.chat_run_buffers.pop(run_id, None)
def _resolve_agent_config_id(self, message: ChannelMessage) -> int:
async def _resolve_agent_config_id(self, message: ChannelMessage, db=None) -> int:
channel_id = message.identity.channel_id
content = message.content.strip()
@ -268,9 +494,29 @@ class MessageRouter:
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:
return int(channel_default)
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:
@ -278,6 +524,15 @@ class MessageRouter:
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, {})
@ -291,3 +546,15 @@ class MessageRouter:
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()

View File

@ -3,12 +3,13 @@ from __future__ import annotations
import uuid as uuid_lib
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from yuxi.channels.models import ChannelMessage
from yuxi.storage.postgres.models_channels import ChannelUserMapping, ChannelThreadMapping
from yuxi.storage.postgres.models_business import User
from yuxi.storage.postgres.models_channels import ChannelThreadMapping, ChannelUserMapping
from yuxi.utils.datetime_utils import utc_now_naive
from yuxi.utils.logging_config import logger
@ -28,32 +29,41 @@ class SessionMapper:
return mapping.internal_user_id
internal_user_id = f"ch_{identity.channel_id}_{uuid_lib.uuid4().hex[:8]}"
username = f"{identity.channel_id}_{identity.channel_user_id}"
source = f"{USER_SOURCE_PREFIX}{identity.channel_id}"
try:
user = User(
username=f"{identity.channel_id}_{identity.channel_user_id}",
stmt = (
pg_insert(User)
.values(
username=username,
user_id=internal_user_id,
password_hash="",
role="user",
department_id=self.department_id,
source=source,
)
self.db.add(user)
await self.db.flush()
.on_conflict_do_update(
index_elements=["username"],
set_={"username": username},
)
.returning(User.user_id)
)
result = await self.db.execute(stmt)
actual_user_id = result.scalar_one()
try:
mapping = ChannelUserMapping(
channel_id=identity.channel_id,
channel_user_id=identity.channel_user_id,
internal_user_id=internal_user_id,
internal_user_id=actual_user_id,
)
self.db.add(mapping)
await self.db.commit()
logger.info(
f"Created channel user mapping: {identity.channel_id}/{identity.channel_user_id} -> {internal_user_id}"
f"Created channel user mapping: {identity.channel_id}/{identity.channel_user_id} -> {actual_user_id}"
)
return internal_user_id
return actual_user_id
except IntegrityError:
await self.db.rollback()

View File

@ -0,0 +1,123 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass, field
from datetime import datetime, timedelta
from enum import StrEnum
from typing import Any
from yuxi.utils.datetime_utils import utc_now_naive
class BindingType(StrEnum):
AGENT = "agent"
SUBAGENT = "subagent"
ACP = "acp"
CONVERSATION = "conversation"
@dataclass
class ThreadBinding:
thread_id: str
binding_type: BindingType
target_id: str
created_at: datetime = field(default_factory=utc_now_naive)
expires_at: datetime | None = None
metadata: dict[str, Any] = field(default_factory=dict)
@property
def is_expired(self) -> bool:
if self.expires_at is None:
return False
return utc_now_naive() > self.expires_at
class ThreadBindingManager:
def __init__(
self,
default_ttl_hours: int = 24,
):
self._bindings: dict[str, ThreadBinding] = {}
self._default_ttl = timedelta(hours=default_ttl_hours) if default_ttl_hours else None
self._listeners: list[Callable] = []
def bind(
self,
thread_id: str,
binding_type: BindingType,
target_id: str,
ttl_hours: int | None = None,
metadata: dict[str, Any] | None = None,
) -> ThreadBinding:
expires = None
if ttl_hours is not None:
expires = utc_now_naive() + timedelta(hours=ttl_hours)
elif self._default_ttl:
expires = utc_now_naive() + self._default_ttl
binding = ThreadBinding(
thread_id=thread_id,
binding_type=binding_type,
target_id=target_id,
expires_at=expires,
metadata=metadata or {},
)
self._bindings[f"{thread_id}:{binding_type.value}"] = binding
self._notify("bind", binding)
return binding
def unbind(self, thread_id: str, binding_type: BindingType) -> bool:
key = f"{thread_id}:{binding_type.value}"
if key in self._bindings:
binding = self._bindings.pop(key)
self._notify("unbind", binding)
return True
return False
def get_binding(self, thread_id: str, binding_type: BindingType) -> ThreadBinding | None:
key = f"{thread_id}:{binding_type.value}"
binding = self._bindings.get(key)
if binding and binding.is_expired:
self.unbind(thread_id, binding_type)
return None
return binding
def list_bindings(
self,
thread_id: str | None = None,
binding_type: BindingType | None = None,
) -> list[ThreadBinding]:
results: list[ThreadBinding] = []
expired_keys: list[str] = []
for key, binding in self._bindings.items():
if binding.is_expired:
expired_keys.append(key)
continue
if thread_id and binding.thread_id != thread_id:
continue
if binding_type and binding.binding_type != binding_type:
continue
results.append(binding)
for key in expired_keys:
self._bindings.pop(key, None)
return results
def add_listener(self, listener: Callable[[str, ThreadBinding], None]) -> None:
self._listeners.append(listener)
def remove_listener(self, listener: Callable[[str, ThreadBinding], None]) -> None:
if listener in self._listeners:
self._listeners.remove(listener)
def _notify(self, event: str, binding: ThreadBinding) -> None:
for listener in self._listeners:
try:
listener(event, binding)
except Exception:
pass