项目文件夹

文件
2026-07-13 13:22:34 +08:00

329 行
13 KiB
Python

import functools
import logging
import time
import uuid
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any
from autogen import Agent, ConversableAgent
from autogen.logger.base_logger import BaseLogger
from openai.types.chat import ChatCompletion
from mlflow.entities.span import NoOpSpan, Span, SpanType
from mlflow.entities.span_event import SpanEvent
from mlflow.entities.span_status import SpanStatus, SpanStatusCode
from mlflow.tracing.constant import SpanAttributeKey, TokenUsageKey
from mlflow.tracing.fluent import start_span_no_context
from mlflow.tracing.utils import capture_function_input_args
from mlflow.utils.autologging_utils import autologging_is_disabled
from mlflow.utils.autologging_utils.safety import safe_patch
# For GroupChat, a single "received_message" events are passed around multiple
# internal layers and thus too verbose if we show them all. Therefore we ignore
# some of the message senders listed below.
_EXCLUDED_MESSAGE_SENDERS = ["chat_manager", "checking_agent"]
_logger = logging.getLogger(__name__)
FLAVOR_NAME = "ag2"
@dataclass
class _PendingSpan:
"""A span waiting for parent relocation, with its end data stored."""
span: Span
outputs: Any
end_time_ns: int
@dataclass
class ChatState:
"""
Represents the state of a chat session.
"""
# The root span object that scopes the entire single chat session. All spans
# such as LLM, function calls, in the chat session should be children of this span.
session_span: Span | None = None
# The last message object in the chat session.
last_message: Any | None = None
# The timestamp (ns) of the last message in the chat session.
last_message_timestamp: int = 0
# LLM/Tool Spans created after the last message in the chat session.
# We consider them as operations for generating the next message and
# re-locate them under the corresponding message span.
# These spans are not ended yet to avoid premature export before parent relocation.
pending_spans: list[_PendingSpan] = field(default_factory=list)
def clear(self):
self.session_span = None
self.last_message = None
self.last_message_timestamp = 0
self.pending_spans = []
def _catch_exception(func):
def wrapper(*args, **kwargs):
try:
return func(*args, **kwargs)
except Exception as e:
_logger.error(f"Error occurred during AutoGen tracing: {e}")
return wrapper
class MlflowAg2Logger(BaseLogger):
def __init__(self):
self._chat_state = ChatState()
def start(self) -> str:
return "session_id"
@_catch_exception
def log_new_agent(self, agent: ConversableAgent, init_args: dict[str, Any]) -> None:
"""
This handler is called whenever a new agent instance is created.
Here we patch the agent's methods to start and end a trace around its chat session.
"""
# TODO: Patch generate_reply() method as well
if hasattr(agent, "initiate_chat"):
safe_patch(
FLAVOR_NAME,
agent.__class__,
"initiate_chat",
# Setting root_only = True because sometimes compounded agent calls initiate_chat()
# method of its sub-agents, which should not start a new trace.
self._get_patch_function(root_only=True),
)
if hasattr(agent, "register_function"):
def patched(original, _self, function_map, **kwargs):
original(_self, function_map, **kwargs)
# Wrap the newly registered tools to start and end a span around its invocation.
for name, f in function_map.items():
if f is not None:
_self._function_map[name] = functools.partial(
self._get_patch_function(span_type=SpanType.TOOL), f
)
safe_patch(FLAVOR_NAME, agent.__class__, "register_function", patched)
def _get_patch_function(self, span_type: str = SpanType.UNKNOWN, root_only: bool = False):
"""
Patch a function to start and end a span around its invocation.
Args:
f: The function to patch.
span_name: The name of the span. If None, the function name is used.
span_type: The type of the span. Default is SpanType.UNKNOWN.
root_only: If True, only create a span if it is the root of the chat session.
When there is an existing root span for the chat session, the function will
not create a new span.
"""
def _wrapper(original, *args, **kwargs):
# If autologging is disabled, just run the original function. This is a safety net to
# prevent patching side effects from being effective after autologging is disabled.
if autologging_is_disabled(FLAVOR_NAME):
return original(*args, **kwargs)
if self._chat_state.session_span is None:
# Create the trace per chat session
span = start_span_no_context(
name=original.__name__,
span_type=span_type,
inputs=capture_function_input_args(original, args, kwargs),
attributes={SpanAttributeKey.MESSAGE_FORMAT: "ag2"},
)
self._chat_state.session_span = span
try:
result = original(*args, **kwargs)
except Exception as e:
result = None
self._record_exception(span, e)
raise e
finally:
# End any pending spans before ending the session
# This ensures they get exported even if an error occurred
for pending in self._chat_state.pending_spans:
pending.span.end(outputs=pending.outputs, end_time_ns=pending.end_time_ns)
span.end(outputs=result)
# Clear the state to start a new chat session
self._chat_state.clear()
elif not root_only:
span = self._start_span_in_session(
name=original.__name__,
span_type=span_type,
inputs=capture_function_input_args(original, args, kwargs),
)
try:
result = original(*args, **kwargs)
except Exception as e:
result = None
self._record_exception(span, e)
raise e
finally:
# Don't end the span yet - defer ending until after parent relocation
# to avoid premature export with incorrect parent_id
end_time_ns = time.time_ns()
self._chat_state.pending_spans.append(_PendingSpan(span, result, end_time_ns))
else:
result = original(*args, **kwargs)
return result
return _wrapper
def _record_exception(self, span: Span, e: Exception):
try:
span.set_status(SpanStatus(SpanStatusCode.ERROR, str(e)))
span.add_event(SpanEvent.from_exception(e))
except Exception as e:
_logger.warning(
"Failed to record exception in span.", exc_info=_logger.isEnabledFor(logging.DEBUG)
)
def _start_span_in_session(
self,
name: str,
span_type: str,
inputs: dict[str, Any],
attributes: dict[str, Any] | None = None,
start_time_ns: int | None = None,
) -> Span:
"""
Start a span in the current chat session.
"""
if self._chat_state.session_span is None:
_logger.warning("Failed to start span. No active chat session.")
return NoOpSpan()
# Add MESSAGE_FORMAT attribute for AG2 spans
attributes = attributes or {}
attributes[SpanAttributeKey.MESSAGE_FORMAT] = "ag2"
return start_span_no_context(
# Tentatively set the parent ID to the session root span, because we
# cannot create a span without a parent span (otherwise it will start
# a new trace). The actual parent will be determined once the chat
# message is received.
parent_span=self._chat_state.session_span,
name=name,
span_type=span_type,
inputs=inputs,
attributes=attributes,
start_time_ns=start_time_ns,
)
@_catch_exception
def log_event(self, source: str | Agent, name: str, **kwargs: dict[str, Any]):
event_end_time = time.time_ns()
if name == "received_message":
if (self._chat_state.last_message is not None) and (
kwargs.get("sender") not in _EXCLUDED_MESSAGE_SENDERS
):
span = self._start_span_in_session(
name=kwargs["sender"],
# Last message is recorded as the input of the next message
inputs=self._chat_state.last_message,
span_type=SpanType.AGENT,
start_time_ns=self._chat_state.last_message_timestamp,
)
# Re-locate the pending spans under this message span BEFORE ending them
# This ensures spans are exported with the correct parent_id
for pending in self._chat_state.pending_spans:
pending.span._span._parent = span._span.context
# Now end the span with its stored outputs and end_time
pending.span.end(outputs=pending.outputs, end_time_ns=pending.end_time_ns)
self._chat_state.pending_spans = []
# End the message span after all children have been relocated and ended
span.end(outputs=kwargs, end_time_ns=event_end_time)
self._chat_state.last_message = kwargs
self._chat_state.last_message_timestamp = event_end_time
@_catch_exception
def log_chat_completion(
self,
invocation_id: uuid.UUID,
client_id: int,
wrapper_id: int,
source: str | Agent,
request: dict[str, float | str | list[dict[str, str]]],
response: str | ChatCompletion,
is_cached: int,
cost: float,
start_time: str,
) -> None:
# The start_time passed from AutoGen is in UTC timezone.
start_dt = datetime.strptime(start_time, "%Y-%m-%d %H:%M:%S.%f")
start_dt = start_dt.replace(tzinfo=timezone.utc)
start_time_ns = int(start_dt.timestamp() * 1e9)
span = self._start_span_in_session(
name="chat_completion",
span_type=SpanType.LLM,
inputs=request,
attributes={
"source": source,
"client_id": client_id,
"invocation_id": invocation_id,
"wrapper_id": wrapper_id,
"cost": cost,
"is_cached": is_cached,
},
start_time_ns=start_time_ns,
)
if model := request.get("model"):
span.set_attribute(SpanAttributeKey.MODEL, model)
if isinstance(model, str):
match model.split("/", 1):
case [provider, _]:
span.set_attribute(SpanAttributeKey.MODEL_PROVIDER, provider)
if usage := self._parse_usage(response):
span.set_attribute(SpanAttributeKey.CHAT_USAGE, usage)
# Defer ending until after parent relocation
# to avoid premature export with incorrect parent_id
end_time_ns = time.time_ns()
self._chat_state.pending_spans.append(_PendingSpan(span, response, end_time_ns))
def _parse_usage(self, output: Any) -> dict[str, int] | None:
usage = getattr(output, "usage", None)
if usage is None:
return None
input_tokens = usage.prompt_tokens
output_tokens = usage.completion_tokens
total_tokens = usage.total_tokens
if total_tokens is None and None not in (input_tokens, output_tokens):
total_tokens = input_tokens + output_tokens
return {
TokenUsageKey.INPUT_TOKENS: input_tokens,
TokenUsageKey.OUTPUT_TOKENS: output_tokens,
TokenUsageKey.TOTAL_TOKENS: total_tokens,
}
# The following methods are not used but are required to implement the BaseLogger interface.
@_catch_exception
def log_function_use(self, *args: Any, **kwargs: Any):
pass
@_catch_exception
def log_new_wrapper(self, wrapper, init_args):
pass
@_catch_exception
def log_new_client(self, client, wrapper, init_args):
pass
@_catch_exception
def stop(self) -> None:
pass
@_catch_exception
def get_connection(self):
pass