import logging import time from typing import Any from yuxi.channel.message.models import DispatchResult, UnifiedMessage from yuxi.services.langfuse_service import ( get_langfuse_client, is_langfuse_enabled, ) logger = logging.getLogger(__name__) class ChannelTrace: """渠道消息处理链路的 Langfuse trace 包装器。 复用现有 langfuse_service 的 Langfuse 客户端,创建实际 Langfuse span 以追踪消息处理每个阶段(分发、校验、规则匹配等)的耗时和状态。 """ def __init__(self, msg: UnifiedMessage, dispatch_result: DispatchResult | None = None): self._msg = msg self._dispatch_result = dispatch_result self._client = get_langfuse_client() self._enabled = is_langfuse_enabled() and self._client is not None self._trace_id: str | None = None self._trace_obj: Any = None self._trace_start_time: float | None = None self._live_spans: dict[str, Any] = {} @property def trace_name(self) -> str: peer_kind = self._msg.sender.kind.value return f"channel:{self._msg.channel_type}:{peer_kind}" @property def trace_id(self) -> str | None: return self._trace_id @property def span_count(self) -> int: return len(self._live_spans) async def start(self) -> "ChannelTrace": if not self._enabled: return self try: self._trace_start_time = time.monotonic() self._trace_id = self._client.create_trace_id() self._trace_obj = self._client.trace( id=self._trace_id, name=self.trace_name, metadata=self._build_base_metadata(), tags=self._build_base_tags(), ) except Exception: logger.exception("Failed to create channel trace: %s", self.trace_name) self._enabled = False return self async def add_span( self, name: str, metadata: dict[str, Any] | None = None, input_data: Any = None, output_data: Any = None, level: str = "DEFAULT", status_message: str | None = None, parent_span_id: str | None = None, ) -> str | None: """创建 Langfuse span 并追踪其生命周期。 返回 span_id 供后续 end_span() 结束使用。 调用者必须在完成处理后调用 end_span() 或依赖 finish() 统一结束。 """ if not self._enabled or self._trace_obj is None: return None try: parent = self._live_spans.get(parent_span_id) if parent_span_id else self._trace_obj span = parent.span( name=name, input=input_data, output=output_data, metadata=metadata or {}, level=level, status_message=status_message, ) span_id: str = span.id self._live_spans[span_id] = span return span_id except Exception: logger.exception("Failed to create span '%s' for trace: %s", name, self.trace_name) return None async def end_span( self, span_id: str, output_data: Any = None, status_message: str | None = None, level: str | None = None, metadata: dict[str, Any] | None = None, ) -> None: """结束指定 span,记录输出与最终状态。""" span = self._live_spans.pop(span_id, None) if span is None: return try: update_kwargs: dict[str, Any] = {} if output_data is not None: update_kwargs["output"] = output_data if level is not None: update_kwargs["level"] = level if status_message is not None: update_kwargs["status_message"] = status_message if metadata is not None: update_kwargs["metadata"] = metadata if update_kwargs: span.update(**update_kwargs) span.end() except Exception: logger.exception("Failed to end span '%s' for trace: %s", span_id, self.trace_name) async def finish(self, error: str | None = None) -> None: if not self._enabled or self._trace_obj is None: return try: pending = dict(self._live_spans) self._live_spans.clear() for span_id, span in pending.items(): try: span.end() except Exception: logger.exception("Failed to end pending span '%s'", span_id) duration_ms = None if self._trace_start_time is not None: duration_ms = (time.monotonic() - self._trace_start_time) * 1000 metadata: dict[str, Any] = { **self._build_base_metadata(), "span_count": len(pending), "dispatch_success": error is None, **(self._build_dispatch_metadata() if self._dispatch_result else {}), } if duration_ms is not None: metadata["duration_ms"] = round(duration_ms, 2) self._trace_obj.update( output=error or "success", metadata=metadata, ) except Exception: logger.exception("Failed to finish channel trace: %s", self.trace_name) def _build_base_metadata(self) -> dict[str, Any]: return { "source": "channel", "channel_type": self._msg.channel_type, "account_id": self._msg.account_id, "peer_kind": self._msg.sender.kind.value, "peer_id": self._msg.sender.id, "msg_id": self._msg.msg_id, "message_type": self._msg.message_type.value, "feature": "channel_dispatch", } def _build_dispatch_metadata(self) -> dict[str, Any]: if self._dispatch_result is None: return {} return { "agent_config_id": str(self._dispatch_result.agent_config_id) if self._dispatch_result.agent_config_id else None, "session_key": getattr(self._dispatch_result, "session_key", None), "thread_id": self._dispatch_result.thread_id, "matched_by": getattr(self._dispatch_result, "matched_by", None), } def _build_base_tags(self) -> list[str]: tags = [ "yuxi", "channel", f"channel:{self._msg.channel_type}", f"peer_kind:{self._msg.sender.kind.value}", ] if self._dispatch_result and self._dispatch_result.agent_config_id: tags.append(f"agent_config:{self._dispatch_result.agent_config_id}") return tags async def create_channel_trace( msg: UnifiedMessage, dispatch_result: DispatchResult | None = None, ) -> ChannelTrace: trace = ChannelTrace(msg, dispatch_result) await trace.start() return trace