diff --git a/assert_ai/core/otel.py b/assert_ai/core/otel.py index 8d2f9088..c92952b6 100644 --- a/assert_ai/core/otel.py +++ b/assert_ai/core/otel.py @@ -861,6 +861,314 @@ def _genai_tool_result_str(value: Any) -> str: return json.dumps(parsed) +def _coerce_token_count(value: Any) -> int: + """Coerce token-count attributes to ints, defaulting missing/bad values to 0.""" + if value is None: + return 0 + try: + return int(value) + except (TypeError, ValueError): + return 0 + + +def _truncate_content(value: Any, max_content_chars: int) -> str: + """Stringify and truncate content for judge/tree serialization fields.""" + if value is None: + return "" + if isinstance(value, str): + text = value + elif isinstance(value, (dict, list)): + text = json.dumps(value) + else: + text = str(value) + return text[:max_content_chars] + + +def _messages_text(value: Any, *, roles: set[str | None] | None = None) -> str: + """Extract text from GenAI message objects, optionally filtering by role.""" + messages = _coerce_json(value) + if messages is None: + return "" + if isinstance(messages, str): + return messages + if isinstance(messages, dict): + messages = [messages] + if not isinstance(messages, list): + return "" + + texts: list[str] = [] + for msg in messages: + if isinstance(msg, str): + if roles is None: + texts.append(msg) + continue + if not isinstance(msg, dict): + continue + role = msg.get("role") + if roles is not None and role not in roles: + continue + text = _message_text(msg) + if text: + texts.append(text) + return "\n".join(texts) + + +def _genai_input_text(attrs: dict[str, Any]) -> str: + """Extract user/request text from GenAI input messages.""" + value = attrs.get(_GENAI_INPUT_MESSAGES_KEY) + if value is None: + value = attrs.get(_OPENCLAW_INPUT_MESSAGES_KEY) + text = _messages_text(value, roles={"user"}) + if text: + return text + # Some emitters serialize the prompt as a plain string, or omit roles. + return _messages_text(value, roles=None) + + +def _genai_output_text(span: OTelSpan) -> str: + """Extract assistant/model text from GenAI output messages or span events.""" + attrs = span.attributes + text, _ = _genai_messages_content(attrs.get(_GENAI_OUTPUT_MESSAGES_KEY)) + if not text: + text, _ = _genai_messages_content(attrs.get(_OPENCLAW_OUTPUT_MESSAGES_KEY)) + if not text: + event_text, _, _ = _genai_extract_events(span.events) + text = event_text + return text + + +def _span_input_value(span: OTelSpan) -> Any: + """Convention-aware input/query value for direct extraction helpers.""" + if span.convention == "gen_ai": + if span.kind == "TOOL": + value = span.attributes.get(_GENAI_TOOL_CALL_ARGS_KEY) + if value is None: + value = span.attributes.get(_OPENCLAW_TOOL_INPUT_KEY) + return _truncate_content(_coerce_json(value), 10_000) + return _genai_input_text(span.attributes) + return span.attributes.get(_INPUT_VALUE_KEY, "") + + +def _span_output_value(span: OTelSpan) -> Any: + """Convention-aware output/response value for direct extraction helpers.""" + if span.convention == "gen_ai": + if span.kind == "TOOL": + return _span_tool_result(span) + return _genai_output_text(span) + return span.attributes.get(_OUTPUT_VALUE_KEY, "") + + +def _span_model_name(span: OTelSpan) -> str: + """Convention-aware model name with OpenInference precedence.""" + if span.convention == "gen_ai": + return ( + span.attributes.get(_GENAI_RESPONSE_MODEL_KEY) + or span.attributes.get(_GENAI_REQUEST_MODEL_KEY) + or "" + ) + return span.attributes.get(_LLM_MODEL_KEY, "") + + +def _span_input_tokens(span: OTelSpan) -> int: + if span.convention == "gen_ai": + return _coerce_token_count(span.attributes.get(_GENAI_INPUT_TOKENS_KEY)) + return _coerce_token_count(span.attributes.get(_LLM_INPUT_TOKENS_KEY, 0)) + + +def _span_output_tokens(span: OTelSpan) -> int: + if span.convention == "gen_ai": + return _coerce_token_count(span.attributes.get(_GENAI_OUTPUT_TOKENS_KEY)) + return _coerce_token_count(span.attributes.get(_LLM_OUTPUT_TOKENS_KEY, 0)) + + +def _span_node_name(span: OTelSpan) -> str: + return span.attributes.get(_LANGGRAPH_NODE_KEY, span.name) + + +def _span_tool_name(span: OTelSpan) -> str: + if span.convention == "gen_ai": + return span.attributes.get(_GENAI_TOOL_NAME_KEY) or span.name + return span.attributes.get(_TOOL_NAME_KEY, span.name) + + +def _span_tool_call_id(span: OTelSpan) -> str | None: + if span.convention == "gen_ai": + return span.attributes.get(_GENAI_TOOL_CALL_ID_KEY) + return None + + +def _span_tool_args(span: OTelSpan) -> Any: + if span.convention == "gen_ai": + value = span.attributes.get(_GENAI_TOOL_CALL_ARGS_KEY) + if value is None: + value = span.attributes.get(_OPENCLAW_TOOL_INPUT_KEY) + return _genai_tool_args(value) + + tool_input = span.attributes.get(_INPUT_VALUE_KEY, "") + try: + return json.loads(tool_input) if tool_input else {} + except (json.JSONDecodeError, TypeError): + return {"raw": tool_input} + + +def _span_tool_result(span: OTelSpan) -> str: + if span.convention == "gen_ai": + value = span.attributes.get(_GENAI_TOOL_CALL_RESULT_KEY) + if value is None: + value = span.attributes.get(_OPENCLAW_TOOL_OUTPUT_KEY) + return _genai_tool_result_str(value) + return span.attributes.get(_OUTPUT_VALUE_KEY, "") or "" + + +def _span_requested_tool_calls(span: OTelSpan) -> list[dict[str, Any]]: + """Tool calls requested from GenAI model/agent messages/events.""" + if span.convention != "gen_ai" or span.kind == "TOOL": + return [] + _, attr_tool_calls = _genai_messages_content( + span.attributes.get(_GENAI_OUTPUT_MESSAGES_KEY) + ) + if not attr_tool_calls: + _, attr_tool_calls = _genai_messages_content( + span.attributes.get(_OPENCLAW_OUTPUT_MESSAGES_KEY) + ) + _, event_tool_calls, _ = _genai_extract_events(span.events) + return [*attr_tool_calls, *event_tool_calls] + + +def _span_requested_tool_calls_for_serialization( + span: OTelSpan, + *, + max_content_chars: int, +) -> list[dict[str, Any]]: + """GenAI model/agent-requested tool calls for tree/chunked serialization.""" + if span.convention != "gen_ai" or span.kind == "TOOL": + return [] + + _, attr_tool_calls = _genai_messages_content( + span.attributes.get(_GENAI_OUTPUT_MESSAGES_KEY) + ) + if not attr_tool_calls: + _, attr_tool_calls = _genai_messages_content( + span.attributes.get(_OPENCLAW_OUTPUT_MESSAGES_KEY) + ) + _, event_tool_calls, event_results = _genai_extract_events(span.events) + + results_by_id: dict[str, list[Any]] = {} + for call_id, result in event_results: + results_by_id.setdefault(call_id, []).append(result) + + serialized: list[dict[str, Any]] = [] + for call in [*attr_tool_calls, *event_tool_calls]: + item: dict[str, Any] = { + "name": call.get("name") or "tool", + "arguments": call.get("args", {}), + } + call_id = call.get("call_id") + if call_id: + item["tool_call_id"] = call_id + pending_results = results_by_id.get(call_id, []) + if pending_results: + item["tool_result"] = _truncate_content( + pending_results.pop(0), + max_content_chars, + ) + serialized.append(item) + return serialized + + +def _direct_tool_call_entry(call: dict[str, Any]) -> dict[str, Any]: + return { + "name": call.get("name") or "tool", + "arguments": call.get("args", {}), + } + + +def _remember_requested_call_id( + pending_requested_call_ids: dict[str, int], + call_id: str | None, +) -> None: + if call_id: + pending_requested_call_ids[call_id] = ( + pending_requested_call_ids.get(call_id, 0) + 1 + ) + + +def _consume_requested_call_id( + pending_requested_call_ids: dict[str, int], + call_id: str | None, +) -> bool: + if not call_id: + return False + count = pending_requested_call_ids.get(call_id, 0) + if count <= 0: + return False + if count == 1: + del pending_requested_call_ids[call_id] + else: + pending_requested_call_ids[call_id] = count - 1 + return True + + +def _append_requested_tool_call( + tool_calls: list[dict[str, Any]], + pending_requested_call_ids: dict[str, int], + call: dict[str, Any], +) -> None: + """Append a model/agent-requested tool call and track it for TOOL-span dedupe.""" + tool_calls.append(_direct_tool_call_entry(call)) + _remember_requested_call_id(pending_requested_call_ids, call.get("call_id")) + + +def _append_tool_span_call( + tool_calls: list[dict[str, Any]], + pending_requested_call_ids: dict[str, int], + span: OTelSpan, +) -> None: + """Append a TOOL span unless it is the execution half of a request already listed.""" + if _consume_requested_call_id(pending_requested_call_ids, _span_tool_call_id(span)): + return + tool_calls.append({ + "name": _span_tool_name(span), + "arguments": _span_tool_args(span), + }) + + +def _format_session_tool_call(call: dict[str, Any]) -> str: + args = call.get("args", {}) + if isinstance(args, str): + args_text = args + else: + args_text = json.dumps(args) + return f"{call.get('name') or 'tool'}({args_text})" + + +def _append_requested_session_tool_call( + tool_calls: list[str], + pending_requested_call_ids: dict[str, int], + call: dict[str, Any], +) -> None: + tool_calls.append(_format_session_tool_call(call)) + _remember_requested_call_id(pending_requested_call_ids, call.get("call_id")) + + +def _append_session_tool_span_call( + tool_calls: list[str], + pending_requested_call_ids: dict[str, int], + span: OTelSpan, +) -> None: + if _consume_requested_call_id(pending_requested_call_ids, _span_tool_call_id(span)): + return + tool_calls.append(f"{_span_tool_name(span)}({_span_input_value(span)})") + + +def _truncate_tool_args(value: Any, max_content_chars: int) -> Any: + if isinstance(value, dict) and isinstance(value.get("raw"), str): + truncated = dict(value) + truncated["raw"] = value["raw"][:max_content_chars] + return truncated + return value + + # ── Span validation ─────────────────────────────────────────────── @@ -1277,6 +1585,9 @@ def extract_span_inputs( - response: the output from this span - model: the LLM model used (if available) - tokens: input/output token counts (if available) + + Supports both OpenInference attributes and OTel GenAI (gen_ai.*) fallbacks, + while preserving OpenInference precedence for dual-emitting spans. """ results = [] for span in spans: @@ -1285,13 +1596,13 @@ def extract_span_inputs( results.append({ "span_id": span.span_id, "trace_id": span.trace_id, - "query": span.attributes.get(_INPUT_VALUE_KEY, ""), - "response": span.attributes.get(_OUTPUT_VALUE_KEY, ""), - "model": span.attributes.get(_LLM_MODEL_KEY, ""), - "input_tokens": span.attributes.get(_LLM_INPUT_TOKENS_KEY, 0), - "output_tokens": span.attributes.get(_LLM_OUTPUT_TOKENS_KEY, 0), + "query": _span_input_value(span), + "response": _span_output_value(span), + "model": _span_model_name(span), + "input_tokens": _span_input_tokens(span), + "output_tokens": _span_output_tokens(span), "latency_ms": span.latency_ms, - "node": span.attributes.get(_LANGGRAPH_NODE_KEY, span.name), + "node": _span_node_name(span), }) return results @@ -1329,24 +1640,39 @@ def extract_trajectory_inputs( node_path: list[str] = [] total_input_tokens = 0 total_output_tokens = 0 + pending_requested_call_ids: dict[str, int] = {} for span in group_spans: if span.kind == "LLM": if not user_input: - user_input = span.attributes.get(_INPUT_VALUE_KEY, "") - node = span.attributes.get(_LANGGRAPH_NODE_KEY, span.name) + user_input = _span_input_value(span) + node = _span_node_name(span) if node and node not in node_path: node_path.append(node) - total_input_tokens += span.attributes.get(_LLM_INPUT_TOKENS_KEY, 0) - total_output_tokens += span.attributes.get(_LLM_OUTPUT_TOKENS_KEY, 0) + total_input_tokens += _span_input_tokens(span) + total_output_tokens += _span_output_tokens(span) + for call in _span_requested_tool_calls(span): + _append_requested_tool_call( + tool_calls, + pending_requested_call_ids, + call, + ) elif span.kind == "TOOL": - tool_name = span.attributes.get(_TOOL_NAME_KEY, span.name) - tool_input = span.attributes.get(_INPUT_VALUE_KEY, "") - try: - args = json.loads(tool_input) if tool_input else {} - except (json.JSONDecodeError, TypeError): - args = {"raw": tool_input} - tool_calls.append({"name": tool_name, "arguments": args}) + _append_tool_span_call( + tool_calls, + pending_requested_call_ids, + span, + ) + else: + node = _span_node_name(span) + if node and node not in node_path: + node_path.append(node) + for call in _span_requested_tool_calls(span): + _append_requested_tool_call( + tool_calls, + pending_requested_call_ids, + call, + ) results.append({ "trace_id": group_id, @@ -1387,20 +1713,36 @@ def extract_session_inputs( output_messages: list[str] = [] tool_calls: list[str] = [] trace_ids: set[str] = set() + pending_requested_call_ids: dict[str, int] = {} for span in session_spans: trace_ids.add(span.trace_id) if span.kind == "LLM": - inp = span.attributes.get(_INPUT_VALUE_KEY, "") - out = span.attributes.get(_OUTPUT_VALUE_KEY, "") + inp = _span_input_value(span) + out = _span_output_value(span) if inp and inp not in user_inputs: user_inputs.append(inp) if out: output_messages.append(out) + for call in _span_requested_tool_calls(span): + _append_requested_session_tool_call( + tool_calls, + pending_requested_call_ids, + call, + ) elif span.kind == "TOOL": - tool_name = span.attributes.get(_TOOL_NAME_KEY, span.name) - tool_input = span.attributes.get(_INPUT_VALUE_KEY, "") - tool_calls.append(f"{tool_name}({tool_input})") + _append_session_tool_span_call( + tool_calls, + pending_requested_call_ids, + span, + ) + else: + for call in _span_requested_tool_calls(span): + _append_requested_session_tool_call( + tool_calls, + pending_requested_call_ids, + call, + ) results.append({ "session_id": session_id, @@ -1451,31 +1793,37 @@ def to_dict( d: dict[str, Any] = { "span_id": self.span.span_id, "type": self.span.kind, - "name": self.span.attributes.get(_LANGGRAPH_NODE_KEY, self.span.name), + "name": _span_node_name(self.span), "latency_ms": round(self.span.latency_ms, 1), } if self.span.kind == "LLM": - d["model"] = self.span.attributes.get(_LLM_MODEL_KEY, "") + d["model"] = _span_model_name(self.span) d["tokens"] = { - "input": self.span.attributes.get(_LLM_INPUT_TOKENS_KEY, 0), - "output": self.span.attributes.get(_LLM_OUTPUT_TOKENS_KEY, 0), + "input": _span_input_tokens(self.span), + "output": _span_output_tokens(self.span), } if self.span.kind == "TOOL": - d["tool_name"] = self.span.attributes.get(_TOOL_NAME_KEY, self.span.name) - tool_input = self.span.attributes.get(_INPUT_VALUE_KEY, "") - tool_output = self.span.attributes.get(_OUTPUT_VALUE_KEY, "") - try: - d["tool_args"] = json.loads(tool_input) if tool_input else {} - except (json.JSONDecodeError, TypeError): - d["tool_args"] = {"raw": tool_input[:max_content_chars]} - d["tool_result"] = tool_output[:max_content_chars] if tool_output else "" + d["tool_name"] = _span_tool_name(self.span) + d["tool_args"] = _truncate_tool_args( + _span_tool_args(self.span), + max_content_chars, + ) + result = _span_tool_result(self.span) + d["tool_result"] = _truncate_content(result, max_content_chars) + + requested_tool_calls = _span_requested_tool_calls_for_serialization( + self.span, + max_content_chars=max_content_chars, + ) + if requested_tool_calls: + d["tool_calls"] = requested_tool_calls if include_input: - inp = self.span.attributes.get(_INPUT_VALUE_KEY, "") - d["input"] = inp[:max_content_chars] if inp else "" + inp = _span_input_value(self.span) + d["input"] = _truncate_content(inp, max_content_chars) if include_output: - out = self.span.attributes.get(_OUTPUT_VALUE_KEY, "") - d["output"] = out[:max_content_chars] if out else "" + out = _span_output_value(self.span) + d["output"] = _truncate_content(out, max_content_chars) if self.children: d["children"] = [ @@ -1629,7 +1977,7 @@ def extract_for_judge( chunks: list[dict[str, Any]] = [] for root in tree: chunk = { - "agent": root.span.attributes.get(_LANGGRAPH_NODE_KEY, root.span.name), + "agent": _span_node_name(root.span), "type": root.span.kind, "span_count": root.size, "tree": root.to_dict( diff --git a/tests/test_otel_genai.py b/tests/test_otel_genai.py index fb0526bd..b81aafbf 100644 --- a/tests/test_otel_genai.py +++ b/tests/test_otel_genai.py @@ -899,5 +899,317 @@ def test_standalone_execute_tool_result_does_not_seed_later_duplicate_id(self): self.assertEqual([t["edit"]["tool_result"] for t in tools], ["first-result", ""]) +class TestGenAIDirectExtractionFollowup(unittest.TestCase): + """Issue #241: direct extraction helpers must not leave pure gen_ai spans blank.""" + + def _span(self, *, kind="LLM", span_id="s1", parent=None, name="span", attrs=None, start=0, events=None): + from assert_ai.core.otel import OTelSpan + return OTelSpan( + trace_id="genai_direct", + span_id=span_id, + parent_span_id=parent, + name=name, + kind=kind, + start_time_ns=start, + end_time_ns=start + 1_000_000, + attributes=attrs or {}, + events=events or [], + ) + + def _genai_llm(self, *, span_id="llm1", parent=None, start=0, text="Answer."): + return self._span( + kind="LLM", + span_id=span_id, + parent=parent, + name="chat gpt-4o", + start=start, + attrs={ + "gen_ai.operation.name": "chat", + "gen_ai.request.model": "gpt-4o-mini", + "gen_ai.response.model": "gpt-4o-2024-08-06", + "gen_ai.usage.input_tokens": 12, + "gen_ai.usage.output_tokens": 7, + "gen_ai.input.messages": json.dumps([ + {"role": "user", "content": "What is the weather?"} + ]), + "gen_ai.output.messages": json.dumps([ + {"role": "assistant", "content": text} + ]), + "langgraph.node": "planner", + }, + ) + + def _genai_tool(self, *, span_id="tool1", parent=None, start=0, call_id="call_1"): + return self._span( + kind="TOOL", + span_id=span_id, + parent=parent, + name="execute_tool lookup", + start=start, + attrs={ + "gen_ai.operation.name": "execute_tool", + "gen_ai.tool.name": "lookup", + "gen_ai.tool.call.id": call_id, + "gen_ai.tool.call.arguments": json.dumps({"q": "weather"}), + "gen_ai.tool.call.result": json.dumps({"result": "sunny"}), + }, + ) + + def test_extract_span_inputs_reads_genai_fields(self): + from assert_ai.core.otel import extract_span_inputs + + row = extract_span_inputs([self._genai_llm()])[0] + self.assertEqual(row["query"], "What is the weather?") + self.assertEqual(row["response"], "Answer.") + self.assertEqual(row["model"], "gpt-4o-2024-08-06") + self.assertEqual(row["input_tokens"], 12) + self.assertEqual(row["output_tokens"], 7) + self.assertEqual(row["node"], "planner") + + def test_extract_trajectory_inputs_reads_genai_tools_and_tokens(self): + from assert_ai.core.otel import extract_trajectory_inputs + + llm = self._span( + kind="LLM", + span_id="llm1", + start=0, + attrs={ + "gen_ai.operation.name": "chat", + "gen_ai.request.model": "gpt-4o", + "gen_ai.usage.input_tokens": 20, + "gen_ai.usage.output_tokens": 5, + "gen_ai.input.messages": json.dumps([ + {"role": "user", "content": "Find docs"} + ]), + "gen_ai.output.messages": json.dumps([{ + "role": "assistant", + "parts": [ + {"type": "text", "content": "I'll look."}, + { + "type": "tool_call", + "id": "call_doc", + "name": "search_docs", + "arguments": {"query": "GenAI spans"}, + }, + ], + }]), + "langgraph.node": "researcher", + }, + ) + tool = self._span( + kind="TOOL", + span_id="tool1", + start=1_000, + attrs={ + "gen_ai.operation.name": "execute_tool", + "gen_ai.tool.name": "search_docs", + "gen_ai.tool.call.id": "call_doc", + "gen_ai.tool.call.arguments": json.dumps({"query": "GenAI spans"}), + "gen_ai.tool.call.result": "found", + }, + ) + row = extract_trajectory_inputs([tool, llm])[0] + self.assertEqual(row["user_input"], "Find docs") + self.assertEqual(row["total_tokens"], {"input": 20, "output": 5}) + self.assertEqual(json.loads(row["node_path"]), ["researcher"]) + self.assertEqual( + json.loads(row["tool_calls"]), + [{"name": "search_docs", "arguments": {"query": "GenAI spans"}}], + ) + + def test_extract_session_inputs_reads_genai_messages_and_tool_calls(self): + from assert_ai.core.otel import extract_session_inputs + + row = extract_session_inputs([ + self._genai_llm(span_id="llm1"), + self._genai_tool(span_id="tool1", start=1_000), + ])[0] + self.assertEqual(json.loads(row["user_inputs"]), ["What is the weather?"]) + self.assertEqual(json.loads(row["output_messages"]), ["Answer."]) + self.assertEqual(json.loads(row["tool_calls"]), ['lookup({"q": "weather"})']) + + def test_span_node_to_dict_reads_genai_fields(self): + from assert_ai.core.otel import SpanNode + + llm_dict = SpanNode(self._genai_llm()).to_dict(include_input=True) + self.assertEqual(llm_dict["model"], "gpt-4o-2024-08-06") + self.assertEqual(llm_dict["tokens"], {"input": 12, "output": 7}) + self.assertEqual(llm_dict["input"], "What is the weather?") + self.assertEqual(llm_dict["output"], "Answer.") + + tool_dict = SpanNode(self._genai_tool()).to_dict(include_input=True) + self.assertEqual(tool_dict["tool_name"], "lookup") + self.assertEqual(tool_dict["tool_args"], {"q": "weather"}) + self.assertIn("sunny", tool_dict["tool_result"]) + + def test_tree_mode_auto_selection_preserves_pure_genai_content(self): + from assert_ai.core.otel import extract_for_judge, ExtractionMode + + event_carried_tool_call = { + "name": "gen_ai.choice", + "attributes": [{"key": "message", "value": {"stringValue": json.dumps({ + "role": "assistant", + "content": "Need a lookup.", + "tool_calls": [{ + "id": "event_call", + "type": "function", + "function": { + "name": "search_docs", + "arguments": json.dumps({"query": "GenAI spans"}), + }, + }], + })}}], + } + event_carried_tool_result = { + "name": "gen_ai.tool.message", + "attributes": [ + {"key": "id", "value": {"stringValue": "event_call"}}, + {"key": "content", "value": {"stringValue": "docs found"}}, + ], + } + spans = [self._span( + kind="AGENT", + span_id="root", + name="invoke_agent", + attrs={"gen_ai.operation.name": "invoke_agent", "langgraph.node": "agent"}, + )] + spans.append(self._genai_llm( + span_id="llm0", + parent="root", + start=1_000, + text="Answer 0.", + )) + spans.append(self._span( + kind="LLM", + span_id="llm_event_tool", + parent="root", + name="chat gpt-4o", + start=1_500, + attrs={ + "gen_ai.operation.name": "chat", + "gen_ai.request.model": "gpt-4o", + "langgraph.node": "tool_requester", + }, + events=[event_carried_tool_call, event_carried_tool_result], + )) + spans.append(self._genai_tool(span_id="tool0", parent="root", start=2_000)) + for i in range(7): + spans.append(self._genai_llm( + span_id=f"llm_extra_{i}", + parent="root", + start=3_000 + i, + text=f"Extra answer {i}.", + )) + + result = extract_for_judge(spans) + self.assertEqual(result["mode"], ExtractionMode.TREE) + root = result["representation"][0] + llm_child = next(c for c in root["children"] if c["span_id"] == "llm0") + tool_requester = next(c for c in root["children"] if c["span_id"] == "llm_event_tool") + tool_child = next(c for c in root["children"] if c["span_id"] == "tool0") + self.assertEqual(llm_child["model"], "gpt-4o-2024-08-06") + self.assertEqual(llm_child["output"], "Answer 0.") + self.assertEqual(tool_requester["tool_calls"], [{ + "name": "search_docs", + "arguments": {"query": "GenAI spans"}, + "tool_call_id": "event_call", + "tool_result": "docs found", + }]) + self.assertEqual(tool_child["tool_name"], "lookup") + self.assertEqual(tool_child["tool_args"], {"q": "weather"}) + self.assertIn("sunny", tool_child["tool_result"]) + + def test_repeated_call_ids_are_not_collapsed_in_direct_helpers(self): + from assert_ai.core.otel import extract_trajectory_inputs, extract_session_inputs + + def choice_span(span_id, q, start): + return self._span( + kind="LLM", + span_id=span_id, + start=start, + attrs={ + "gen_ai.operation.name": "chat", + "gen_ai.request.model": "gpt-4o", + "gen_ai.input.messages": json.dumps([ + {"role": "user", "content": f"lookup {q}"} + ]), + "gen_ai.output.messages": json.dumps([{ + "role": "assistant", + "tool_calls": [{ + "id": "dup", + "type": "function", + "function": { + "name": "lookup", + "arguments": json.dumps({"q": q}), + }, + }], + }]), + "session.id": "sess_dup", + }, + ) + + spans = [choice_span("llm1", "first", 0), choice_span("llm2", "second", 1_000)] + traj_calls = json.loads(extract_trajectory_inputs(spans)[0]["tool_calls"]) + self.assertEqual( + traj_calls, + [ + {"name": "lookup", "arguments": {"q": "first"}}, + {"name": "lookup", "arguments": {"q": "second"}}, + ], + ) + session_calls = json.loads(extract_session_inputs(spans)[0]["tool_calls"]) + self.assertEqual(session_calls, ['lookup({"q": "first"})', 'lookup({"q": "second"})']) + + def test_tool_arg_raw_value_still_truncates_in_tree_serialization(self): + from assert_ai.core.otel import OTelSpan, SpanNode + + span = OTelSpan( + trace_id="t1", + span_id="tool_raw", + parent_span_id=None, + name="tool", + kind="TOOL", + start_time_ns=0, + end_time_ns=1_000_000, + attributes={ + "openinference.span.kind": "TOOL", + "tool.name": "bad_args_tool", + "input.value": "x" * 2000, + "output.value": "ok", + }, + ) + d = SpanNode(span).to_dict(max_content_chars=50) + self.assertEqual(d["tool_args"], {"raw": "x" * 50}) + + def test_dual_emitting_spans_keep_openinference_precedence(self): + from assert_ai.core.otel import extract_span_inputs, SpanNode + + span = self._span( + kind="LLM", + attrs={ + "openinference.span.kind": "LLM", + "input.value": "openinference input", + "output.value": "openinference output", + "llm.model_name": "oi-model", + "llm.token_count.prompt": 3, + "llm.token_count.completion": 4, + "gen_ai.operation.name": "chat", + "gen_ai.request.model": "genai-model", + "gen_ai.input.messages": json.dumps([ + {"role": "user", "content": "genai input"} + ]), + "gen_ai.output.messages": json.dumps([ + {"role": "assistant", "content": "genai output"} + ]), + }, + ) + + row = extract_span_inputs([span])[0] + self.assertEqual(row["query"], "openinference input") + self.assertEqual(row["response"], "openinference output") + self.assertEqual(row["model"], "oi-model") + self.assertEqual(SpanNode(span).to_dict()["output"], "openinference output") + + if __name__ == "__main__": unittest.main()