From 5d2e8b76e40caef6ee56c768e31c3a524dd8cd22 Mon Sep 17 00:00:00 2001 From: kerry Date: Tue, 22 Sep 2026 20:44:28 +0000 Subject: [PATCH] test(integration): keep Databricks streaming usage and cache reads in the client stream and spend log (Pylon #7094) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/contracts.json | 3 + .../providers/test_databricks_chat_wire.py | 129 ++++++++++++++++++ 2 files changed, 132 insertions(+) create mode 100644 tests/integration/providers/test_databricks_chat_wire.py diff --git a/tests/integration/contracts.json b/tests/integration/contracts.json index e49bba7e90c..968018efdd5 100644 --- a/tests/integration/contracts.json +++ b/tests/integration/contracts.json @@ -241,6 +241,9 @@ "tests/integration/providers/test_xiaomi_mimo_wire.py::test_xiaomi_mimo_tool_call_is_forwarded_and_returned": [ "other.provider_wire.xiaomi_mimo.tool_call_survives_translation" ], + "tests/integration/providers/test_databricks_chat_wire.py::test_databricks_stream_final_usage_chunk_reaches_client_and_spend_log": [ + "other.provider_wire.databricks.stream_usage_and_cache_reads_reach_client_and_spend_log" + ], "tests/integration/providers/test_fal_ai_video_wire.py::test_fal_h3_video_create_uses_canonical_body_and_status_path": [ "other.provider_wire.fal_ai.video_queue_create_status_and_content_download" ], diff --git a/tests/integration/providers/test_databricks_chat_wire.py b/tests/integration/providers/test_databricks_chat_wire.py new file mode 100644 index 00000000000..614382a77f0 --- /dev/null +++ b/tests/integration/providers/test_databricks_chat_wire.py @@ -0,0 +1,129 @@ +import json +import uuid +from collections.abc import Mapping +from typing import Final + +import pytest +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, wire_server +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter + +_BACKEND: Final = "databricks-glm-5-2" +_API_KEY: Final = "synthetic-databricks-key" +_PROMPT: Final = "Summarise the cached briefing in one sentence." +_PROVIDER_USAGE: Final[Mapping[str, JsonValue]] = { + "prompt_tokens": 12011, + "completion_tokens": 8, + "total_tokens": 12019, + "cache_read_input_tokens": 12002, + "cache_creation_input_tokens": 0, +} +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +class _PromptTokensDetails(BaseModel): + model_config = ConfigDict(extra="ignore") + cached_tokens: int | None = None + + +class _Usage(BaseModel): + model_config = ConfigDict(extra="ignore") + prompt_tokens: int + completion_tokens: int + total_tokens: int + prompt_tokens_details: _PromptTokensDetails | None = None + + +class _Delta(BaseModel): + model_config = ConfigDict(extra="ignore") + content: str | None = None + + +class _Choice(BaseModel): + model_config = ConfigDict(extra="ignore") + delta: _Delta + + +class _Chunk(BaseModel): + model_config = ConfigDict(extra="ignore") + id: str + choices: tuple[_Choice, ...] + usage: _Usage | None = None + + +def _frame(identity: str, choices: list[Mapping[str, object]], usage: Mapping[str, JsonValue] | None = None) -> bytes: + value: Final = { + "id": identity, + "object": "chat.completion.chunk", + "created": 1, + "model": _BACKEND, + "choices": choices, + **({} if usage is None else {"usage": usage}), + } + return b"data: " + json.dumps(value).encode() + b"\n\n" + + +@pytest.mark.covers("other.provider_wire.databricks.stream_usage_and_cache_reads_reach_client_and_spend_log") +def test_databricks_stream_final_usage_chunk_reaches_client_and_spend_log(gateway: Gateway) -> None: + identity: Final = f"databricks-stream-{uuid.uuid4().hex}" + frames: Final = ( + _frame( + identity, [{"index": 0, "delta": {"role": "assistant", "content": "The briefing "}, "finish_reason": None}] + ), + _frame(identity, [{"index": 0, "delta": {"content": "is short."}, "finish_reason": None}]), + _frame(identity, [{"index": 0, "delta": {}, "finish_reason": "stop"}]), + _frame(identity, [], usage=_PROVIDER_USAGE), + b"data: [DONE]\n\n", + ) + + def respond(request: Request) -> Reply: + assert request.method == "POST" + assert request.target == "/chat/completions" + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["model"] == _BACKEND + assert body["messages"] == [{"role": "user", "content": _PROMPT}] + assert body["stream"] is True + return Reply(content_type="text/event-stream", chunks=frames) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"databricks/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + with gateway.client.stream( + "POST", + "/v1/chat/completions", + json={ + "model": model, + "messages": [{"role": "user", "content": _PROMPT}], + "stream": True, + "stream_options": {"include_usage": True}, + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + ) as response: + assert response.status_code == 200, response.read() + lines: Final = tuple(line for line in response.iter_lines() if line.startswith("data: ")) + assert lines[-1] == "data: [DONE]", lines + chunks: Final = tuple(_Chunk.model_validate_json(line.removeprefix("data: ")) for line in lines[:-1]) + assert {chunk.id for chunk in chunks} == {identity} + assert ( + "".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices) + == "The briefing is short." + ) + usages: Final = tuple(chunk.usage for chunk in chunks if chunk.usage is not None) + assert len(usages) == 1, lines + assert ( + usages[0].prompt_tokens, + usages[0].completion_tokens, + usages[0].total_tokens, + usages[0].prompt_tokens_details.cached_tokens if usages[0].prompt_tokens_details is not None else None, + ) == (12011, 8, 12019, 12002), lines + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")] + rows: Final = eventually( + lambda: read_rows( + 'SELECT prompt_tokens, completion_tokens, total_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (identity,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert (rows[0]["prompt_tokens"], rows[0]["completion_tokens"], rows[0]["total_tokens"]) == (12011, 8, 12019)