This commit is contained in:
Christian Mantha
2026-03-02 19:10:52 -05:00
commit 2ca0b9ef7c
28907 changed files with 5233713 additions and 0 deletions

View File

@@ -0,0 +1,56 @@
from mlflow.utils.annotations import experimental
from mlflow.utils.autologging_utils import autologging_integration
FLAVOR_NAME = "autogen"
@experimental
def autolog(
log_traces: bool = True,
disable: bool = False,
silent: bool = False,
):
"""
Enables (or disables) and configures autologging from Autogen to MLflow. Currently, MLflow
only supports tracing for Autogen agents.
Args:
log_traces: If ``True``, traces are logged for Autogen agents by using runtime logging.
If ``False``, no traces are collected during inference. Default to ``True``.
disable: If ``True``, disables the Autogen autologging. Default to ``False``.
silent: If ``True``, suppress all event logs and warnings from MLflow during Autogen
autologging. If ``False``, show all events and warnings.
"""
from autogen import runtime_logging
from mlflow.autogen.autogen_logger import MlflowAutogenLogger
# NB: The @autologging_integration annotation is used for adding shared logic. However, one
# caveat is that the wrapped function is NOT executed when disable=True is passed. This prevents
# us from running cleaning up logging when autologging is turned off. To workaround this, we
# annotate _autolog() instead of this entrypoint, and define the cleanup logic outside it.
# TODO: since this implementation is inconsistent, explore a universal way to solve the issue.
if log_traces and not disable:
runtime_logging.start(logger=MlflowAutogenLogger())
else:
runtime_logging.stop()
_autolog(log_traces=log_traces, disable=disable, silent=silent)
# This is required by mlflow.autolog()
autolog.integration_name = FLAVOR_NAME
@autologging_integration(FLAVOR_NAME)
def _autolog(
log_traces: bool = True,
disable: bool = False,
silent: bool = False,
):
"""
This is a dummy function only for the purpose of adding the autologging_integration annotation.
We cannot add the annotation directly to the autolog() function above due to the reason
mentioned in the comment above. Note that this function MUST declare the same signature as the
autolog(), otherwise the annotation will not work properly.
"""

View File

@@ -0,0 +1,294 @@
import functools
import logging
import time
import uuid
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any, Optional, Union
from autogen import Agent, ConversableAgent
from autogen.logger.base_logger import BaseLogger
from openai.types.chat import ChatCompletion
from mlflow import MlflowClient
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.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 = "autogen"
@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: Optional[Span] = None
# The last message object in the chat session.
last_message: Optional[Any] = 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.
pending_spans: list[Span] = 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 MlflowAutogenLogger(BaseLogger):
def __init__(self):
self._client = MlflowClient()
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):
original(_self, function_map)
# 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 = self._client.start_trace(
name=original.__name__,
span_type=span_type,
inputs=capture_function_input_args(original, args, kwargs),
)
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:
self._client.end_trace(
request_id=span.request_id, outputs=result, status=span.status
)
# 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:
self._client.end_span(
request_id=span.request_id,
span_id=span.span_id,
outputs=result,
status=span.status,
)
self._chat_state.pending_spans.append(span)
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: Optional[dict[str, Any]] = None,
start_time_ns: Optional[int] = 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()
return self._client.start_span(
request_id=self._chat_state.session_span.request_id,
# 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_id=self._chat_state.session_span.span_id,
name=name,
span_type=span_type,
inputs=inputs,
attributes=attributes,
start_time_ns=start_time_ns,
)
@_catch_exception
def log_event(self, source: Union[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,
)
self._client.end_span(
request_id=span.request_id,
span_id=span.span_id,
outputs=kwargs,
end_time_ns=event_end_time,
)
# Re-locate the pended spans under this message span
for child_span in self._chat_state.pending_spans:
child_span._span._parent = span._span.context
self._chat_state.pending_spans = []
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: Union[str, Agent],
request: dict[str, Union[float, str, list[dict[str, str]]]],
response: Union[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,
)
self._client.end_span(
request_id=span.request_id,
span_id=span.span_id,
outputs=response,
end_time_ns=time.time_ns(),
)
self._chat_state.pending_spans.append(span)
# 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