Initialize trace adapter (#103)

This commit is contained in:
Yuge Zhang
2025-09-20 05:57:12 -07:00
committed by GitHub
parent 3eb725fade
commit a63197355c
11 changed files with 348 additions and 29 deletions
+6
View File
@@ -0,0 +1,6 @@
# Copyright (c) Microsoft. All rights reserved.
from .base import Adapter, TraceAdapter
from .triplet import TraceTripletAdapter
__all__ = ["TraceAdapter", "Adapter", "TraceTripletAdapter"]
+85
View File
@@ -0,0 +1,85 @@
# Copyright (c) Microsoft. All rights reserved.
from typing import Generic, List, TypeVar
from opentelemetry.sdk.trace import ReadableSpan
T_from = TypeVar("T_from")
T_to = TypeVar("T_to")
class Adapter(Generic[T_from, T_to]):
"""Base class for synchronous adapters that convert data from one format to another.
This class defines a simple protocol for transformation:
- The `__call__` method makes adapters callable, so they can be used like functions.
- Subclasses must implement the `adapt` method to define the actual conversion logic.
Type parameters:
T_from: The source data type (input).
T_to: The target data type (output).
Example:
>>> class IntToStrAdapter(Adapter[int, str]):
... def adapt(self, source: int) -> str:
... return str(source)
...
>>> adapter = IntToStrAdapter()
>>> adapter(42)
'42'
"""
def __call__(self, source: T_from, /) -> T_to:
"""Convert the data to the target format.
This method delegates to `adapt` and allows the adapter
to be invoked as a function.
Args:
source: Input data in the source format.
Returns:
Data converted to the target format.
"""
return self.adapt(source)
def adapt(self, source: T_from, /) -> T_to:
"""Convert the data to the target format.
Subclasses should override this method with the concrete
transformation logic.
Args:
source: Input data in the source format.
Returns:
Data converted to the target format.
Raises:
NotImplementedError: If the method is not implemented
in a subclass.
"""
raise NotImplementedError("Adapter.adapt() is not implemented")
class TraceAdapter(Adapter[List[ReadableSpan], T_to], Generic[T_to]):
"""Base class for adapters that convert OpenTelemetry trace spans into other formats.
This class specializes `Adapter` for working with OpenTelemetry `ReadableSpan`
objects. It expects a list of spans as input and produces a custom target format
(e.g., reinforcement learning training data, SFT datasets, logs, metrics).
Subclasses should override `adapt` to define the desired conversion.
Type parameters:
T_to: The target data type that spans should be converted into.
Example:
>>> class TraceToDictAdapter(TraceAdapter[dict]):
... def adapt(self, spans: List[ReadableSpan]) -> dict:
... return {"count": len(spans)}
...
>>> adapter = TraceToDictAdapter()
>>> adapter([span1, span2])
{'count': 2}
"""
+218
View File
@@ -0,0 +1,218 @@
# Copyright (c) Microsoft. All rights reserved.
import json
from collections import defaultdict
from typing import Any, Dict, Generator, List, Optional, Sequence, TypedDict, Union
from openai.types.chat.chat_completion_function_tool_param import ChatCompletionFunctionToolParam
from openai.types.chat.chat_completion_message import ChatCompletionMessage
from openai.types.chat.chat_completion_message_function_tool_call import ChatCompletionMessageFunctionToolCall, Function
from openai.types.shared_params import FunctionDefinition
from opentelemetry.sdk.trace import ReadableSpan
from pydantic import BaseModel
from .base import TraceAdapter
class OpenAIMessages(BaseModel):
messages: List[ChatCompletionMessage]
tools: Optional[List[ChatCompletionFunctionToolParam]] = None
class _RawSpanInfo(TypedDict):
prompt: List[Dict[str, Any]]
completion: List[Dict[str, Any]]
request: Dict[str, Any]
response: Dict[str, Any]
def group_genai_dict(data: Dict[str, Any], prefix: str) -> Union[Dict[str, Any], List[Any]]:
"""
Convert a flat dict with keys like 'gen_ai.prompt.0.role'
into structured nested dicts or lists under the given prefix.
Args:
data: Flat dictionary (keys are dotted paths).
prefix: Top-level key to extract (e.g., 'gen_ai.prompt').
Returns:
A nested dict (if no index detected) or list (if indexed).
"""
result: Union[Dict[str, Any], List[Any]] = {}
# Collect keys that match the prefix
relevant = {k[len(prefix) + 1 :]: v for k, v in data.items() if k.startswith(prefix + ".")}
# Detect if we have numeric indices (-> list) or not (-> dict)
indexed = any(part.split(".")[0].isdigit() for part in relevant.keys())
if indexed:
# Group by index
grouped: Dict[int, Dict[str, Any]] = defaultdict(dict)
for k, v in relevant.items():
parts = k.split(".")
if not parts[0].isdigit():
continue
idx, rest = int(parts[0]), ".".join(parts[1:])
grouped[idx][rest] = v
# Recursively build
result = []
for i in sorted(grouped.keys()):
result.append(group_genai_dict({f"{prefix}.{rest}": val for rest, val in grouped[i].items()}, prefix))
else:
# No indices: build dict
nested: Dict[str, Any] = defaultdict(dict)
for k, v in relevant.items():
if "." in k:
head, _tail = k.split(".", 1)
nested[head][f"{prefix}.{k}"] = v
else:
result[k] = v
# Recurse into nested dicts
for head, subdict in nested.items():
result[head] = group_genai_dict(subdict, prefix + "." + head)
return result
def convert_to_openai_messages(
prompt_completion_list: List[_RawSpanInfo], tool_requests: List[Dict[str, Any]]
) -> Generator[OpenAIMessages, None, None]:
"""
Convert raw tool call traces + prompt/completion list
into OpenAI fine-tuning JSONL format (tool calling style).
Since promopt-completions sometimes do not contain the generated tool calls,
the tool call requests need to be provided separately.
The tool calls are then matched in a first-come-first-served basis to the tool call requests.
https://learn.microsoft.com/en-us/azure/ai-foundry/openai/how-to/fine-tuning-functions
"""
for pc_entry in prompt_completion_list:
messages: List[ChatCompletionMessage] = []
tools: List[ChatCompletionFunctionToolParam] = []
# Extract messages
for msg in pc_entry["prompt"]:
role = msg["role"]
if role == "assistant" and "tool_calls" in msg:
# Use the tool_calls directly
tool_calls: Sequence[ChatCompletionMessageFunctionToolCall] = []
for call in msg["tool_calls"]:
function = Function(name=call["name"], arguments=call["arguments"])
tool_calls.append(
ChatCompletionMessageFunctionToolCall(
id=call["id"],
type="function",
function=function,
)
)
messages.append(ChatCompletionMessage(role="assistant", tool_calls=list(tool_calls)))
else:
# Normal user/system/tool content
message = ChatCompletionMessage(
role=role, content=msg.get("content", ""), tool_calls=msg.get("tool_calls", None)
)
messages.append(message)
# Extract completions (assistant outputs after tool responses)
for comp in pc_entry["completion"]:
if comp.get("role") == "assistant":
if comp.get("content"):
message = ChatCompletionMessage(role="assistant", content=comp["content"])
messages.append(message)
elif comp.get("finish_reason") == "tool_calls":
if len(tool_requests) == 0:
raise ValueError("No tool requests available for tool_calls completion")
tool_req = tool_requests.pop(0)
# FIXME: this is a hack because tracing frameworks did not report the tool call properly
message = ChatCompletionMessage(
role="assistant",
tool_calls=[
ChatCompletionMessageFunctionToolCall(
id=tool_req["call"]["id"],
type=tool_req["call"]["type"],
function=Function(name=tool_req["name"], arguments=tool_req["parameters"]),
)
],
)
messages.append(message)
else:
raise ValueError(f"Unsupported assistant completion: {comp}")
# Build tools definitions (if available)
if "functions" in pc_entry["request"]:
for fn in pc_entry["request"]["functions"]:
tools.append(
ChatCompletionFunctionToolParam(
type="function",
function=FunctionDefinition(
name=fn["name"],
description=fn.get("description", ""),
parameters=(
json.loads(fn["parameters"]) if isinstance(fn["parameters"], str) else fn["parameters"]
),
),
)
)
yield OpenAIMessages(messages=messages, tools=tools)
else:
yield OpenAIMessages(messages=messages, tools=tools)
class TraceMessagesAdapter(TraceAdapter[List[OpenAIMessages]]):
"""
Adapter that converts OpenTelemetry trace spans into OpenAI-compatible message format.
This adapter processes trace spans containing LLM conversation data and transforms them
into structured OpenAI message format suitable for fine-tuning or analysis. It extracts
prompts, completions, tool calls, and function definitions from trace attributes and
reconstructs the conversation flow.
The adapter handles:
- Converting flat trace attributes into structured message objects
- Extracting and matching tool calls with their corresponding requests
- Building proper OpenAI ChatCompletionMessage objects with roles, content, and tool calls
- Generating function definitions for tools used in conversations
Returns:
List[OpenAIMessages]: A list of structured message conversations with associated tools
"""
def adapt(self, source: List[ReadableSpan], /) -> List[OpenAIMessages]:
raw_tool_calls: List[Dict[str, Any]] = []
raw_prompt_completions: List[_RawSpanInfo] = []
for span in source:
if span.attributes is None:
continue
attributes = {k: v for k, v in span.attributes.items()}
# Otherwise we strip all the tool calls and prompts and responses
tool_call = group_genai_dict(dict(attributes), "tool")
if not isinstance(tool_call, dict):
raise ValueError(f"Extracted tool call from trace is not a dict: {tool_call}")
if tool_call:
raw_tool_calls.append(tool_call)
# Get all related information from the trace span
prompt = group_genai_dict(attributes, "gen_ai.prompt")
completion = group_genai_dict(attributes, "gen_ai.completion")
request = group_genai_dict(attributes, "gen_ai.request")
response = group_genai_dict(attributes, "gen_ai.response")
if not isinstance(prompt, list):
raise ValueError(f"Extracted prompt from trace is not a list: {prompt}")
if not isinstance(completion, list):
raise ValueError(f"Extracted completion from trace is not a list: {completion}")
if not isinstance(request, dict):
raise ValueError(f"Extracted request from trace is not a dict: {request}")
if not isinstance(response, dict):
raise ValueError(f"Extracted response from trace is not a dict: {response}")
if prompt or completion or request or response:
raw_prompt_completions.append(
_RawSpanInfo(prompt=prompt, completion=completion, request=request, response=response)
)
return list(convert_to_openai_messages(raw_prompt_completions, raw_tool_calls))
@@ -11,6 +11,8 @@ from pydantic import BaseModel
from agentlightning.types import Triplet
from .base import TraceAdapter
class Transition(BaseModel):
"""
@@ -510,9 +512,9 @@ class TraceTree:
)
class TripletExporter:
class TraceTripletAdapter(TraceAdapter[List[Triplet]]):
"""
A class to export triplet data from OpenTelemetry spans.
An adapter to convert OpenTelemetry spans to triplet data.
Attributes:
repair_hierarchy: When `repair_hierarchy` is set to True, the trace will be repaired with the time information.
@@ -537,9 +539,9 @@ class TripletExporter:
self.exclude_llm_call_in_reward = exclude_llm_call_in_reward
self.reward_match = reward_match
def export(self, spans: List[ReadableSpan]) -> List[Triplet]:
def adapt(self, source: List[ReadableSpan], /) -> List[Triplet]:
"""Convert OpenTelemetry spans to a list of Triplet objects."""
trace_tree = TraceTree.from_spans(spans)
trace_tree = TraceTree.from_spans(source)
if self.repair_hierarchy:
trace_tree.repair_hierarchy()
trajectory = trace_tree.to_trajectory(
+4 -4
View File
@@ -7,9 +7,9 @@ from typing import Any, Dict, List, Optional, cast
from opentelemetry.sdk.trace import ReadableSpan
from .adapter import TraceTripletAdapter
from .client import AgentLightningClient
from .litagent import LitAgent, is_v0_1_rollout_api
from .tracer import TripletExporter
from .tracer.base import BaseTracer
from .types import ParallelWorkerBase, Rollout, RolloutRawResult, Triplet
@@ -37,7 +37,7 @@ class AgentRunner(ParallelWorkerBase):
agent: LitAgent[Any],
client: AgentLightningClient,
tracer: BaseTracer,
triplet_exporter: TripletExporter,
triplet_exporter: TraceTripletAdapter,
worker_id: Optional[int] = None,
max_tasks: Optional[int] = None,
):
@@ -108,9 +108,9 @@ class AgentRunner(ParallelWorkerBase):
trace = [json.loads(readable_span.to_json()) for readable_span in spans]
trace_spans = spans
# Always extract triplets from the trace using TripletExporter
# Always extract triplets from the trace using TraceTripletAdapter
if trace_spans:
triplets = self.triplet_exporter.export(trace_spans)
triplets = self.triplet_exporter(trace_spans)
# If the agent has triplets, use the last one for final reward if not set
if triplets and triplets[-1].reward is not None and final_reward is None:
+1 -2
View File
@@ -2,6 +2,5 @@
from .agentops import AgentOpsTracer
from .base import BaseTracer
from .triplet import TripletExporter
__all__ = ["AgentOpsTracer", "BaseTracer", "TripletExporter"]
__all__ = ["AgentOpsTracer", "BaseTracer"]
+1 -1
View File
@@ -38,7 +38,7 @@ class BaseTracer(ParallelWorkerBase):
# Process the trace data
if trace_tree:
rl_triplets = TripletExporter().export(spans)
rl_triplets = TraceTripletAdapter().adapt(spans)
# ... do something with the triplets
```
"""
+7 -7
View File
@@ -9,13 +9,13 @@ import time
import warnings
from typing import Any, Dict, List, Optional, TypeVar, Union
from .adapter import TraceTripletAdapter
from .algorithm.base import BaseAlgorithm
from .client import AgentLightningClient
from .litagent import LitAgent
from .runner import AgentRunner
from .tracer.agentops import AgentOpsTracer
from .tracer.base import BaseTracer
from .tracer.triplet import TripletExporter
from .types import Dataset, ParallelWorkerBase
logger = logging.getLogger(__name__)
@@ -41,7 +41,7 @@ class Trainer(ParallelWorkerBase):
tracer: A tracer instance, or a string pointing to the class full name or a dictionary with a 'type' key
that specifies the class full name and other initialization parameters.
If None, a default `AgentOpsTracer` will be created with the current settings.
triplet_exporter: An instance of `TripletExporter` to export triplets from traces,
triplet_exporter: An instance of `TraceTripletAdapter` to export triplets from traces,
or a dictionary with the initialization parameters for the exporter.
algorithm: An instance of `BaseAlgorithm` to use for training.
"""
@@ -54,7 +54,7 @@ class Trainer(ParallelWorkerBase):
max_tasks: Optional[int] = None,
daemon: bool = True,
tracer: Union[BaseTracer, str, Dict[str, Any], None] = None,
triplet_exporter: Union[TripletExporter, Dict[str, Any], None] = None,
triplet_exporter: Union[TraceTripletAdapter, Dict[str, Any], None] = None,
algorithm: Union[BaseAlgorithm, str, Dict[str, Any], None] = None,
):
super().__init__()
@@ -65,15 +65,15 @@ class Trainer(ParallelWorkerBase):
self._client: AgentLightningClient | None = None # Will be initialized in fit method
self.tracer = self._make_tracer(tracer)
if isinstance(triplet_exporter, TripletExporter):
if isinstance(triplet_exporter, TraceTripletAdapter):
self.triplet_exporter = triplet_exporter
elif isinstance(triplet_exporter, dict):
self.triplet_exporter = TripletExporter(**triplet_exporter)
self.triplet_exporter = TraceTripletAdapter(**triplet_exporter)
elif triplet_exporter is None:
self.triplet_exporter = TripletExporter()
self.triplet_exporter = TraceTripletAdapter()
else:
raise ValueError(
f"Invalid triplet_exporter type: {type(triplet_exporter)}. Expected TripletExporter, dict, or None."
f"Invalid triplet_exporter type: {type(triplet_exporter)}. Expected TraceTripletAdapter, dict, or None."
)
self.algorithm = self._make_algorithm(algorithm)
+4
View File
@@ -14,6 +14,10 @@ dependencies = [
"uvicorn",
"fastapi",
"aiohttp",
"opentelemetry-api>=1.35",
"opentelemetry-sdk>=1.35",
"pydantic>=2.11",
"openai",
]
[project.optional-dependencies]
+12 -7
View File
@@ -1,13 +1,16 @@
# Copyright (c) Microsoft. All rights reserved.
from contextlib import contextmanager
from typing import Any, Iterator, List, Optional
from typing import Any, Iterator, List, Optional, cast
import pytest
from agentlightning import LitAgent, ResourcesUpdate, Task
from agentlightning.adapter import TraceTripletAdapter
from agentlightning.client import AgentLightningClient
from agentlightning.runner import AgentRunner
from agentlightning.tracer import BaseTracer, TripletExporter
from agentlightning.tracer import BaseTracer
from agentlightning.types import Rollout
class DummyTracer(BaseTracer):
@@ -60,7 +63,7 @@ class HookAgent(LitAgent[Any]):
super().__init__()
self.start_called = False
self.end_called = False
self.end_rollout = None
self.end_rollout: Rollout | None = None
def training_rollout(self, task: Any, resources: Any, rollout: Any) -> float:
return 0.5
@@ -81,12 +84,13 @@ def test_runner_calls_hooks():
agent = HookAgent()
client = DummyClient()
tracer = DummyTracer()
runner = AgentRunner(agent, client, tracer, TripletExporter()) # type: ignore
runner = AgentRunner(agent, cast(AgentLightningClient, client), tracer, TraceTripletAdapter())
assert runner.run() is True
assert agent.start_called
assert agent.end_called
assert agent.end_rollout.final_reward == 0.5 # type: ignore
assert agent.end_rollout is not None
assert agent.end_rollout.final_reward == 0.5
@pytest.mark.asyncio
@@ -94,9 +98,10 @@ async def test_runner_calls_hooks_async():
agent = HookAgent()
client = DummyAsyncClient()
tracer = DummyTracer()
runner = AgentRunner(agent, client, tracer, TripletExporter()) # type: ignore
runner = AgentRunner(agent, cast(AgentLightningClient, client), tracer, TraceTripletAdapter())
assert await runner.run_async() is True
assert agent.start_called
assert agent.end_called
assert agent.end_rollout.final_reward == 0.5 # type: ignore
assert agent.end_rollout is not None
assert agent.end_rollout.final_reward == 0.5
+4 -4
View File
@@ -62,10 +62,10 @@ from openai import AsyncOpenAI, OpenAI
from pydantic import BaseModel, Field
from typing_extensions import TypedDict
from agentlightning.adapter.triplet import TraceTree, TraceTripletAdapter
from agentlightning.reward import reward
from agentlightning.tracer.agentops import AgentOpsTracer, LightningSpanProcessor
from agentlightning.tracer.http import HttpTracer
from agentlightning.tracer.triplet import TraceTree, TripletExporter
from agentlightning.types import Triplet
USE_OPENAI = os.environ.get("USE_OPENAI", "false").lower() == "true"
@@ -695,9 +695,9 @@ def run_with_agentops_tracer() -> None:
assert_expected_pairs_in_tree(tree.names_tuple(), AGENTOPS_EXPECTED_TREES[agent_func.__name__])
# for triplet in TripletExporter().export(tracer.get_last_trace()):
# for triplet in TripleTraceTripletAdaptertExporter().adapt(tracer.get_last_trace()):
# print(triplet)
triplets = TripletExporter().export(tracer.get_last_trace())
triplets = TraceTripletAdapter().adapt(tracer.get_last_trace())
assert (
len(triplets) == AGENTOPS_EXPECTED_TRIPLETS_NUMBER[agent_func.__name__]
), f"Expected {AGENTOPS_EXPECTED_TRIPLETS_NUMBER[agent_func.__name__]} triplets, but got: {triplets}"
@@ -803,7 +803,7 @@ def test_run_with_agentops_tracer(agent_func):
assert_expected_pairs_in_tree(tree.names_tuple(), AGENTOPS_EXPECTED_TREES[agent_func.__name__])
triplets = TripletExporter().export(tracer.get_last_trace())
triplets = TraceTripletAdapter().adapt(tracer.get_last_trace())
assert (
len(triplets) == AGENTOPS_EXPECTED_TRIPLETS_NUMBER[agent_func.__name__]
), f"Expected {AGENTOPS_EXPECTED_TRIPLETS_NUMBER[agent_func.__name__]} triplets, but got: {triplets}"