From 2cba7089fbe4c2138c4cf134f4dd94faaebd189f Mon Sep 17 00:00:00 2001 From: Jeff Dupont Date: Wed, 23 Sep 2026 14:15:13 -0700 Subject: [PATCH 1/3] [AIC-3210] feat(AIC-3210): add graph().stream() with OTel parity and graph-streaming example Co-authored-by: Cursor --- examples/graph_streaming.py | 63 ++ main.py | 2 + .../src/launchdarkly_ai_server/__init__.py | 2 + .../launchdarkly_ai_server/conversation.py | 28 + .../src/launchdarkly_ai_server/graph.py | 711 +++++++++++++++--- .../src/launchdarkly_ai_server/tracking.py | 23 +- .../src/launchdarkly_ai_server/types.py | 10 + packages/client/tests/test_graph_stream.py | 700 +++++++++++++++++ packages/client/tests/test_tracking.py | 7 +- 9 files changed, 1449 insertions(+), 97 deletions(-) create mode 100644 examples/graph_streaming.py create mode 100644 packages/client/tests/test_graph_stream.py diff --git a/examples/graph_streaming.py b/examples/graph_streaming.py new file mode 100644 index 00000000..af175507 --- /dev/null +++ b/examples/graph_streaming.py @@ -0,0 +1,63 @@ +""" +Example: ``graph().stream()`` — the streaming counterpart to ``examples/graph_example.py``. + +Model text goes to stdout; node boundaries go to stderr so stdout stays a clean transcript. + +Like ``examples/streaming.py``, the generator is built inside ``conversation_id`` and iterated +*outside* it. That is the shape a server produces when it hands a stream to a transport, and it +is what exercises call-time binding: an async generator body does not run until the first +``__anext__``, so both the conversation id and the ``ld.ai.graph`` span's OTel parent have to be +captured when ``stream()`` is called, not when iteration starts. + +Writes no JSON output file, same as the single-config streaming example. + +Usage (via main.py): + python main.py graph-streaming "" +""" + +from __future__ import annotations + +import json +import sys + +import examples.register # noqa: F401 – side-effect: populate global_registry +from examples.utils import new_context, new_conversation_id +from launchdarkly_ai_server import conversation_id, global_registry, graph + + +async def run(key: str, user_input: str) -> None: + conversation = new_conversation_id("graph-streaming-example") + print(f"[conversation] {conversation}", file=sys.stderr) + + with conversation_id(conversation): + stream = graph( + key, + registry=global_registry, + ).stream(user_input, new_context(), {"user_id": "user-123"}) + + async for event in stream: + if event["type"] == "chunk": + sys.stdout.write(event.get("text", "")) + sys.stdout.flush() + elif event["type"] == "node_start": + print(f"\n[node_start] {event['nodeKey']}", file=sys.stderr) + elif event["type"] == "node_done": + print( + f"\n[node_done] {event['nodeKey']} usage={json.dumps(event.get('usage'))}", + file=sys.stderr, + ) + elif event["type"] == "handoff": + print( + f"[handoff] {event['sourceKey']} -> {event['targetKey']}", + file=sys.stderr, + ) + else: + # Final event — usage aggregated across nodes, plus graph judge results when configured. + sys.stdout.write("\n\n") + print("Usage:", json.dumps(event.get("usage"), indent=2, default=str)) + if event.get("judgeResults"): + print( + "Judge results:", + json.dumps(event["judgeResults"], indent=2, default=str), + ) + sys.stdout.write("\n") diff --git a/main.py b/main.py index fc5c1e5b..0bbaefe1 100644 --- a/main.py +++ b/main.py @@ -10,6 +10,7 @@ python main.py judge launch-darkly-documentation-summarizer "What is the LaunchDarkly AI SDK?" python main.py graph my-agent-graph "What is the LaunchDarkly AI SDK?" python main.py graph-history my-agent-graph "" + python main.py graph-streaming travel-agent-flow "I was double charged for my flight" python main.py openai-only my-openai-flag "Tell me about feature flags" python main.py langchain my-langchain-flag "Tell me about feature flags" python main.py claude-agents launch-darkly-documentation-summarizer "What is the LaunchDarkly AI SDK?" @@ -46,6 +47,7 @@ "streaming": "examples.streaming", "graph": "examples.graph_example", "graph-history": "examples.graph_history", + "graph-streaming": "examples.graph_streaming", "conversation": "examples.conversation", "history": "examples.history", "judge": "examples.judge_example", diff --git a/packages/client/src/launchdarkly_ai_server/__init__.py b/packages/client/src/launchdarkly_ai_server/__init__.py index e05df909..5381ba9f 100644 --- a/packages/client/src/launchdarkly_ai_server/__init__.py +++ b/packages/client/src/launchdarkly_ai_server/__init__.py @@ -73,6 +73,7 @@ GraphEdge, GraphNode, GraphOptions, + GraphStreamEvent, GraphTopology, HandlerResult, HandlerStreamEvent, @@ -136,6 +137,7 @@ "GraphEdge", "GraphNode", "GraphOptions", + "GraphStreamEvent", "GraphTopology", "HandlerResult", "HandlerStreamEvent", diff --git a/packages/client/src/launchdarkly_ai_server/conversation.py b/packages/client/src/launchdarkly_ai_server/conversation.py index 56b5eb05..d3af0e9e 100644 --- a/packages/client/src/launchdarkly_ai_server/conversation.py +++ b/packages/client/src/launchdarkly_ai_server/conversation.py @@ -200,6 +200,34 @@ def bind_conversation_id( return _stream_with_bound_id(generator, conversation) +async def bind_span_context( + generator: AsyncGenerator[Any, None], + ctx: otel_context.Context, +) -> AsyncGenerator[Any, None]: + """Re-enter ``ctx`` around every step of ``generator``. + + Sibling of :func:`bind_conversation_id`, which deliberately carries only the conversation id + and leaves span parenting alone. A generator body suspends at each ``yield``, so the context + has to be re-applied on every ``__anext__`` — wrapping the body once is not enough. + """ + try: + while True: + token = otel_context.attach(ctx) + try: + item = await generator.__anext__() + except StopAsyncIteration: + return + finally: + otel_context.detach(token) + yield item + finally: + token = otel_context.attach(ctx) + try: + await generator.aclose() + finally: + otel_context.detach(token) + + @asynccontextmanager async def with_judge_evaluation(name: str) -> AsyncIterator[RecordEvaluation]: """Hold the judge ``invoke_agent`` span open until ``record`` runs. diff --git a/packages/client/src/launchdarkly_ai_server/graph.py b/packages/client/src/launchdarkly_ai_server/graph.py index 4546c3b8..317eff73 100644 --- a/packages/client/src/launchdarkly_ai_server/graph.py +++ b/packages/client/src/launchdarkly_ai_server/graph.py @@ -6,15 +6,17 @@ import re as _re import time import uuid -from collections.abc import Callable +from collections.abc import AsyncGenerator, Callable from typing import Any +from .conversation import bind_conversation_id, bind_span_context from .registry import resolve_handlers, resolve_tools from .types import ( AiConfigRep, GraphDefinition, GraphEdge, GraphNode, + GraphStreamEvent, JudgeResult, LDContext, NativeTool, @@ -24,7 +26,7 @@ UsageDict, VariationMeta, ) -from .utils import model_stamps_from_meta, select_handler, to_ld_context +from .utils import end_span_once, model_stamps_from_meta, select_handler, to_ld_context logger = logging.getLogger(__name__) @@ -33,7 +35,8 @@ def _sanitize_name(key: str) -> str: - return _re.sub(r"[^a-z0-9_-]", "_", key, flags=_re.IGNORECASE)[:64] + # Match TS sanitizeName: hyphens become underscores (tool names must be [a-zA-Z0-9_]). + return _re.sub(r"[^a-zA-Z0-9_]", "_", key)[:64] def _disabled_definition(key: str) -> GraphDefinition: @@ -46,6 +49,12 @@ async def _run_node_disabled(*args: Any, **kwargs: Any) -> Any: async def _route_disabled(*args: Any, **kwargs: Any) -> Any: raise ValueError(f'Agent graph "{key}" is disabled') + async def _stream_route_disabled( + *args: Any, **kwargs: Any + ) -> AsyncGenerator[GraphStreamEvent, None]: + raise ValueError(f'Agent graph "{key}" is disabled') + yield # pragma: no cover + return GraphDefinition( key=key, enabled=False, @@ -58,6 +67,7 @@ async def _route_disabled(*args: Any, **kwargs: Any) -> Any: edges_from=lambda k: [], run_node=_run_node_disabled, route=_route_disabled, + stream_route=_stream_route_disabled, traverse=_traverse_noop, reverse_traverse=_traverse_noop, ) @@ -259,35 +269,13 @@ async def run_node( ) raise - # ── route ───────────────────────────────────────────────────────────────── - # For nodes with zero/one outgoing edge, delegates to run_node and returns - # the sole successor as `next`. For multi-edge nodes, injects synthetic - # handoff tools so the model picks the next agent. Mirrors TS route(). + # ── build_handoff_routing ───────────────────────────────────────────────── + # Shared by route and stream_route so multi-edge descriptions stay identical. - async def route( + def build_handoff_routing( node: GraphNode, - input: str = "", - opts: dict[str, Any] | None = None, + out_edges: list[GraphEdge], ) -> dict[str, Any]: - opts = opts or {} - handlers: list[ProviderHandler] = options.get("handlers") or [] - if not handlers: - raise ValueError( - "route is not available when no handlers were provided — use a " - "framework-native runner (to_openai_agents, to_lang_graph, to_claude_agents) instead." - ) - - out_edges = edges_from(node.key) - - # Zero/one outgoing edge: run node directly; report sole child as next. - if len(out_edges) <= 1: - res = await run_node(node, input, opts) - next_node = nodes.get(out_edges[0].target_key) if out_edges else None - return {**res, "next": next_node} - - handler = select_handler(node.config, node.meta, handlers, strict=False) - tool_handlers = opts.get("tool_handlers") or options.get("tool_handlers") - chosen: list[str] = [] handoff_tools: dict[str, Any] = {} handoff_handlers: dict[str, Any] = {} @@ -332,7 +320,7 @@ def _fn(*a: Any, **kw: Any) -> str: handoff_handlers[tool_name] = _make_handoff_fn(target_key) - route_config: AiConfigRep = { + routed_config: AiConfigRep = { **node.config, "instructions": ( (node.config.get("instructions") or "") @@ -348,12 +336,52 @@ def _fn(*a: Any, **kw: Any) -> str: }, } + return { + "routed_config": routed_config, + "handoff_handlers": handoff_handlers, + "chosen": lambda: chosen[0] if chosen else None, + } + + # ── route ───────────────────────────────────────────────────────────────── + # For nodes with zero/one outgoing edge, delegates to run_node and returns + # the sole successor as `next`. For multi-edge nodes, injects synthetic + # handoff tools so the model picks the next agent. Mirrors TS route(). + + async def route( + node: GraphNode, + input: str = "", + opts: dict[str, Any] | None = None, + ) -> dict[str, Any]: + opts = opts or {} + handlers: list[ProviderHandler] = options.get("handlers") or [] + if not handlers: + raise ValueError( + "route is not available when no handlers were provided — use a " + "framework-native runner (to_openai_agents, to_lang_graph, to_claude_agents) instead." + ) + + out_edges = edges_from(node.key) + + # Zero/one outgoing edge: run node directly; report sole child as next. + if len(out_edges) <= 1: + res = await run_node(node, input, opts) + next_node = nodes.get(out_edges[0].target_key) if out_edges else None + return {**res, "next": next_node} + + handler = select_handler(node.config, node.meta, handlers, strict=False) + tool_handlers = opts.get("tool_handlers") or options.get("tool_handlers") + + routing = build_handoff_routing(node, out_edges) + routed_config = routing["routed_config"] + handoff_handlers = routing["handoff_handlers"] + chosen = routing["chosen"] + merged_tool_handlers = {**(tool_handlers or {}), **handoff_handlers} try: result = await execute_and_track( config_key=node.key, - config=route_config, + config=routed_config, meta=node.meta, user_context=context, handler=handler, @@ -382,7 +410,8 @@ def _fn(*a: Any, **kw: Any) -> str: graph_key=key, ) - next_node = nodes.get(chosen[0]) if chosen else None + chosen_key = chosen() + next_node = nodes.get(chosen_key) if chosen_key else None if next_node: get_client().track( @@ -403,14 +432,279 @@ def _fn(*a: Any, **kw: Any) -> str: "next": next_node, } except Exception: - if chosen: + chosen_key = chosen() + if chosen_key: get_client().track( "$ld:ai:graph:handoff_failure", ld_ctx, { **graph_track_data, "sourceKey": node.key, - "targetKey": chosen[0], + "targetKey": chosen_key, + }, + 1, + ) + raise + + # ── stream_node / stream_route ──────────────────────────────────────────── + # Streaming counterparts to run_node / route. Python async generators cannot + # return a value (unlike JS yield*), so callers pass an ``outcome`` dict that + # is populated with the ProviderResponse / RouteResult fields when done. + + async def stream_node( + node: GraphNode, + input: str = "", + opts: dict[str, Any] | None = None, + outcome: dict[str, Any] | None = None, + ) -> AsyncGenerator[GraphStreamEvent, None]: + from .tracking import execute_and_stream + + opts = opts or {} + handlers: list[ProviderHandler] = options.get("handlers") or [] + if not handlers: + raise ValueError( + "stream_node is not available when no handlers were provided — use a " + "framework-native runner (to_openai_agents, to_lang_graph, to_claude_agents) instead." + ) + handler = select_handler(node.config, node.meta, handlers, strict=False) + tool_handlers = opts.get("tool_handlers") or options.get("tool_handlers") + from_node: GraphNode | None = opts.get("from") + + yield {"type": "node_start", "nodeKey": node.key} + + try: + response = "" + usage: dict[str, Any] = {"input": 0, "output": 0, "total": 0} + track_data: TrackData = { + "runId": str(uuid.uuid4()), + "configKey": node.key, + "variationKey": ( + node.meta.get("variationKey", "") + if isinstance(node.meta, dict) + else "" + ), + "version": ( + node.meta.get("version", 1) if isinstance(node.meta, dict) else 1 + ), + "modelName": (node.config.get("model") or {}).get("name", ""), + "providerName": (node.config.get("provider") or {}).get("name", ""), + "graphKey": key, + } + + async for event in execute_and_stream( + config_key=node.key, + config=node.config, + meta=node.meta, + user_context=context, + handler=handler, + user_input=input, + tool_handlers=tool_handlers, + variables=opts.get("variables"), + graph_key=key, + history=opts.get("history"), + ): + if event.get("type") == "chunk": + yield { + "type": "chunk", + "text": event["text"], + "nodeKey": node.key, + } + else: + response = event.get("response", "") + usage = event.get("usage") or usage + track_data = event.get("track_data") or track_data + + judge_results = await run_judges( + config=node.config, + user_context=context, + handler=handler, + handlers=handlers, + user_input=input, + llm_response=response, + base_track_data=track_data, + tool_handlers=tool_handlers, + graph_key=key, + ) + + if from_node: + get_client().track( + "$ld:ai:graph:handoff_success", + ld_ctx, + { + **graph_track_data, + "sourceKey": from_node.key, + "targetKey": node.key, + }, + 1, + ) + + yield { + "type": "node_done", + "nodeKey": node.key, + "response": response, + "usage": usage, + } + if outcome is not None: + outcome.clear() + outcome.update( + { + "response": response, + "usage": usage, + "judge_results": judge_results, + "track_data": track_data, + } + ) + except Exception: + if from_node: + get_client().track( + "$ld:ai:graph:handoff_failure", + ld_ctx, + { + **graph_track_data, + "sourceKey": from_node.key, + "targetKey": node.key, + }, + 1, + ) + raise + + async def stream_route( + node: GraphNode, + input: str = "", + opts: dict[str, Any] | None = None, + outcome: dict[str, Any] | None = None, + ) -> AsyncGenerator[GraphStreamEvent, None]: + from .tracking import execute_and_stream + + opts = opts or {} + handlers: list[ProviderHandler] = options.get("handlers") or [] + if not handlers: + raise ValueError( + "stream_route is not available when no handlers were provided — use a " + "framework-native runner (to_openai_agents, to_lang_graph, to_claude_agents) instead." + ) + + out_edges = edges_from(node.key) + + if len(out_edges) <= 1: + node_outcome: dict[str, Any] = {} + async for event in stream_node(node, input, opts, node_outcome): + yield event + next_node = nodes.get(out_edges[0].target_key) if out_edges else None + if outcome is not None: + outcome.clear() + outcome.update({**node_outcome, "next": next_node}) + return + + handler = select_handler(node.config, node.meta, handlers, strict=False) + tool_handlers = opts.get("tool_handlers") or options.get("tool_handlers") + + routing = build_handoff_routing(node, out_edges) + routed_config = routing["routed_config"] + handoff_handlers = routing["handoff_handlers"] + chosen = routing["chosen"] + + yield {"type": "node_start", "nodeKey": node.key} + + try: + response = "" + usage: dict[str, Any] = {"input": 0, "output": 0, "total": 0} + track_data: TrackData = { + "runId": str(uuid.uuid4()), + "configKey": node.key, + "variationKey": ( + node.meta.get("variationKey", "") + if isinstance(node.meta, dict) + else "" + ), + "version": ( + node.meta.get("version", 1) if isinstance(node.meta, dict) else 1 + ), + "modelName": (node.config.get("model") or {}).get("name", ""), + "providerName": (node.config.get("provider") or {}).get("name", ""), + "graphKey": key, + } + + merged_tool_handlers = {**(tool_handlers or {}), **handoff_handlers} + + async for event in execute_and_stream( + config_key=node.key, + config=routed_config, + meta=node.meta, + user_context=context, + handler=handler, + user_input=input, + tool_handlers=merged_tool_handlers, + variables=opts.get("variables"), + graph_key=key, + history=opts.get("history"), + ): + if event.get("type") == "chunk": + yield { + "type": "chunk", + "text": event["text"], + "nodeKey": node.key, + } + else: + response = event.get("response", "") + usage = event.get("usage") or usage + track_data = event.get("track_data") or track_data + + # Judge against the node's original config, not the routing-augmented one. + judge_results = await run_judges( + config=node.config, + user_context=context, + handler=handler, + handlers=handlers, + user_input=input, + llm_response=response, + base_track_data=track_data, + tool_handlers=tool_handlers, + graph_key=key, + ) + + chosen_key = chosen() + next_node = nodes.get(chosen_key) if chosen_key else None + + if next_node: + get_client().track( + "$ld:ai:graph:handoff_success", + ld_ctx, + { + **graph_track_data, + "sourceKey": node.key, + "targetKey": next_node.key, + }, + 1, + ) + + yield { + "type": "node_done", + "nodeKey": node.key, + "response": response, + "usage": usage, + } + if outcome is not None: + outcome.clear() + outcome.update( + { + "response": response, + "usage": usage, + "judge_results": judge_results, + "track_data": track_data, + "next": next_node, + } + ) + except Exception: + chosen_key = chosen() + if chosen_key: + get_client().track( + "$ld:ai:graph:handoff_failure", + ld_ctx, + { + **graph_track_data, + "sourceKey": node.key, + "targetKey": chosen_key, }, 1, ) @@ -498,6 +792,7 @@ async def reverse_traverse(fn: Any, ctx: dict[str, Any] | None = None) -> Any: edges_from=edges_from, run_node=run_node, route=route, + stream_route=stream_route, traverse=traverse, reverse_traverse=reverse_traverse, ) @@ -550,6 +845,8 @@ async def invoke( variables: dict[str, Any] | None = None, history: list[dict[str, Any]] | None = None, ) -> ProviderGraphResponse: + from opentelemetry import trace + from .judges import run_judges from .lifecycle import get_client @@ -591,10 +888,227 @@ async def invoke( if not graph_def.enabled: raise ValueError(f'Agent graph "{self._key}" is disabled') + tracer = trace.get_tracer("@launchdarkly/ai-server") + with tracer.start_as_current_span("ld.ai.graph") as span: + span.set_attribute("ld.ai.graph.key", self._key) + + start_time = time.monotonic() + path: list[str] = [] + total_usage = {"input": 0, "output": 0, "total": 0} + resolved_input = user_input or "" + + try: + current: GraphNode | None = graph_def.root + previous_node: GraphNode | None = None + current_input = resolved_input + last: dict[str, Any] | None = None + visited: set[str] = set() + steps = 0 + + while current and steps < MAX_TRAVERSAL_DEPTH: + steps += 1 + opts: dict[str, Any] = {"variables": variables} + if previous_node: + opts["from"] = previous_node + # History seeds the entry point only. After the root hop, nodes + # stay oriented through the string threading built below, so + # history is not re-sent to downstream handlers. + elif history: + opts["history"] = history + + res = await graph_def.route(current, current_input, opts) + path.append(current.key) + total_usage["input"] += ( + res["usage"].get("input", 0) + if isinstance(res["usage"], dict) + else 0 + ) + total_usage["output"] += ( + res["usage"].get("output", 0) + if isinstance(res["usage"], dict) + else 0 + ) + total_usage["total"] += ( + res["usage"].get("total", 0) + if isinstance(res["usage"], dict) + else 0 + ) + last = res + + next_node = res.get("next") + if not next_node or next_node.key in visited: + break + visited.add(current.key) + previous_node = current + current = next_node + current_input = "\n\n".join( + [ + f"[Original request]\n{resolved_input}", + f"[Previous agent response]\n{res['response']}", + ] + ) + + final_response = (last or {}).get("response", "") + + elapsed_ms = int((time.monotonic() - start_time) * 1000) + client = get_client() + client.track( + "$ld:ai:graph:duration:total", ld_ctx, graph_track_data, elapsed_ms + ) + if total_usage["total"] > 0: + client.track( + "$ld:ai:graph:total_tokens", + ld_ctx, + graph_track_data, + total_usage["total"], + ) + client.track( + "$ld:ai:graph:path", + ld_ctx, + {**graph_track_data, "path": path}, + len(path), + ) + client.track( + "$ld:ai:graph:invocation_success", ld_ctx, graph_track_data, 1 + ) + + # Optional graph-level judge run against the final response. + judge_results: dict[str, JudgeResult] | None = None + graph_judge: str | None = resolved_options.get("graph_judge") + root_node = graph_def.root + if graph_judge and root_node and resolved_handlers: + judge_handler = select_handler( + root_node.config, + root_node.meta, + resolved_handlers, + strict=False, + ) + judge_results = await run_judges( + config={ + "judgeConfiguration": { + "judges": [{"key": graph_judge, "samplingRate": 1}] + } + }, + user_context=context, + handler=judge_handler, + handlers=resolved_handlers, + user_input=resolved_input, + llm_response=final_response, + base_track_data=graph_track_data, + tool_handlers=resolved_tools, + graph_key=self._key, + ) + + return ProviderGraphResponse( + response=final_response, + # Named rather than splatted, so a new UsageDict member cannot silently arrive + # here from a dict that has no business filling it. Graph totals carry no cache + # breakdown: they are a sum across nodes, and the per-node detail is on the node's + # own spans. + usage=UsageDict( + input=total_usage["input"], + output=total_usage["output"], + total=total_usage["total"], + ), + judge_results=judge_results, + ) + + except Exception: + elapsed_ms = int((time.monotonic() - start_time) * 1000) + client = get_client() + client.track( + "$ld:ai:graph:duration:total", ld_ctx, graph_track_data, elapsed_ms + ) + client.track( + "$ld:ai:graph:invocation_failure", ld_ctx, graph_track_data, 1 + ) + raise + + def stream( + self, + user_input: str | None, + context: LDContext, + variables: dict[str, Any] | None = None, + history: list[dict[str, Any]] | None = None, + ) -> AsyncGenerator[GraphStreamEvent, None]: + """Stream graph traversal events. + + Deliberately not an ``async def`` with ``yield``: a generator body does not run until the + first ``__anext__``, by which point a ``conversation_id`` / caller span scope wrapped around + this call may have already exited. Binding the conversation id and capturing the OTel parent + here — at call time — matches ``config().stream()`` and the TypeScript graph stream. + """ + from opentelemetry import context as otel_context + + caller_context = otel_context.get_current() + return bind_conversation_id( + self._stream_events( + user_input, context, variables, history, caller_context + ) + ) + + async def _stream_events( + self, + user_input: str | None, + context: LDContext, + variables: dict[str, Any] | None, + history: list[dict[str, Any]] | None, + caller_context: Any, + ) -> AsyncGenerator[GraphStreamEvent, None]: + from opentelemetry import context as otel_context + from opentelemetry import trace + from opentelemetry.trace import Status, StatusCode, set_span_in_context + + from .judges import run_judges + from .lifecycle import get_client + + resolved_input = user_input or "" + ld_ctx = to_ld_context(get_client(), context) + + resolved_handlers = resolve_handlers( + self._options.get("registry"), self._options.get("handlers") + ) + resolved_tools = resolve_tools( + self._options.get("registry"), self._options.get("tool_handlers") + ) + resolved_options = { + **self._options, + "handlers": resolved_handlers, + "tool_handlers": resolved_tools, + } + + if not resolved_handlers: + raise ValueError( + "graph().stream() requires handlers to be provided. Pass handlers in options, or " + "use resolve_graph() with a framework-native runner." + ) + + try: + cache_key: str | None = json.dumps(context, sort_keys=True) + except (TypeError, ValueError): + cache_key = None + if cache_key is not None and cache_key in self._cache: + built = self._cache[cache_key] + else: + built = await _build_graph(self._key, context, resolved_options) + if cache_key is not None: + if len(self._cache) >= MAX_GRAPH_CACHE_SIZE: + self._cache.pop(next(iter(self._cache))) + self._cache[cache_key] = built + graph_def, graph_track_data = built + + if not graph_def.enabled: + raise ValueError(f'Agent graph "{self._key}" is disabled') + + tracer = trace.get_tracer("@launchdarkly/ai-server") + span = tracer.start_span("ld.ai.graph", context=caller_context) + span.set_attribute("ld.ai.graph.key", self._key) + span_context = set_span_in_context(span, caller_context) + ended: set[int] = set() start_time = time.monotonic() + path: list[str] = [] total_usage = {"input": 0, "output": 0, "total": 0} - resolved_input = user_input or "" try: current: GraphNode | None = graph_def.root @@ -606,44 +1120,54 @@ async def invoke( while current and steps < MAX_TRAVERSAL_DEPTH: steps += 1 - opts: dict[str, Any] = {"variables": variables} + route_opts: dict[str, Any] = {"variables": variables} if previous_node: - opts["from"] = previous_node + route_opts["from"] = previous_node # History seeds the entry point only. After the root hop, nodes # stay oriented through the string threading built below, so # history is not re-sent to downstream handlers. elif history: - opts["history"] = history + route_opts["history"] = history + + outcome: dict[str, Any] = {} + async for event in bind_span_context( + graph_def.stream_route( + current, current_input, route_opts, outcome + ), + span_context, + ): + yield event - res = await graph_def.route(current, current_input, opts) path.append(current.key) + usage = outcome.get("usage") or {} total_usage["input"] += ( - res["usage"].get("input", 0) - if isinstance(res["usage"], dict) - else 0 + usage.get("input", 0) if isinstance(usage, dict) else 0 ) total_usage["output"] += ( - res["usage"].get("output", 0) - if isinstance(res["usage"], dict) - else 0 + usage.get("output", 0) if isinstance(usage, dict) else 0 ) total_usage["total"] += ( - res["usage"].get("total", 0) - if isinstance(res["usage"], dict) - else 0 + usage.get("total", 0) if isinstance(usage, dict) else 0 ) - last = res + last = outcome - next_node = res.get("next") + next_node = outcome.get("next") if not next_node or next_node.key in visited: break + + yield { + "type": "handoff", + "sourceKey": current.key, + "targetKey": next_node.key, + } + visited.add(current.key) previous_node = current current = next_node current_input = "\n\n".join( [ f"[Original request]\n{resolved_input}", - f"[Previous agent response]\n{res['response']}", + f"[Previous agent response]\n{outcome.get('response', '')}", ] ) @@ -669,55 +1193,66 @@ async def invoke( ) client.track("$ld:ai:graph:invocation_success", ld_ctx, graph_track_data, 1) - # Optional graph-level judge run against the final response. judge_results: dict[str, JudgeResult] | None = None graph_judge: str | None = resolved_options.get("graph_judge") root_node = graph_def.root if graph_judge and root_node and resolved_handlers: - judge_handler = select_handler( - root_node.config, - root_node.meta, - resolved_handlers, - strict=False, - ) - judge_results = await run_judges( - config={ - "judgeConfiguration": { - "judges": [{"key": graph_judge, "samplingRate": 1}] - } - }, - user_context=context, - handler=judge_handler, - handlers=resolved_handlers, - user_input=resolved_input, - llm_response=final_response, - base_track_data=graph_track_data, - tool_handlers=resolved_tools, - graph_key=self._key, - ) - - return ProviderGraphResponse( - response=final_response, - # Named rather than splatted, so a new UsageDict member cannot silently arrive - # here from a dict that has no business filling it. Graph totals carry no cache - # breakdown: they are a sum across nodes, and the per-node detail is on the node's - # own spans. - usage=UsageDict( - input=total_usage["input"], - output=total_usage["output"], - total=total_usage["total"], - ), - judge_results=judge_results, - ) + # Re-enter the graph span explicitly. bind_span_context only covers the + # delegated per-node generator; this call runs in the generator body. + token = otel_context.attach(span_context) + try: + judge_handler = select_handler( + root_node.config, + root_node.meta, + resolved_handlers, + strict=False, + ) + results = await run_judges( + config={ + "judgeConfiguration": { + "judges": [{"key": graph_judge, "samplingRate": 1}] + } + }, + user_context=context, + handler=judge_handler, + handlers=resolved_handlers, + user_input=resolved_input, + llm_response=final_response, + base_track_data=graph_track_data, + tool_handlers=resolved_tools, + graph_key=self._key, + ) + finally: + otel_context.detach(token) + if results: + judge_results = results + + span.set_status(Status(StatusCode.OK)) + end_span_once(span, ended) + + done_event: GraphStreamEvent = { + "type": "done", + "response": final_response, + "usage": total_usage, + } + if judge_results: + done_event["judgeResults"] = judge_results + yield done_event - except Exception: + except Exception as err: elapsed_ms = int((time.monotonic() - start_time) * 1000) client = get_client() client.track( "$ld:ai:graph:duration:total", ld_ctx, graph_track_data, elapsed_ms ) client.track("$ld:ai:graph:invocation_failure", ld_ctx, graph_track_data, 1) + span.record_exception(err) + span.set_status(Status(StatusCode.ERROR, str(err))) + end_span_once(span, ended) raise + finally: + end_span_once(span, ended, abandoned=True) + def graph( diff --git a/packages/client/src/launchdarkly_ai_server/tracking.py b/packages/client/src/launchdarkly_ai_server/tracking.py index 9beae0ba..4aeeedd5 100644 --- a/packages/client/src/launchdarkly_ai_server/tracking.py +++ b/packages/client/src/launchdarkly_ai_server/tracking.py @@ -80,14 +80,23 @@ def stub(*args: Any, **kwargs: Any) -> None: def _make_regular_wrapper( tool_name: str, original: Callable[..., Any] ) -> Callable[..., Any]: + # Synthetic graph handoff tools are sync and skipped for tool_call metrics. + # Keep them sync so a bare handoff() records the chosen edge (stream and + # invoke share the same surface; an async wrapper would leave `chosen` empty). + if tool_name.startswith("__handoff_"): + + def handoff_wrapper(*args: Any, **kwargs: Any) -> Any: + return original(*args, **kwargs) + + return handoff_wrapper + async def wrapper(*args: Any, **kwargs: Any) -> Any: - if not tool_name.startswith("__handoff_"): - get_client().track( - "$ld:ai:tool_call", - user_context, - {**track_data, "toolKey": tool_name}, - 1, - ) + get_client().track( + "$ld:ai:tool_call", + user_context, + {**track_data, "toolKey": tool_name}, + 1, + ) result = original(*args, **kwargs) return await result if inspect.isawaitable(result) else result diff --git a/packages/client/src/launchdarkly_ai_server/types.py b/packages/client/src/launchdarkly_ai_server/types.py index dcbc5701..86660bff 100644 --- a/packages/client/src/launchdarkly_ai_server/types.py +++ b/packages/client/src/launchdarkly_ai_server/types.py @@ -327,6 +327,14 @@ class JudgeRunResult: StreamDoneEvent = dict[str, Any] # {"type": "done", "response": str, "usage": ..., ...} StreamEvent = dict[str, Any] # StreamChunkEvent | StreamDoneEvent +# Graph stream events (camelCase public fields, matching the TypeScript SDK) +# node_start: {type, nodeKey} +# chunk: {type, text, nodeKey} +# node_done: {type, nodeKey, response, usage} +# handoff: {type, sourceKey, targetKey} +# done: {type, response, usage, judgeResults?} +GraphStreamEvent = dict[str, Any] + # Internal execute stream event that also carries track_data ExecuteStreamDoneEvent = dict[str, Any] ExecuteStreamEvent = dict[str, Any] @@ -392,6 +400,7 @@ def __init__( edges_from: Callable[[str], list[GraphEdge]], run_node: Callable[..., Any], route: Callable[..., Any], + stream_route: Callable[..., Any], traverse: Callable[..., Any], reverse_traverse: Callable[..., Any], ) -> None: @@ -406,6 +415,7 @@ def __init__( self.edges_from = edges_from self.run_node = run_node self.route = route + self.stream_route = stream_route self.traverse = traverse self.reverse_traverse = reverse_traverse diff --git a/packages/client/tests/test_graph_stream.py b/packages/client/tests/test_graph_stream.py new file mode 100644 index 00000000..010e77f2 --- /dev/null +++ b/packages/client/tests/test_graph_stream.py @@ -0,0 +1,700 @@ +""" +Tests for §3.15a ``graph().stream()``. +Reference: TESTING.md §3.15a, Appendix A.4 / A.13. +""" + +from __future__ import annotations + +from collections.abc import AsyncGenerator, Iterator +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from opentelemetry import trace +from opentelemetry.sdk.trace import ReadableSpan, TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +import launchdarkly_ai_server.lifecycle as lifecycle_module +from launchdarkly_ai_server import ProviderHandler, graph +from launchdarkly_ai_server.conversation import ( + GEN_AI_CONVERSATION_ID, + ConversationIdSpanProcessor, + conversation_id, +) + +CONTEXT = {"kind": "user", "key": "u1"} + +_exporter = InMemorySpanExporter() +_provider = TracerProvider() +_provider.add_span_processor(ConversationIdSpanProcessor()) +_provider.add_span_processor(SimpleSpanProcessor(_exporter)) +# Same as the JS suite's setGlobalTracerProvider: SDK code uses trace.get_tracer(), so the +# test provider must be global or ld.ai.graph spans never reach this exporter. +trace.set_tracer_provider(_provider) +_tracer = _provider.get_tracer("@launchdarkly/ai-server") + + +def _node_variation(instructions: str = "Be helpful.") -> dict[str, Any]: + return { + "model": {"name": "gpt-4"}, + "provider": {"name": "TestProvider"}, + "instructions": instructions, + "_ldMeta": { + "enabled": True, + "variationKey": "v1", + "version": 1, + "mode": "messages", + }, + } + + +def _make_client(graph_variation: dict | None = None) -> MagicMock: + c = MagicMock() + c.track = MagicMock() + c.flush = AsyncMock() + c.close = AsyncMock() + + graph_var = graph_variation or { + "root": "root-node", + "edges": {"root-node": [{"key": "leaf-node"}]}, + } + + node_instructions = { + "root-node": "I am root", + "leaf-node": "I am leaf", + "agent-a": "I am A", + "agent-b": "I am B", + } + + async def variation_side_effect(key: str, ctx: dict, default: Any) -> Any: + if key == "graph-key": + return graph_var + return _node_variation(node_instructions.get(key, "Be helpful.")) + + c.variation = AsyncMock(side_effect=variation_side_effect) + return c + + +def _make_streaming_handler( + chunks: list[str] | None = None, + usage: dict | None = None, +) -> ProviderHandler: + _chunks = chunks or ["Hi", "!"] + _usage = usage or {"input_tokens": 2, "output_tokens": 3} + + async def fn(config, user_input, tool_handlers, variables, history=None) -> dict: # type: ignore[override] + return {"output": "".join(_chunks), "usage": _usage} + + async def stream_fn( + config, user_input, tool_handlers, variables, history=None + ) -> AsyncGenerator: # type: ignore[override] + for c in _chunks: + yield {"type": "chunk", "text": c} + yield {"type": "done", "output": "".join(_chunks), "usage": _usage} + + return ProviderHandler( + fn=fn, provides_for=("TestProvider", "messages"), stream_fn=stream_fn + ) + + +def _make_blocking_handler(response: str = "blocked") -> ProviderHandler: + async def fn(config, user_input, tool_handlers, variables, history=None) -> dict: # type: ignore[override] + return { + "output": response, + "usage": {"input_tokens": 1, "output_tokens": 1}, + } + + return ProviderHandler(fn=fn, provides_for=("TestProvider", "messages")) # type: ignore[arg-type] + + +def _make_branch_picking_stream_handler(pick_target: str) -> ProviderHandler: + sanitized = "".join(c if c.isalnum() or c == "_" else "_" for c in pick_target) + _usage = {"input_tokens": 1, "output_tokens": 1} + received: list[tuple[Any, Any]] = [] + + async def fn(config, user_input, tool_handlers, variables, history=None) -> dict: # type: ignore[override] + return {"output": "ok", "usage": _usage} + + async def stream_fn( + config, user_input, tool_handlers, variables, history=None + ) -> AsyncGenerator: # type: ignore[override] + received.append((config, tool_handlers)) + handoff = (tool_handlers or {}).get(f"__handoff_{sanitized}") + if handoff: + handoff() + yield {"type": "chunk", "text": "ok"} + yield {"type": "done", "output": "ok", "usage": _usage} + + h = ProviderHandler( + fn=fn, provides_for=("TestProvider", "messages"), stream_fn=stream_fn + ) + h._test_received = received # type: ignore[attr-defined] + return h + + +def _span_creating_stream_handler() -> ProviderHandler: + _chunks = ["ok"] + _usage = {"input_tokens": 1, "output_tokens": 1} + + async def fn(config, user_input, tool_handlers, variables, history=None) -> dict: # type: ignore[override] + return {"output": "ok", "usage": _usage} + + async def stream_fn( + config, user_input, tool_handlers, variables, history=None + ) -> AsyncGenerator: # type: ignore[override] + root = _tracer.start_span("handler.invoke_agent") + for c in _chunks: + with trace.use_span(root, end_on_exit=False): + chat = _tracer.start_span("handler.chat") + chat.end() + yield {"type": "chunk", "text": c} + root.end() + yield {"type": "done", "output": "".join(_chunks), "usage": _usage} + + return ProviderHandler( + fn=fn, provides_for=("TestProvider", "messages"), stream_fn=stream_fn + ) + + +async def _collect(gen: Any) -> list[Any]: + return [e async for e in gen] + + +def _track_names(client: MagicMock) -> list[str]: + return [c[0][0] for c in client.track.call_args_list] + + +def _track_payload(client: MagicMock, name: str) -> Any: + for c in client.track.call_args_list: + if c[0][0] == name: + return c[0][2] if len(c[0]) > 2 else c[1] + return None + + +def _finished() -> list[ReadableSpan]: + return list(_exporter.get_finished_spans()) + + +@pytest.fixture +def mock_ld_client() -> Iterator[MagicMock]: + client = _make_client() + lifecycle_module._set_client_for_testing(client) + yield client + lifecycle_module._reset_for_testing() + + +@pytest.fixture(autouse=True) +def _reset_exporter() -> Iterator[None]: + _exporter.clear() + yield + _exporter.clear() + + +# --------------------------------------------------------------------------- +# Setup / errors / traversal events +# --------------------------------------------------------------------------- + + +class TestGraphStream: + def test_returns_async_generator(self, mock_ld_client: MagicMock) -> None: + gen = graph("graph-key", handlers=[_make_streaming_handler()]).stream( + "hi", CONTEXT + ) + assert hasattr(gen, "__aiter__") + + async def test_throws_when_graph_disabled(self, mock_ld_client: MagicMock) -> None: + mock_ld_client.variation = AsyncMock(return_value={"edges": {}}) + gen = graph("graph-key", handlers=[_make_streaming_handler()]).stream( + "hi", CONTEXT + ) + with pytest.raises((ValueError, RuntimeError), match="disabled"): + await _collect(gen) + + async def test_throws_when_no_handlers(self, mock_ld_client: MagicMock) -> None: + gen = graph("graph-key").stream("hi", CONTEXT) + with pytest.raises((ValueError, RuntimeError)): + await _collect(gen) + + async def test_emits_node_start_in_order(self, mock_ld_client: MagicMock) -> None: + events = await _collect( + graph("graph-key", handlers=[_make_streaming_handler(["ok"])]).stream( + "hi", CONTEXT + ) + ) + starts = [e for e in events if e["type"] == "node_start"] + assert [e["nodeKey"] for e in starts] == ["root-node", "leaf-node"] + + async def test_forwards_chunks_tagged_with_node_key( + self, mock_ld_client: MagicMock + ) -> None: + events = await _collect( + graph("graph-key", handlers=[_make_streaming_handler(["Hi", "!"])]).stream( + "hi", CONTEXT + ) + ) + chunks = [e for e in events if e["type"] == "chunk"] + assert chunks == [ + {"type": "chunk", "text": "Hi", "nodeKey": "root-node"}, + {"type": "chunk", "text": "!", "nodeKey": "root-node"}, + {"type": "chunk", "text": "Hi", "nodeKey": "leaf-node"}, + {"type": "chunk", "text": "!", "nodeKey": "leaf-node"}, + ] + + async def test_emits_node_done_with_response_and_usage( + self, mock_ld_client: MagicMock + ) -> None: + events = await _collect( + graph( + "graph-key", + handlers=[ + _make_streaming_handler( + ["Hi", "!"], {"input_tokens": 2, "output_tokens": 3} + ) + ], + ).stream("hi", CONTEXT) + ) + dones = [e for e in events if e["type"] == "node_done"] + assert len(dones) == 2 + assert dones[0] == { + "type": "node_done", + "nodeKey": "root-node", + "response": "Hi!", + "usage": {"input": 2, "output": 3, "total": 5}, + } + assert dones[1]["nodeKey"] == "leaf-node" + assert dones[1]["usage"] == {"input": 2, "output": 3, "total": 5} + + async def test_emits_handoff_between_node_done_and_next_start( + self, mock_ld_client: MagicMock + ) -> None: + events = await _collect( + graph("graph-key", handlers=[_make_streaming_handler(["ok"])]).stream( + "hi", CONTEXT + ) + ) + types = [e["type"] for e in events] + idx = types.index("handoff") + assert events[idx] == { + "type": "handoff", + "sourceKey": "root-node", + "targetKey": "leaf-node", + } + assert types[idx - 1] == "node_done" + assert types[idx + 1] == "node_start" + + async def test_final_done_has_leaf_response_and_aggregate_usage( + self, mock_ld_client: MagicMock + ) -> None: + events = await _collect( + graph( + "graph-key", + handlers=[ + _make_streaming_handler( + ["final"], {"input_tokens": 2, "output_tokens": 3} + ) + ], + ).stream("hi", CONTEXT) + ) + done = events[-1] + assert done["type"] == "done" + assert done["response"] == "final" + assert done["usage"] == {"input": 4, "output": 6, "total": 10} + assert "path" not in done + assert "nodes" not in done + + async def test_lifecycle_events_before_final_done( + self, mock_ld_client: MagicMock + ) -> None: + events = await _collect( + graph("graph-key", handlers=[_make_streaming_handler(["a"])]).stream( + "hi", CONTEXT + ) + ) + assert events[-1]["type"] == "done" + assert all(e["type"] != "done" for e in events[:-1]) + + async def test_tracks_invocation_success(self, mock_ld_client: MagicMock) -> None: + await _collect( + graph("graph-key", handlers=[_make_streaming_handler(["ok"])]).stream( + "hi", CONTEXT + ) + ) + assert "$ld:ai:graph:invocation_success" in _track_names(mock_ld_client) + + async def test_tracks_duration_total(self, mock_ld_client: MagicMock) -> None: + await _collect( + graph("graph-key", handlers=[_make_streaming_handler(["ok"])]).stream( + "hi", CONTEXT + ) + ) + assert "$ld:ai:graph:duration:total" in _track_names(mock_ld_client) + + async def test_tracks_path(self, mock_ld_client: MagicMock) -> None: + await _collect( + graph("graph-key", handlers=[_make_streaming_handler(["ok"])]).stream( + "hi", CONTEXT + ) + ) + assert "$ld:ai:graph:path" in _track_names(mock_ld_client) + + async def test_tracks_handoff_success(self, mock_ld_client: MagicMock) -> None: + await _collect( + graph("graph-key", handlers=[_make_streaming_handler(["ok"])]).stream( + "hi", CONTEXT + ) + ) + payload = _track_payload(mock_ld_client, "$ld:ai:graph:handoff_success") + assert payload is not None + assert payload["sourceKey"] == "root-node" + assert payload["targetKey"] == "leaf-node" + + async def test_tracks_invocation_failure_and_rethrows( + self, mock_ld_client: MagicMock + ) -> None: + async def fn(config, user_input, tool_handlers, variables, history=None) -> dict: # type: ignore[override] + return {"output": "x", "usage": {"input_tokens": 1, "output_tokens": 1}} + + async def stream_fn( + config, user_input, tool_handlers, variables, history=None + ) -> AsyncGenerator: # type: ignore[override] + raise RuntimeError("stream boom") + yield # pragma: no cover + + h = ProviderHandler( + fn=fn, provides_for=("TestProvider", "messages"), stream_fn=stream_fn + ) + with pytest.raises(RuntimeError, match="stream boom"): + await _collect(graph("graph-key", handlers=[h]).stream("hi", CONTEXT)) + assert "$ld:ai:graph:invocation_failure" in _track_names(mock_ld_client) + + async def test_generation_success_includes_graph_key( + self, mock_ld_client: MagicMock + ) -> None: + await _collect( + graph("graph-key", handlers=[_make_streaming_handler(["ok"])]).stream( + "hi", CONTEXT + ) + ) + success = [ + c + for c in mock_ld_client.track.call_args_list + if c[0][0] == "$ld:ai:generation:success" + ] + assert len(success) >= 2 + for call in success: + assert call[0][2]["graphKey"] == "graph-key" + + async def test_falls_back_to_blocking_handler( + self, mock_ld_client: MagicMock + ) -> None: + events = await _collect( + graph("graph-key", handlers=[_make_blocking_handler("blocked")]).stream( + "hi", CONTEXT + ) + ) + chunks = [e for e in events if e["type"] == "chunk"] + assert chunks == [ + {"type": "chunk", "text": "blocked", "nodeKey": "root-node"}, + {"type": "chunk", "text": "blocked", "nodeKey": "leaf-node"}, + ] + assert events[-1]["type"] == "done" + assert events[-1]["response"] == "blocked" + + async def test_includes_graph_judge_results_on_done( + self, mock_ld_client: MagicMock + ) -> None: + judge_data = { + "graph-judge": { + "usage": {"input": 1, "output": 1, "total": 2}, + "response": "ok", + "score": 0.8, + } + } + with patch( + "launchdarkly_ai_server.judges.run_judges", + new_callable=AsyncMock, + return_value=judge_data, + ) as run_judges: + events = await _collect( + graph( + "graph-key", + handlers=[_make_streaming_handler(["final"])], + graph_judge="graph-judge", + ).stream("hi", CONTEXT) + ) + assert run_judges.await_count >= 1 + assert events[-1]["type"] == "done" + assert events[-1]["judgeResults"] == judge_data + + async def test_omits_judge_results_when_empty( + self, mock_ld_client: MagicMock + ) -> None: + with patch( + "launchdarkly_ai_server.judges.run_judges", + new_callable=AsyncMock, + return_value={}, + ): + events = await _collect( + graph( + "graph-key", + handlers=[_make_streaming_handler(["final"])], + graph_judge="graph-judge", + ).stream("hi", CONTEXT) + ) + assert "judgeResults" not in events[-1] + + +# --------------------------------------------------------------------------- +# Multi-edge routing +# --------------------------------------------------------------------------- + + +class TestGraphStreamMultiEdge: + @pytest.fixture + def mock_ld_client(self) -> Iterator[MagicMock]: + client = _make_client( + { + "root": "root-node", + "edges": { + "root-node": [{"key": "agent-a"}, {"key": "agent-b"}], + }, + } + ) + lifecycle_module._set_client_for_testing(client) + yield client + lifecycle_module._reset_for_testing() + + async def test_model_pick_emits_handoff_success( + self, mock_ld_client: MagicMock + ) -> None: + h = _make_branch_picking_stream_handler("agent-b") + await _collect(graph("graph-key", handlers=[h]).stream("hi", CONTEXT)) + handoffs = [ + c + for c in mock_ld_client.track.call_args_list + if c[0][0] == "$ld:ai:graph:handoff_success" + ] + assert len(handoffs) >= 1 + assert handoffs[0][0][2]["sourceKey"] == "root-node" + assert handoffs[0][0][2]["targetKey"] == "agent-b" + + async def test_handoff_failure_when_node_throws_after_choice( + self, mock_ld_client: MagicMock + ) -> None: + sanitized = "agent_a" + _usage = {"input_tokens": 1, "output_tokens": 1} + + async def fn(config, user_input, tool_handlers, variables, history=None) -> dict: # type: ignore[override] + return {"output": "ok", "usage": _usage} + + async def stream_fn( + config, user_input, tool_handlers, variables, history=None + ) -> AsyncGenerator: # type: ignore[override] + handoff = (tool_handlers or {}).get(f"__handoff_{sanitized}") + if handoff: + handoff() + raise RuntimeError("boom after choice") + yield # pragma: no cover + + h = ProviderHandler( + fn=fn, provides_for=("TestProvider", "messages"), stream_fn=stream_fn + ) + with pytest.raises(RuntimeError, match="boom after choice"): + await _collect(graph("graph-key", handlers=[h]).stream("hi", CONTEXT)) + payload = _track_payload(mock_ld_client, "$ld:ai:graph:handoff_failure") + assert payload is not None + assert payload["sourceKey"] == "root-node" + assert payload["targetKey"] == "agent-a" + + async def test_judges_receive_original_config( + self, mock_ld_client: MagicMock + ) -> None: + h = _make_branch_picking_stream_handler("agent-b") + with patch( + "launchdarkly_ai_server.judges.run_judges", + new_callable=AsyncMock, + return_value={}, + ) as run_judges: + await _collect(graph("graph-key", handlers=[h]).stream("hi", CONTEXT)) + + root_call = next( + ( + c + for c in run_judges.await_args_list + if (c.kwargs.get("config") or {}).get("instructions") == "I am root" + ), + None, + ) + assert root_call is not None + judged = root_call.kwargs["config"] + assert judged["instructions"] == "I am root" + assert not any(k.startswith("__handoff_") for k in (judged.get("tools") or {})) + + async def test_handoff_tools_and_instructions_match_invoke( + self, mock_ld_client: MagicMock + ) -> None: + stream_h = _make_branch_picking_stream_handler("agent-b") + await _collect(graph("graph-key", handlers=[stream_h]).stream("hi", CONTEXT)) + stream_cfg, stream_tools = stream_h._test_received[0] # type: ignore[attr-defined] + + invoke_received: list[tuple[Any, Any]] = [] + + async def invoke_fn( + config, user_input, tool_handlers, variables, history=None + ) -> dict: # type: ignore[override] + invoke_received.append((config, tool_handlers)) + handoff = (tool_handlers or {}).get("__handoff_agent_b") + if handoff: + handoff() + return {"output": "ok", "usage": {"input_tokens": 1, "output_tokens": 1}} + + invoke_h = ProviderHandler( + fn=invoke_fn, provides_for=("TestProvider", "messages") + ) # type: ignore[arg-type] + await graph("graph-key", handlers=[invoke_h]).invoke("hi", CONTEXT) + invoke_cfg, invoke_tools = invoke_received[0] + + assert stream_cfg["tools"]["__handoff_agent_a"]["description"] == invoke_cfg[ + "tools" + ]["__handoff_agent_a"]["description"] + assert stream_cfg["tools"]["__handoff_agent_b"]["description"] == invoke_cfg[ + "tools" + ]["__handoff_agent_b"]["description"] + assert stream_cfg["instructions"] == invoke_cfg["instructions"] + assert stream_tools["__handoff_agent_b"]() == invoke_tools["__handoff_agent_b"]() + + +# --------------------------------------------------------------------------- +# Conversation id + OTel parenting + abandonment (§3.15a / A.4) +# --------------------------------------------------------------------------- + + +class TestGraphStreamOtel: + async def test_stamps_conversation_id_when_bound_at_call_time( + self, mock_ld_client: MagicMock + ) -> None: + with conversation_id("thread-graph-stream"): + gen = graph( + "graph-key", handlers=[_span_creating_stream_handler()] + ).stream("hi", CONTEXT) + await _collect(gen) + + graph_spans = [s for s in _finished() if s.name == "ld.ai.graph"] + assert len(graph_spans) >= 1 + assert graph_spans[0].attributes + assert ( + graph_spans[0].attributes.get(GEN_AI_CONVERSATION_ID) + == "thread-graph-stream" + ) + + async def test_handler_spans_nest_under_graph_on_stream( + self, mock_ld_client: MagicMock + ) -> None: + await _collect( + graph("graph-key", handlers=[_span_creating_stream_handler()]).stream( + "hi", CONTEXT + ) + ) + spans = _finished() + graph_spans = [s for s in spans if s.name == "ld.ai.graph"] + handler_spans = [s for s in spans if s.name.startswith("handler.")] + assert len(graph_spans) >= 1 + assert len(handler_spans) >= 1 + gctx = graph_spans[0].get_span_context() + for hs in handler_spans: + assert hs.get_span_context().trace_id == gctx.trace_id + # Ancestor: parent chain reaches the graph span + parent = hs.parent + assert parent is not None + + async def test_handler_spans_nest_under_graph_on_invoke( + self, mock_ld_client: MagicMock + ) -> None: + async def fn(config, user_input, tool_handlers, variables, history=None) -> dict: # type: ignore[override] + root = _tracer.start_span("handler.invoke_agent") + with trace.use_span(root, end_on_exit=False): + chat = _tracer.start_span("handler.chat") + chat.end() + root.end() + return {"output": "ok", "usage": {"input_tokens": 1, "output_tokens": 1}} + + h = ProviderHandler(fn=fn, provides_for=("TestProvider", "messages")) # type: ignore[arg-type] + await graph("graph-key", handlers=[h]).invoke("hi", CONTEXT) + + spans = _finished() + graph_spans = [s for s in spans if s.name == "ld.ai.graph"] + handler_spans = [s for s in spans if s.name.startswith("handler.")] + assert len(graph_spans) >= 1 + assert len(handler_spans) >= 1 + gctx = graph_spans[0].get_span_context() + for hs in handler_spans: + assert hs.get_span_context().trace_id == gctx.trace_id + + async def test_graph_parent_captured_at_stream_call_time( + self, mock_ld_client: MagicMock + ) -> None: + caller = _tracer.start_span("caller") + with trace.use_span(caller, end_on_exit=False): + gen = graph( + "graph-key", handlers=[_make_streaming_handler(["ok"])] + ).stream("hi", CONTEXT) + caller.end() + await _collect(gen) + + graph_spans = [s for s in _finished() if s.name == "ld.ai.graph"] + assert len(graph_spans) >= 1 + assert graph_spans[0].parent is not None + assert graph_spans[0].parent.span_id == caller.get_span_context().span_id + + async def test_abandoned_on_consumer_break(self, mock_ld_client: MagicMock) -> None: + gen = graph( + "graph-key", handlers=[_make_streaming_handler(["a", "b"])] + ).stream("hi", CONTEXT) + async for event in gen: + if event["type"] == "chunk": + break + await gen.aclose() + + graph_spans = [s for s in _finished() if s.name == "ld.ai.graph"] + assert len(graph_spans) >= 1 + attrs = graph_spans[0].attributes or {} + assert attrs.get("launchdarkly.stream.abandoned") is True + assert "$ld:ai:graph:invocation_success" not in _track_names(mock_ld_client) + + async def test_graph_judge_spans_nest_under_graph( + self, mock_ld_client: MagicMock + ) -> None: + async def run_judges_with_span(*args: Any, **kwargs: Any) -> dict: + span = _tracer.start_span("graph.judge") + span.end() + return { + "graph-judge": { + "usage": {"input": 1, "output": 1, "total": 2}, + "response": "ok", + "score": 0.9, + } + } + + with patch( + "launchdarkly_ai_server.judges.run_judges", + new_callable=AsyncMock, + side_effect=run_judges_with_span, + ): + await _collect( + graph( + "graph-key", + handlers=[_make_streaming_handler(["final"])], + graph_judge="graph-judge", + ).stream("hi", CONTEXT) + ) + + spans = _finished() + graph_spans = [s for s in spans if s.name == "ld.ai.graph"] + judge_spans = [s for s in spans if s.name == "graph.judge"] + assert len(graph_spans) >= 1 + assert len(judge_spans) >= 1 + assert ( + judge_spans[0].get_span_context().trace_id + == graph_spans[0].get_span_context().trace_id + ) diff --git a/packages/client/tests/test_tracking.py b/packages/client/tests/test_tracking.py index 65a31389..8a625d3f 100644 --- a/packages/client/tests/test_tracking.py +++ b/packages/client/tests/test_tracking.py @@ -70,15 +70,18 @@ def test_native_tool_instance_preserved_on_stub(self) -> None: stub = wrapped["search"] assert getattr(stub, NATIVE_TOOL_KEY) is native - async def test_handoff_prefix_skips_tracking(self) -> None: + def test_handoff_prefix_skips_tracking(self) -> None: mock_client = _make_mock_client() original = MagicMock(return_value=None) with patch("launchdarkly_ai_server.lifecycle._client", mock_client): wrapped = wrap_tool_handlers( {"__handoff_leaf": original}, CONTEXT, TRACK_DATA ) - await wrapped["__handoff_leaf"]() + # Handoff tools stay sync so stream/invoke routing can record the choice + # with a bare call (an async wrapper would leave `chosen` empty). + wrapped["__handoff_leaf"]() mock_client.track.assert_not_called() + original.assert_called_once() def test_undefined_tool_handlers(self) -> None: result = wrap_tool_handlers(None, CONTEXT, TRACK_DATA) From b03ceb1a7101ed2e63921e9687c8eab685cdae24 Mon Sep 17 00:00:00 2001 From: Jeff Dupont Date: Wed, 23 Sep 2026 14:49:01 -0700 Subject: [PATCH 2/3] [AIC-3210] fix(AIC-3210): keep stream_route off GraphDefinition and stop claiming the global tracer MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit CI on #104 failed three separate ways. All three are fixed here. 1. `stream_route` was added as a required keyword-only argument to `GraphDefinition.__init__`, which broke every direct construction of the type: 63 tests across claude-agents, langchain-agents and openai-agents failed with `TypeError: missing 1 required keyword-only argument`. Those packages consume a `GraphDefinition` but never build one outside tests, so the fix is at the source rather than in their helpers — none of their test files are touched. `_build_graph` now returns `stream_route` alongside the definition, mirroring `buildGraph`'s `{ def, graphTrackData, streamRoute }` in the JS SDK and keeping `resolve_graph()`'s contract stable. TESTING.md never lists stream_route on the GraphDefinition surface, and A.2's live-vs- disabled key parity still holds. 2. `test_graph_stream.py` called `trace.set_tracer_provider()` at import. That global is first-writer-wins, and claude-agents' `test_handler.py` already claims it at import and collects first alphabetically, so the new file's exporter received nothing and 6 OTel assertions failed on `0 >= 1`. The suite now patches `get_tracer` per test instead — `graph.py` imports `trace` inside its functions, so that is the seam. Both file orders now pass; previously each order broke whichever file lost the race. 3. The two new files were never formatted, failing `ruff format --check`. Verified locally: uv lock, ruff check, ruff format --check, mypy packages/*/src, uv run pytest (1328 passed, 11 skipped), and uv build for every package all exit 0. Co-Authored-By: Claude Opus 5 (1M context) --- .../src/launchdarkly_ai_server/graph.py | 56 ++++++++++-------- .../src/launchdarkly_ai_server/types.py | 2 - packages/client/tests/test_graph_stream.py | 58 +++++++++++-------- 3 files changed, 68 insertions(+), 48 deletions(-) diff --git a/packages/client/src/launchdarkly_ai_server/graph.py b/packages/client/src/launchdarkly_ai_server/graph.py index 317eff73..444e9009 100644 --- a/packages/client/src/launchdarkly_ai_server/graph.py +++ b/packages/client/src/launchdarkly_ai_server/graph.py @@ -49,12 +49,6 @@ async def _run_node_disabled(*args: Any, **kwargs: Any) -> Any: async def _route_disabled(*args: Any, **kwargs: Any) -> Any: raise ValueError(f'Agent graph "{key}" is disabled') - async def _stream_route_disabled( - *args: Any, **kwargs: Any - ) -> AsyncGenerator[GraphStreamEvent, None]: - raise ValueError(f'Agent graph "{key}" is disabled') - yield # pragma: no cover - return GraphDefinition( key=key, enabled=False, @@ -67,12 +61,24 @@ async def _stream_route_disabled( edges_from=lambda k: [], run_node=_run_node_disabled, route=_route_disabled, - stream_route=_stream_route_disabled, traverse=_traverse_noop, reverse_traverse=_traverse_noop, ) +def _disabled_stream_route(key: str) -> Callable[..., Any]: + """The ``stream_route`` a disabled graph gets. Kept beside the definition it + pairs with, but off ``GraphDefinition`` — see ``_build_graph``'s return.""" + + async def _stream_route_disabled( + *args: Any, **kwargs: Any + ) -> AsyncGenerator[GraphStreamEvent, None]: + raise ValueError(f'Agent graph "{key}" is disabled') + yield # pragma: no cover + + return _stream_route_disabled + + async def _fetch_graph_variation( key: str, context: LDContext, @@ -99,7 +105,7 @@ async def _build_graph( key: str, context: LDContext, options: dict[str, Any], -) -> tuple[GraphDefinition, TrackData]: +) -> tuple[GraphDefinition, TrackData, Callable[..., Any]]: from .judges import run_judges from .lifecycle import extract_variation, get_client from .tracking import execute_and_track @@ -123,7 +129,11 @@ async def _build_graph( } if not enabled or not topology: - return _disabled_definition(key), graph_track_data + return ( + _disabled_definition(key), + graph_track_data, + _disabled_stream_route(key), + ) raw_edges_map: dict[str, list[dict[str, Any]]] = topology.get("edges") or {} edges: list[GraphEdge] = [] @@ -162,7 +172,11 @@ def edges_from(node_key: str) -> list[GraphEdge]: ) except Exception as exc: logger.error("Graph node variation failed: %s", exc) - return _disabled_definition(key), graph_track_data + return ( + _disabled_definition(key), + graph_track_data, + _disabled_stream_route(key), + ) root_node = nodes.get(topology["root"]) @@ -792,12 +806,11 @@ async def reverse_traverse(fn: Any, ctx: dict[str, Any] | None = None) -> Any: edges_from=edges_from, run_node=run_node, route=route, - stream_route=stream_route, traverse=traverse, reverse_traverse=reverse_traverse, ) - return graph_def, graph_track_data + return graph_def, graph_track_data, stream_route async def resolve_graph( @@ -822,7 +835,7 @@ async def resolve_graph( "tool_handlers": resolved_tools, "registry": registry, } - graph_def, _ = await _build_graph(key, context, options) + graph_def, _, _ = await _build_graph(key, context, options) return graph_def @@ -836,7 +849,9 @@ def __init__( ) -> None: self._key = key self._options = options - self._cache: dict[str, tuple[GraphDefinition, TrackData]] = {} + self._cache: dict[ + str, tuple[GraphDefinition, TrackData, Callable[..., Any]] + ] = {} async def invoke( self, @@ -883,7 +898,7 @@ async def invoke( # Evict an arbitrary entry to keep the cache bounded. self._cache.pop(next(iter(self._cache))) self._cache[cache_key] = built - graph_def, graph_track_data = built + graph_def, graph_track_data, _ = built if not graph_def.enabled: raise ValueError(f'Agent graph "{self._key}" is disabled') @@ -1042,9 +1057,7 @@ def stream( caller_context = otel_context.get_current() return bind_conversation_id( - self._stream_events( - user_input, context, variables, history, caller_context - ) + self._stream_events(user_input, context, variables, history, caller_context) ) async def _stream_events( @@ -1095,7 +1108,7 @@ async def _stream_events( if len(self._cache) >= MAX_GRAPH_CACHE_SIZE: self._cache.pop(next(iter(self._cache))) self._cache[cache_key] = built - graph_def, graph_track_data = built + graph_def, graph_track_data, stream_route = built if not graph_def.enabled: raise ValueError(f'Agent graph "{self._key}" is disabled') @@ -1131,9 +1144,7 @@ async def _stream_events( outcome: dict[str, Any] = {} async for event in bind_span_context( - graph_def.stream_route( - current, current_input, route_opts, outcome - ), + stream_route(current, current_input, route_opts, outcome), span_context, ): yield event @@ -1254,7 +1265,6 @@ async def _stream_events( end_span_once(span, ended, abandoned=True) - def graph( key: str, *, diff --git a/packages/client/src/launchdarkly_ai_server/types.py b/packages/client/src/launchdarkly_ai_server/types.py index 86660bff..7248f022 100644 --- a/packages/client/src/launchdarkly_ai_server/types.py +++ b/packages/client/src/launchdarkly_ai_server/types.py @@ -400,7 +400,6 @@ def __init__( edges_from: Callable[[str], list[GraphEdge]], run_node: Callable[..., Any], route: Callable[..., Any], - stream_route: Callable[..., Any], traverse: Callable[..., Any], reverse_traverse: Callable[..., Any], ) -> None: @@ -415,7 +414,6 @@ def __init__( self.edges_from = edges_from self.run_node = run_node self.route = route - self.stream_route = stream_route self.traverse = traverse self.reverse_traverse = reverse_traverse diff --git a/packages/client/tests/test_graph_stream.py b/packages/client/tests/test_graph_stream.py index 010e77f2..785b5a2e 100644 --- a/packages/client/tests/test_graph_stream.py +++ b/packages/client/tests/test_graph_stream.py @@ -29,9 +29,6 @@ _provider = TracerProvider() _provider.add_span_processor(ConversationIdSpanProcessor()) _provider.add_span_processor(SimpleSpanProcessor(_exporter)) -# Same as the JS suite's setGlobalTracerProvider: SDK code uses trace.get_tracer(), so the -# test provider must be global or ld.ai.graph spans never reach this exporter. -trace.set_tracer_provider(_provider) _tracer = _provider.get_tracer("@launchdarkly/ai-server") @@ -187,7 +184,12 @@ def mock_ld_client() -> Iterator[MagicMock]: @pytest.fixture(autouse=True) def _reset_exporter() -> Iterator[None]: _exporter.clear() - yield + # Hand the SDK this file's tracer without touching the global provider: + # set_tracer_provider is process-wide and first-writer-wins, and another package's + # suite already claims it at import time, so registering here exports nothing. + # graph.py imports `trace` inside its functions, so the seam is get_tracer itself. + with patch.object(trace, "get_tracer", return_value=_tracer): + yield _exporter.clear() @@ -352,7 +354,9 @@ async def test_tracks_handoff_success(self, mock_ld_client: MagicMock) -> None: async def test_tracks_invocation_failure_and_rethrows( self, mock_ld_client: MagicMock ) -> None: - async def fn(config, user_input, tool_handlers, variables, history=None) -> dict: # type: ignore[override] + async def fn( + config, user_input, tool_handlers, variables, history=None + ) -> dict: # type: ignore[override] return {"output": "x", "usage": {"input_tokens": 1, "output_tokens": 1}} async def stream_fn( @@ -485,7 +489,9 @@ async def test_handoff_failure_when_node_throws_after_choice( sanitized = "agent_a" _usage = {"input_tokens": 1, "output_tokens": 1} - async def fn(config, user_input, tool_handlers, variables, history=None) -> dict: # type: ignore[override] + async def fn( + config, user_input, tool_handlers, variables, history=None + ) -> dict: # type: ignore[override] return {"output": "ok", "usage": _usage} async def stream_fn( @@ -555,14 +561,18 @@ async def invoke_fn( await graph("graph-key", handlers=[invoke_h]).invoke("hi", CONTEXT) invoke_cfg, invoke_tools = invoke_received[0] - assert stream_cfg["tools"]["__handoff_agent_a"]["description"] == invoke_cfg[ - "tools" - ]["__handoff_agent_a"]["description"] - assert stream_cfg["tools"]["__handoff_agent_b"]["description"] == invoke_cfg[ - "tools" - ]["__handoff_agent_b"]["description"] + assert ( + stream_cfg["tools"]["__handoff_agent_a"]["description"] + == invoke_cfg["tools"]["__handoff_agent_a"]["description"] + ) + assert ( + stream_cfg["tools"]["__handoff_agent_b"]["description"] + == invoke_cfg["tools"]["__handoff_agent_b"]["description"] + ) assert stream_cfg["instructions"] == invoke_cfg["instructions"] - assert stream_tools["__handoff_agent_b"]() == invoke_tools["__handoff_agent_b"]() + assert ( + stream_tools["__handoff_agent_b"]() == invoke_tools["__handoff_agent_b"]() + ) # --------------------------------------------------------------------------- @@ -575,9 +585,9 @@ async def test_stamps_conversation_id_when_bound_at_call_time( self, mock_ld_client: MagicMock ) -> None: with conversation_id("thread-graph-stream"): - gen = graph( - "graph-key", handlers=[_span_creating_stream_handler()] - ).stream("hi", CONTEXT) + gen = graph("graph-key", handlers=[_span_creating_stream_handler()]).stream( + "hi", CONTEXT + ) await _collect(gen) graph_spans = [s for s in _finished() if s.name == "ld.ai.graph"] @@ -611,7 +621,9 @@ async def test_handler_spans_nest_under_graph_on_stream( async def test_handler_spans_nest_under_graph_on_invoke( self, mock_ld_client: MagicMock ) -> None: - async def fn(config, user_input, tool_handlers, variables, history=None) -> dict: # type: ignore[override] + async def fn( + config, user_input, tool_handlers, variables, history=None + ) -> dict: # type: ignore[override] root = _tracer.start_span("handler.invoke_agent") with trace.use_span(root, end_on_exit=False): chat = _tracer.start_span("handler.chat") @@ -636,9 +648,9 @@ async def test_graph_parent_captured_at_stream_call_time( ) -> None: caller = _tracer.start_span("caller") with trace.use_span(caller, end_on_exit=False): - gen = graph( - "graph-key", handlers=[_make_streaming_handler(["ok"])] - ).stream("hi", CONTEXT) + gen = graph("graph-key", handlers=[_make_streaming_handler(["ok"])]).stream( + "hi", CONTEXT + ) caller.end() await _collect(gen) @@ -648,9 +660,9 @@ async def test_graph_parent_captured_at_stream_call_time( assert graph_spans[0].parent.span_id == caller.get_span_context().span_id async def test_abandoned_on_consumer_break(self, mock_ld_client: MagicMock) -> None: - gen = graph( - "graph-key", handlers=[_make_streaming_handler(["a", "b"])] - ).stream("hi", CONTEXT) + gen = graph("graph-key", handlers=[_make_streaming_handler(["a", "b"])]).stream( + "hi", CONTEXT + ) async for event in gen: if event["type"] == "chunk": break From 30428e68fc83ecc7fe6b40dd29bea23341e2ff37 Mon Sep 17 00:00:00 2001 From: Jeff Dupont Date: Thu, 24 Sep 2026 07:45:40 -0700 Subject: [PATCH 3/3] [AIC-3210] fix(AIC-3210): mark a cancelled graph stream as run.cancelled A timeout or task.cancel() was stamping launchdarkly.stream.abandoned on the graph span while handler spans recorded launchdarkly.run.cancelled. Co-authored-by: Cursor --- .../src/launchdarkly_ai_server/graph.py | 12 ++++- packages/client/tests/test_graph_stream.py | 54 +++++++++++++++++++ 2 files changed, 65 insertions(+), 1 deletion(-) diff --git a/packages/client/src/launchdarkly_ai_server/graph.py b/packages/client/src/launchdarkly_ai_server/graph.py index 444e9009..0ef93b08 100644 --- a/packages/client/src/launchdarkly_ai_server/graph.py +++ b/packages/client/src/launchdarkly_ai_server/graph.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import inspect as _inspect import json import logging @@ -1123,6 +1124,12 @@ async def _stream_events( path: list[str] = [] total_usage = {"input": 0, "output": 0, "total": 0} + # A consumer that stops reading unwinds as GeneratorExit, and that is abandonment. + # A timeout or task.cancel() unwinds as CancelledError. That is a BaseException, so + # `except Exception` never sees it, and the finally below would otherwise stamp + # launchdarkly.stream.abandoned on a run whose handler spans say + # launchdarkly.run.cancelled. cancelled wins inside end_span_once. + cancelled = False try: current: GraphNode | None = graph_def.root previous_node: GraphNode | None = None @@ -1250,6 +1257,9 @@ async def _stream_events( done_event["judgeResults"] = judge_results yield done_event + except asyncio.CancelledError: + cancelled = True + raise except Exception as err: elapsed_ms = int((time.monotonic() - start_time) * 1000) client = get_client() @@ -1262,7 +1272,7 @@ async def _stream_events( end_span_once(span, ended) raise finally: - end_span_once(span, ended, abandoned=True) + end_span_once(span, ended, abandoned=True, cancelled=cancelled) def graph( diff --git a/packages/client/tests/test_graph_stream.py b/packages/client/tests/test_graph_stream.py index 785b5a2e..5d7265bd 100644 --- a/packages/client/tests/test_graph_stream.py +++ b/packages/client/tests/test_graph_stream.py @@ -5,6 +5,7 @@ from __future__ import annotations +import asyncio from collections.abc import AsyncGenerator, Iterator from typing import Any from unittest.mock import AsyncMock, MagicMock, patch @@ -14,6 +15,7 @@ from opentelemetry.sdk.trace import ReadableSpan, TracerProvider from opentelemetry.sdk.trace.export import SimpleSpanProcessor from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.trace import StatusCode import launchdarkly_ai_server.lifecycle as lifecycle_module from launchdarkly_ai_server import ProviderHandler, graph @@ -674,6 +676,58 @@ async def test_abandoned_on_consumer_break(self, mock_ld_client: MagicMock) -> N assert attrs.get("launchdarkly.stream.abandoned") is True assert "$ld:ai:graph:invocation_success" not in _track_names(mock_ld_client) + async def test_cancelled_stream_marks_run_cancelled_not_abandoned( + self, mock_ld_client: MagicMock + ) -> None: + # A consumer that stops reading abandoned the stream. A CancelledError is not that + # choice: a timeout or task.cancel() ends the run underneath the consumer. The handler + # spans already say launchdarkly.run.cancelled for that unwind, so the graph span has + # to say the same thing. Sleeping in the handler, not in the consumer loop, is what + # makes the unwind a CancelledError; a break in the loop body is GeneratorExit. + parked = asyncio.Event() + + async def fn( + config, user_input, tool_handlers, variables, history=None + ) -> dict: # type: ignore[override] + return {"output": "a", "usage": {"input_tokens": 1, "output_tokens": 1}} + + async def stream_fn( + config, user_input, tool_handlers, variables, history=None + ) -> AsyncGenerator: # type: ignore[override] + yield {"type": "chunk", "text": "a"} + parked.set() + await asyncio.sleep(3600) + yield { + "type": "done", + "output": "a", + "usage": {"input_tokens": 1, "output_tokens": 1}, + } + + handler = ProviderHandler( + fn=fn, provides_for=("TestProvider", "messages"), stream_fn=stream_fn + ) + + async def _drain() -> None: + gen = graph("graph-key", handlers=[handler]).stream("hi", CONTEXT) + async for _event in gen: + pass + + task = asyncio.create_task(_drain()) + await asyncio.wait_for(parked.wait(), timeout=2) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + graph_spans = [s for s in _finished() if s.name == "ld.ai.graph"] + assert len(graph_spans) == 1 + attrs = graph_spans[0].attributes or {} + assert attrs.get("launchdarkly.run.cancelled") is True + assert "launchdarkly.stream.abandoned" not in attrs + assert graph_spans[0].status.status_code == StatusCode.UNSET + names = _track_names(mock_ld_client) + assert "$ld:ai:graph:invocation_success" not in names + assert "$ld:ai:graph:invocation_failure" not in names + async def test_graph_judge_spans_nest_under_graph( self, mock_ld_client: MagicMock ) -> None: