Initialize trace adapter (#103)
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from .base import Adapter, TraceAdapter
|
||||
from .triplet import TraceTripletAdapter
|
||||
|
||||
__all__ = ["TraceAdapter", "Adapter", "TraceTripletAdapter"]
|
||||
@@ -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}
|
||||
"""
|
||||
@@ -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(
|
||||
@@ -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:
|
||||
|
||||
@@ -2,6 +2,5 @@
|
||||
|
||||
from .agentops import AgentOpsTracer
|
||||
from .base import BaseTracer
|
||||
from .triplet import TripletExporter
|
||||
|
||||
__all__ = ["AgentOpsTracer", "BaseTracer", "TripletExporter"]
|
||||
__all__ = ["AgentOpsTracer", "BaseTracer"]
|
||||
|
||||
@@ -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
|
||||
```
|
||||
"""
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -14,6 +14,10 @@ dependencies = [
|
||||
"uvicorn",
|
||||
"fastapi",
|
||||
"aiohttp",
|
||||
"opentelemetry-api>=1.35",
|
||||
"opentelemetry-sdk>=1.35",
|
||||
"pydantic>=2.11",
|
||||
"openai",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}"
|
||||
|
||||
Reference in New Issue
Block a user