ForcePilot/backend/package/yuxi/channel/message/langfuse_trace.py

203 lines
6.9 KiB
Python
Raw Normal View History

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