diff --git a/README.md b/README.md index a20b875..84bb619 100644 --- a/README.md +++ b/README.md @@ -227,6 +227,30 @@ go from loading to serving, so the duration of that window is observed there. Th reopens on every later recovery, so a reload after an out-of-memory eviction or a lost node is measured too — not only the first start. +### Tracing + +Every chat completion is one OpenTelemetry span carrying what is needed to explain the route: +tenant, effective and declared privacy class, task, model and revision, adapter, cache result, +score, candidate count, the route reason, and token usage. The response quotes it as +`X-Trace-Id`. Only the OpenTelemetry API is a runtime dependency, so tracing is a no-op until a +collector is configured: + +```bash +python -m pip install -e ".[tracing]" +ROUTER_OTLP_ENDPOINT=http://otel-collector:4318/v1/traces +``` + +Prompts are redacted by privacy class, evaluated after any tenant floor has been applied: + +| Class | Recorded | +|---|---| +| `restricted` | Length only. No digest: a digest of a short or templated prompt can be reversed by guessing. | +| `private` | Length and a SHA-256 digest, so repeats can be correlated. | +| `public` | Length and digest; a bounded prefix of the content only with `ROUTER_TRACE_PROMPT_CONTENT=true`. | + +Completions are never recorded. A failed request records its error type and not its message, +because an engine error can echo the request it rejected. + ## Caching Cache keys bind the tenant, the active model-catalog fingerprint, and the generation @@ -278,6 +302,9 @@ All settings use the `ROUTER_` prefix. | `ROUTER_QUOTA_REQUESTS_PER_MINUTE` | `120` | Per-token sliding-window quota. | | `ROUTER_EXTERNAL_FALLBACK_ENABLED` | `false` | Operator gate for external fallback. | | `ROUTER_REDIS_URL` | _(empty)_ | Shared cache and quota state; in-process when empty. | +| `ROUTER_TENANT_KEYS` | _(empty)_ | `tenant:key` bindings; bare `ROUTER_API_KEYS` keys use the default tenant. | +| `ROUTER_OTLP_ENDPOINT` | _(empty)_ | OTLP/HTTP trace collector; tracing is a no-op when empty. | +| `ROUTER_TRACE_PROMPT_CONTENT` | `false` | Records a bounded prefix of `public` prompts only. | | `ROUTER_BACKEND` | `mock` | `mock` or `vllm`. | | `ROUTER_VLLM_BASE_URL` | `http://127.0.0.1:8001` | vLLM OpenAI-compatible endpoint. | | `ROUTER_BACKEND_TIMEOUT_SECONDS` | `60` | Per-request engine timeout. | diff --git a/pyproject.toml b/pyproject.toml index 1a23bdf..2b425d6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,6 +11,7 @@ requires-python = ">=3.11" dependencies = [ "fastapi>=0.141.1,<1", "httpx>=0.28.1,<1", + "opentelemetry-api>=1.30,<2", "prometheus-client>=0.26.0,<1", "pyyaml>=6.0.3,<7", "pydantic-settings>=2.15.0,<3", @@ -21,8 +22,13 @@ dependencies = [ redis = [ "redis>=8.1.0,<9", ] +tracing = [ + "opentelemetry-exporter-otlp-proto-http>=1.30,<2", + "opentelemetry-sdk>=1.30,<2", +] dev = [ "mypy>=2.3.1,<3", + "opentelemetry-sdk>=1.30,<2", "pytest>=9.1.1,<10", "pytest-asyncio>=1.4.0,<2", "pytest-cov>=7.1.0,<8", @@ -68,3 +74,9 @@ packages = ["llm_router"] module = ["redis.*"] ignore_missing_imports = true follow_imports = "skip" + +[[tool.mypy.overrides]] +# The OTLP exporter ships in the optional tracing extra and is imported lazily, +# only when a collector endpoint is configured. +module = ["opentelemetry.exporter.*"] +ignore_missing_imports = true diff --git a/src/llm_router/app.py b/src/llm_router/app.py index 5bfdc32..9d749b6 100644 --- a/src/llm_router/app.py +++ b/src/llm_router/app.py @@ -12,6 +12,7 @@ import httpx from fastapi import Depends, FastAPI, Header, HTTPException, Request, Response, status from fastapi.responses import JSONResponse, StreamingResponse +from opentelemetry.trace import TracerProvider from llm_router.admission import ( AdmissionController, @@ -60,6 +61,7 @@ strictest_privacy, ) from llm_router.routing import NoEligibleModelError, Router, default_model_profiles +from llm_router.tracing import RequestSpan, Tracing, build_tracer_provider @dataclass(frozen=True) @@ -102,6 +104,7 @@ def create_app( registry: Registry | None = None, redis_client: RedisLike | None = None, engine_stats: EngineStatsCollector | None = None, + tracer_provider: TracerProvider | None = None, ) -> FastAPI: runtime_settings = settings or get_settings() catalog = registry if registry is not None else _load_catalog(runtime_settings.registry_path) @@ -144,6 +147,10 @@ def create_app( else None ) cold_start = ColdStartTracker() + tracing = Tracing( + tracer_provider or build_tracer_provider(runtime_settings), + record_prompt_content=runtime_settings.trace_prompt_content, + ) # One gateway deployment faces one engine target, so engine-level telemetry # and cold starts are attributed to that target rather than to a model. engine_label = runtime_settings.backend @@ -479,54 +486,85 @@ async def _stream_completion( decision: RouteDecision, started: float, queue_seconds: float, + span: RequestSpan, ) -> AsyncIterator[str]: completion_id = f"chatcmpl-{uuid.uuid4().hex}" created = int(time.time()) model_id = decision.profile.id collected: list[str] = [] - yield _chunk(completion_id, created, model_id, delta={"role": "assistant"}) - async for delta in inference_backend.stream(payload, decision): - collected.append(delta) - yield _chunk(completion_id, created, model_id, delta={"content": delta}) - yield _chunk(completion_id, created, model_id, delta={}, finish_reason="stop") - yield "data: [DONE]\n\n" + # The stream outlives the handler, so it owns the span from here and + # ends it whether generation finishes, fails, or the client leaves. + try: + yield _chunk(completion_id, created, model_id, delta={"role": "assistant"}) + async for delta in inference_backend.stream(payload, decision): + collected.append(delta) + yield _chunk(completion_id, created, model_id, delta={"content": delta}) + yield _chunk(completion_id, created, model_id, delta={}, finish_reason="stop") + yield "data: [DONE]\n\n" - text = "".join(collected) - prompt_tokens = max(1, len(prompt) // 4) - completion_tokens = max(1, len(text) // 4) - if payload.routing.structured: - telemetry.record_structured_output(decision, valid=structured_output_valid(text)) - telemetry.record_completion( - decision, - latency_seconds=time.perf_counter() - started, - queue_seconds=queue_seconds, - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - ) - await _store_cache( - payload, - prompt, - cache_key, - subject, - decision, - BackendResult( - text=text, prompt_tokens=prompt_tokens, completion_tokens=completion_tokens - ), - ) + text = "".join(collected) + prompt_tokens = max(1, len(prompt) // 4) + completion_tokens = max(1, len(text) // 4) + span.set_usage(prompt_tokens=prompt_tokens, completion_tokens=completion_tokens) + if payload.routing.structured: + telemetry.record_structured_output(decision, valid=structured_output_valid(text)) + telemetry.record_completion( + decision, + latency_seconds=time.perf_counter() - started, + queue_seconds=queue_seconds, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + ) + await _store_cache( + payload, + prompt, + cache_key, + subject, + decision, + BackendResult( + text=text, prompt_tokens=prompt_tokens, completion_tokens=completion_tokens + ), + ) + except Exception as error: + span.fail(error) + raise + finally: + span.end() @app.post("/v1/chat/completions", response_model=None) async def chat_completions( payload: ChatCompletionRequest, response: Response, principal: Annotated[Principal, Depends(authenticate)], + ) -> ChatCompletionResponse | StreamingResponse: + span = tracing.start_request() + try: + result = await _complete(payload, response, principal, span) + except Exception as error: + span.fail(error) + raise + finally: + if not span.handed_off: + span.end() + trace_id = span.trace_id + if trace_id is not None: + # A response returned directly does not inherit injected headers. + target = result if isinstance(result, Response) else response + target.headers["X-Trace-Id"] = trace_id + return result + + async def _complete( + payload: ChatCompletionRequest, + response: Response, + principal: Principal, + span: RequestSpan, ) -> ChatCompletionResponse | StreamingResponse: started = time.perf_counter() tenant = _entitlement(principal) # Quota and cache are scoped to the tenant, not the credential, so # rotating a key neither resets a quota nor orphans a cache. subject = principal.tenant_id - await _consume_quota(subject, tenant) declared_privacy = payload.routing.privacy if tenant is not None: @@ -543,6 +581,11 @@ async def chat_completions( privacy_raised_from = ( declared_privacy if payload.routing.privacy is not declared_privacy else None ) + # Described before the quota is charged so a rejected request is still + # attributed to its tenant, and after the floor so the trace is redacted + # under the effective class rather than the declared one. + span.set_request(payload, tenant_id=subject, declared_privacy=declared_privacy) + await _consume_quota(subject, tenant) prompt = payload.prompt cache_key = build_cache_key( @@ -550,8 +593,13 @@ async def chat_completions( ) cached = await _lookup_cache(payload, prompt, cache_key, subject) + span.set_cache("miss" if cached is None else cached[0]) if cached is not None: hit_name, entry = cached + span.set_served_from_cache(model_id=entry.model_id, model_revision=entry.model_revision) + span.set_usage( + prompt_tokens=entry.prompt_tokens, completion_tokens=entry.completion_tokens + ) if payload.stream: return _cached_stream(entry, hit_name) response.headers["X-Cache"] = hit_name @@ -583,6 +631,7 @@ async def chat_completions( ), privacy_raised_from=privacy_raised_from, ) + span.set_route(decision) if runtime_settings.cache_enabled: decision_cache.set(prompt, decision.task) telemetry.record_route(decision, privacy=payload.routing.privacy.value) @@ -601,6 +650,7 @@ async def chat_completions( telemetry.inflight_requests.inc() try: if payload.stream: + span.handed_off = True return StreamingResponse( _stream_completion( payload, @@ -610,6 +660,7 @@ async def chat_completions( decision, started, queue_seconds, + span, ), media_type="text/event-stream", headers=_route_headers(decision, "miss"), @@ -621,6 +672,9 @@ async def chat_completions( telemetry.queued_requests.dec() raise + span.set_usage( + prompt_tokens=result.prompt_tokens, completion_tokens=result.completion_tokens + ) if payload.routing.structured: telemetry.record_structured_output(decision, valid=structured_output_valid(result.text)) telemetry.record_completion( diff --git a/src/llm_router/config.py b/src/llm_router/config.py index 7a1471a..ba45041 100644 --- a/src/llm_router/config.py +++ b/src/llm_router/config.py @@ -27,6 +27,11 @@ class Settings(BaseSettings): registry_path: str = "config/registry.yaml" routing_policy_version: str = "v1" redis_url: str = "" + # Traces are exported only when a collector endpoint is set. Prompt content + # is never recorded unless an operator opts in, and then only for public + # requests; restricted and private prompts stay out of traces regardless. + otlp_endpoint: str = "" + trace_prompt_content: bool = False cache_enabled: bool = True cache_ttl_seconds: float = Field(default=300.0, gt=0) cache_max_entries: int = Field(default=1024, ge=1) diff --git a/src/llm_router/tracing.py b/src/llm_router/tracing.py new file mode 100644 index 0000000..9a1f178 --- /dev/null +++ b/src/llm_router/tracing.py @@ -0,0 +1,183 @@ +"""Request tracing with prompt redaction (sections 7.1 and 14). + +Section 7.1 has the gateway attach trace and routing metadata to every request, +and section 14 requires sensitive prompts to be redacted from traces. A span +therefore carries everything needed to explain a route (tenant, task, model, +revision, adapter, cache result, reason) and, by default, nothing of what the +caller actually wrote. + +Only the OpenTelemetry API is a runtime dependency. Without a configured +provider the tracer is a no-op, so tracing costs nothing until an operator +points it at a collector. +""" + +import hashlib +from typing import Any + +from opentelemetry import trace +from opentelemetry.trace import Span, SpanKind, Status, StatusCode, TracerProvider + +from llm_router.config import Settings +from llm_router.models import ChatCompletionRequest, PrivacyClass, RouteDecision + +PROMPT_CONTENT_LIMIT = 512 + + +def prompt_attributes( + prompt: str, privacy: PrivacyClass, *, record_content: bool +) -> dict[str, str | int]: + """Describe a prompt for a trace without disclosing more than its class allows. + + Restricted prompts contribute their length only: even a digest is withheld, + because a digest of a short or templated prompt can be reversed by guessing. + Private prompts add a digest so repeats can be correlated. Content is + recorded only for public prompts, only when an operator opted in, and only + up to a bounded prefix. + """ + + attributes: dict[str, str | int] = {"router.prompt.chars": len(prompt)} + if privacy is PrivacyClass.RESTRICTED: + return attributes + attributes["router.prompt.sha256"] = hashlib.sha256(prompt.encode()).hexdigest() + if privacy is PrivacyClass.PUBLIC and record_content: + attributes["router.prompt.content"] = prompt[:PROMPT_CONTENT_LIMIT] + return attributes + + +class RequestSpan: + """One chat-completion request, annotated as the gateway learns about it.""" + + def __init__(self, span: Span, *, record_content: bool) -> None: + self._span = span + self._record_content = record_content + # A streamed response outlives the handler, so the stream takes over + # ending the span and the handler must not end it first. + self.handed_off = False + + @property + def trace_id(self) -> str | None: + """The trace identifier callers can quote, or None when nothing records.""" + + context = self._span.get_span_context() + return f"{context.trace_id:032x}" if context.is_valid else None + + def _set(self, attributes: dict[str, Any]) -> None: + for name, value in attributes.items(): + if value is not None: + self._span.set_attribute(name, value) + + def set_request( + self, + request: ChatCompletionRequest, + *, + tenant_id: str, + declared_privacy: PrivacyClass, + ) -> None: + """Record the request under its effective privacy class.""" + + privacy = request.routing.privacy + self._set( + { + "gen_ai.operation.name": "chat", + "gen_ai.request.model": request.model, + "gen_ai.request.max_tokens": request.max_tokens, + "gen_ai.request.temperature": request.temperature, + "router.tenant": tenant_id, + "router.privacy": privacy.value, + "router.privacy.declared": declared_privacy.value, + "router.latency_tier": request.routing.latency_tier, + "router.stream": request.stream, + **prompt_attributes(request.prompt, privacy, record_content=self._record_content), + } + ) + + def set_cache(self, result: str) -> None: + self._set({"router.cache": result}) + + def set_route(self, decision: RouteDecision) -> None: + self._set( + { + "gen_ai.response.model": decision.profile.id, + "router.model.revision": decision.profile.revision, + "router.model.local": decision.profile.local, + "router.adapter": decision.adapter_id, + "router.adapter.revision": decision.adapter_revision, + "router.task": decision.task.value, + "router.route.reason": decision.reason, + "router.route.score": decision.score, + "router.route.candidates": decision.candidate_count, + } + ) + + def set_served_from_cache(self, *, model_id: str, model_revision: str) -> None: + self._set({"gen_ai.response.model": model_id, "router.model.revision": model_revision}) + + def set_usage(self, *, prompt_tokens: int, completion_tokens: int) -> None: + self._set( + { + "gen_ai.usage.input_tokens": prompt_tokens, + "gen_ai.usage.output_tokens": completion_tokens, + } + ) + + def fail(self, error: BaseException) -> None: + """Mark the span failed by error type alone. + + The message is withheld: an engine error can echo the request it + rejected, and a trace must not become a side channel for prompt text. + """ + + self._set({"error.type": type(error).__name__}) + self._span.set_status(Status(StatusCode.ERROR)) + + def end(self) -> None: + self._span.end() + + +class Tracing: + """Starts request spans against whichever provider the process configured.""" + + def __init__( + self, provider: TracerProvider | None = None, *, record_prompt_content: bool = False + ) -> None: + self._tracer = (provider or trace.get_tracer_provider()).get_tracer("llm_router") + self._record_content = record_prompt_content + + def start_request(self) -> RequestSpan: + span = self._tracer.start_span("chat", kind=SpanKind.SERVER) + return RequestSpan(span, record_content=self._record_content) + + +def build_tracer_provider(settings: Settings) -> TracerProvider | None: + """Build an exporting provider when a collector endpoint is configured. + + Returns None when tracing is not configured, which leaves the API's no-op + tracer in place. + """ + + if not settings.otlp_endpoint: + return None + return _otlp_provider(settings) # pragma: no cover - needs the tracing extra + + +def _otlp_provider(settings: Settings) -> TracerProvider: # pragma: no cover - needs the extra + try: + from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter + from opentelemetry.sdk.resources import Resource + from opentelemetry.sdk.trace import TracerProvider as SdkTracerProvider + from opentelemetry.sdk.trace.export import BatchSpanProcessor + except ImportError as error: + raise RuntimeError( + "ROUTER_OTLP_ENDPOINT is set but the tracing extra is not installed; " + 'install it with: pip install -e ".[tracing]"' + ) from error + + provider = SdkTracerProvider( + resource=Resource.create( + {"service.name": "llm-gateway", "deployment.environment": settings.environment} + ) + ) + provider.add_span_processor( + BatchSpanProcessor(OTLPSpanExporter(endpoint=settings.otlp_endpoint)) + ) + return provider diff --git a/tests/integration/test_tracing_api.py b/tests/integration/test_tracing_api.py new file mode 100644 index 0000000..881ae09 --- /dev/null +++ b/tests/integration/test_tracing_api.py @@ -0,0 +1,249 @@ +import hashlib +from collections.abc import AsyncIterator + +from fastapi.testclient import TestClient +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 + +from llm_router.app import create_app +from llm_router.backends import BackendResult, BackendUnavailableError, MockInferenceBackend +from llm_router.config import Settings +from llm_router.models import ChatCompletionRequest, PrivacyClass, RouteDecision +from llm_router.registry import load_registry +from llm_router.tracing import PROMPT_CONTENT_LIMIT, build_tracer_provider, prompt_attributes + +CATALOG = load_registry("config/registry.yaml") +SECRET = "patient Jane Doe, MRN 4471-2231, presents with chest pain" + + +class EchoingFailure(MockInferenceBackend): + """Fails with a message that repeats the prompt, as a real engine might.""" + + async def generate( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> BackendResult: + raise BackendUnavailableError(f"engine rejected: {request.prompt}") + + async def stream( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> AsyncIterator[str]: + raise BackendUnavailableError(f"engine rejected: {request.prompt}") + yield "" # pragma: no cover - unreachable, keeps this an async generator + + +def build( + *, backend: object | None = None, **overrides: object +) -> tuple[TestClient, InMemorySpanExporter]: + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + settings = Settings( + api_keys="trace-key", + tenant_keys="clinical-research:clinical-key", + **overrides, # type: ignore[arg-type] + ) + app = create_app( + settings, + registry=CATALOG, + tracer_provider=provider, + backend=backend, # type: ignore[arg-type] + ) + return TestClient(app, raise_server_exceptions=False), exporter + + +def post(client: TestClient, prompt: str, *, key: str = "trace-key", **body: object) -> object: + return client.post( + "/v1/chat/completions", + headers={"Authorization": f"Bearer {key}"}, + json={"model": "auto", "messages": [{"role": "user", "content": prompt}], **body}, + ) + + +def only_span(exporter: InMemorySpanExporter) -> ReadableSpan: + spans = exporter.get_finished_spans() + assert len(spans) == 1 + return spans[0] + + +def recorded_text(span: ReadableSpan) -> str: + """Everything a trace backend would store for the span, as one string.""" + + attributes = " ".join(f"{name}={value}" for name, value in (span.attributes or {}).items()) + events = " ".join(f"{event.name} {dict(event.attributes or {})}" for event in span.events) + return f"{span.name} {attributes} {events} {span.status.description}" + + +def test_restricted_prompts_contribute_their_length_only() -> None: + attributes = prompt_attributes(SECRET, PrivacyClass.RESTRICTED, record_content=True) + + assert attributes == {"router.prompt.chars": len(SECRET)} + + +def test_private_prompts_add_a_digest_but_never_content() -> None: + attributes = prompt_attributes(SECRET, PrivacyClass.PRIVATE, record_content=True) + + assert attributes == { + "router.prompt.chars": len(SECRET), + "router.prompt.sha256": hashlib.sha256(SECRET.encode()).hexdigest(), + } + + +def test_public_content_needs_an_operator_opt_in_and_is_bounded() -> None: + long_prompt = "x" * (PROMPT_CONTENT_LIMIT * 2) + + default = prompt_attributes(long_prompt, PrivacyClass.PUBLIC, record_content=False) + opted_in = prompt_attributes(long_prompt, PrivacyClass.PUBLIC, record_content=True) + + assert "router.prompt.content" not in default + assert opted_in["router.prompt.content"] == "x" * PROMPT_CONTENT_LIMIT + + +def test_a_request_span_carries_the_route_and_its_reason() -> None: + client, exporter = build() + with client: + response = post(client, "Extract the invoice fields", routing={"privacy": "private"}) + + span = only_span(exporter) + attributes = dict(span.attributes or {}) + assert response.status_code == 200 + assert attributes["router.tenant"] == "default" + assert attributes["router.privacy"] == "private" + assert attributes["router.task"] == "extraction" + assert attributes["router.cache"] == "miss" + assert attributes["gen_ai.response.model"] == "small-specialist" + assert attributes["router.model.revision"] == response.json()["routing"]["model_revision"] + assert attributes["router.route.reason"] == response.headers["X-Route-Reason"] + assert attributes["gen_ai.usage.output_tokens"] == response.json()["usage"]["completion_tokens"] + + +def test_the_response_quotes_the_trace_it_was_recorded_under() -> None: + client, exporter = build() + with client: + response = post(client, "Classify this ticket") + + span = only_span(exporter) + assert response.headers["X-Trace-Id"] == f"{span.get_span_context().trace_id:032x}" + + +def test_no_trace_header_is_sent_when_nothing_is_recording() -> None: + with TestClient(create_app(Settings(api_keys="trace-key"))) as client: + response = post(client, "Classify this ticket") + + assert response.status_code == 200 + assert "X-Trace-Id" not in response.headers + + +def test_private_prompt_text_never_reaches_the_trace() -> None: + client, exporter = build(trace_prompt_content=True) + with client: + post(client, SECRET, routing={"privacy": "private"}) + + text = recorded_text(only_span(exporter)) + assert "Jane Doe" not in text + assert "4471-2231" not in text + + +def test_a_tenant_floor_redacts_a_prompt_the_caller_declared_public() -> None: + client, exporter = build(trace_prompt_content=True) + with client: + # Content recording is on and the caller says public, but this tenant's + # floor is restricted, so the trace must be redacted under that class. + post(client, SECRET, key="clinical-key", routing={"privacy": "public"}) + + span = only_span(exporter) + attributes = dict(span.attributes or {}) + assert attributes["router.privacy"] == "restricted" + assert attributes["router.privacy.declared"] == "public" + assert "router.prompt.content" not in attributes + assert "router.prompt.sha256" not in attributes + assert "Jane Doe" not in recorded_text(span) + + +def test_an_engine_error_that_echoes_the_prompt_does_not_leak_it() -> None: + client, exporter = build(backend=EchoingFailure()) + with client: + response = post(client, SECRET, routing={"privacy": "private"}) + + span = only_span(exporter) + assert response.status_code == 502 + assert span.status.status_code is StatusCode.ERROR + assert dict(span.attributes or {})["error.type"] == "BackendUnavailableError" + assert "Jane Doe" not in recorded_text(span) + + +def test_a_quota_rejection_is_still_attributed_to_its_tenant() -> None: + client, exporter = build(quota_requests_per_minute=1) + with client: + post(client, "Classify this ticket") + rejected = post(client, "Classify this ticket again") + + spans = exporter.get_finished_spans() + assert rejected.status_code == 429 + assert len(spans) == 2 + assert dict(spans[1].attributes or {})["router.tenant"] == "default" + assert dict(spans[1].attributes or {})["error.type"] == "QuotaExceededError" + + +def test_a_cache_hit_is_traced_as_one() -> None: + client, exporter = build() + with client: + post(client, "Classify this ticket", routing={"privacy": "public"}) + post(client, "Classify this ticket", routing={"privacy": "public"}) + + first, second = exporter.get_finished_spans() + assert dict(first.attributes or {})["router.cache"] == "miss" + assert dict(second.attributes or {})["router.cache"] == "exact" + assert dict(second.attributes or {})["gen_ai.response.model"] == "small-specialist" + + +def test_a_streamed_request_ends_its_span_after_the_stream_with_usage() -> None: + client, exporter = build() + with client: + with client.stream( + "POST", + "/v1/chat/completions", + headers={"Authorization": "Bearer trace-key"}, + json={ + "model": "auto", + "messages": [{"role": "user", "content": "Summarize the report"}], + "stream": True, + }, + ) as response: + trace_id = response.headers["X-Trace-Id"] + list(response.iter_lines()) + + span = only_span(exporter) + assert trace_id == f"{span.get_span_context().trace_id:032x}" + assert dict(span.attributes or {})["router.stream"] is True + # Usage is only known once generation has finished, so its presence shows + # the span stayed open for the stream rather than ending with the handler. + assert dict(span.attributes or {})["gen_ai.usage.output_tokens"] >= 1 + + +def test_a_stream_that_fails_marks_its_span_failed() -> None: + client, exporter = build(backend=EchoingFailure()) + with client: + with client.stream( + "POST", + "/v1/chat/completions", + headers={"Authorization": "Bearer trace-key"}, + json={ + "model": "auto", + "messages": [{"role": "user", "content": SECRET}], + "stream": True, + }, + ) as response: + try: + list(response.iter_lines()) + except Exception: # the stream aborts mid-flight; that is the point + pass + + span = only_span(exporter) + assert span.status.status_code is StatusCode.ERROR + assert "Jane Doe" not in recorded_text(span) + + +def test_tracing_is_off_without_a_collector_endpoint() -> None: + assert build_tracer_provider(Settings(api_keys="trace-key")) is None