From 571635a43ed46a8721ab9e34d105cdf08558f35f Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 01:53:23 +0200 Subject: [PATCH] fix: align native traces and batch fixtures with upstream --- .../crates/python-bridge/src/routes/traces.rs | 4 +-- .../proxy/batches_endpoints/test_endpoints.py | 14 ++++++++ tests/test_litellm_rust/test_traces.py | 36 ++++++++++++++----- 3 files changed, 44 insertions(+), 10 deletions(-) diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 2e7a6b178a8..7944201fd9f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -120,13 +120,13 @@ impl NativeTraceStorage { fn query<'py>( &self, py: Python<'py>, - query: &str, + sql: &str, #[pyo3(from_py_with = litellm_host_python::from_py_argument)] parameters: BTreeMap< String, Parameter, >, ) -> PyResult> { - let query = ReadQuery::parse(query).map_err(map_error)?; + let query = ReadQuery::parse(sql).map_err(map_error)?; let connection = self.reader.clone().ok_or_else(|| { PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL") })?; diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 2d597abf3b8..8e3d25552e0 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -1038,6 +1038,20 @@ def _raw_batches_request(body: Dict[str, Any]) -> MagicMock: request.headers = {"Content-Type": "application/json"} request.client = MagicMock() request.client.host = "127.0.0.1" + request.scope = { + "type": "http", + "asgi": {"version": "3.0", "spec_version": "2.3"}, + "http_version": "1.1", + "method": "POST", + "scheme": "http", + "path": "/v1/batches", + "raw_path": b"/v1/batches", + "query_string": b"", + "root_path": "", + "headers": [(b"content-type", b"application/json"), (b"host", b"localhost")], + "client": ("127.0.0.1", 54321), + "server": ("localhost", 8000), + } request.body = AsyncMock(return_value=json.dumps(body).encode()) return request diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index 9e6e2e5b0cf..9d2348c9291 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -1,6 +1,7 @@ import base64 import gzip import json +import time from typing import Final from urllib.parse import parse_qs, urlsplit @@ -17,11 +18,14 @@ async def test_trace_reader_projects_connection_and_parameters(recording_server: recording_server.enqueue(ResponseSpec(body={"data": [{"trace_id": "trace-1"}]})) reader_url: Final = recording_server.base_url.replace("http://", "http://reader:p%40ss%2Fword%25@") storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, reader_url + "?database=wrong") - response: Final = json.loads(await storage.query("SELECT {trace_id:String} AS trace_id", {"trace_id": "trace-1"})) + response: Final = json.loads( + await storage.query("trace_spans", {"trace_id": "trace-1", "team_ids": [], "api_key_hash": "", "trace_ref": ""}) + ) request: Final = recording_server.requests[0] parameters: Final = parse_qs(urlsplit(request.path).query) assert response["data"] == [{"trace_id": "trace-1"}] - assert request.raw_body == b"SELECT {trace_id:String} AS trace_id" + assert b"o.TraceId = {trace_id:String}" in request.raw_body + assert b"trace-1" not in request.raw_body assert parameters["database"] == ["trace_test"] assert parameters["param_trace_id"] == ["trace-1"] assert parameters["readonly"] == ["1"] @@ -35,6 +39,13 @@ async def test_trace_reader_rejects_success_status_with_embedded_error(recording recording_server.enqueue(ResponseSpec(body={"data": [], "exception": "query failed"})) storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, recording_server.base_url) with pytest.raises(RuntimeError, match="invalid or failed JSON"): + await storage.query("trace_spans", {"trace_id": "trace-1", "team_ids": [], "api_key_hash": "", "trace_ref": ""}) + + +@pytest.mark.asyncio +async def test_trace_reader_rejects_arbitrary_sql() -> None: + storage: Final = NativeTraceStorage("trace_test", "http://localhost:8123", "http://localhost:8123") + with pytest.raises(ValueError, match="unknown ClickHouse read query"): await storage.query("SELECT 1", {}) @@ -52,7 +63,9 @@ async def test_schema_binding_rejects_non_positive_retention() -> None: @pytest.mark.asyncio -async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement(recording_server: RecordingServer) -> None: +async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement( + recording_server: RecordingServer, +) -> None: recording_server.expected_requests = 2 recording_server.enqueue(ResponseSpec(body="")) recording_server.enqueue(ResponseSpec(status=403, body="denied")) @@ -64,20 +77,27 @@ async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement assert recording_server.requests[0].raw_body.startswith(b"CREATE DATABASE IF NOT EXISTS") assert recording_server.requests[1].raw_body.startswith(b"CREATE TABLE IF NOT EXISTS") assert "readonly" not in parse_qs(urlsplit(recording_server.requests[0].path).query) - assert recording_server.requests[0].headers["authorization"] == "Basic " + base64.b64encode( - b"writer:p@ss/word%" - ).decode() + assert ( + recording_server.requests[0].headers["authorization"] + == "Basic " + base64.b64encode(b"writer:p@ss/word%").decode() + ) @pytest.mark.asyncio async def test_insert_encodes_and_sends_rows(recording_server: RecordingServer) -> None: recording_server.enqueue(ResponseSpec(body="")) storage: Final = NativeTraceStorage("trace_test", recording_server.base_url) + started_ms: Final = time.time_ns() // 1_000_000 await storage.insert_rows("otel_traces", [{"Timestamp": 1_234_567_890, "Input": "hello"}]) + finished_ms: Final = time.time_ns() // 1_000_000 request: Final = recording_server.requests[0] - assert json.loads(gzip.decompress(request.raw_body)) == { + row: Final = json.loads(gzip.decompress(request.raw_body)) + assert started_ms <= row["EngineReceivedMs"] <= finished_ms + assert {key: value for key, value in row.items() if key != "EngineReceivedMs"} == { "Input": "hello", "Timestamp": "1970-01-01T00:00:01.23456789Z", } - assert parse_qs(urlsplit(request.path).query)["query"] == ["INSERT INTO `trace_test`.otel_traces FORMAT JSONEachRow"] + assert parse_qs(urlsplit(request.path).query)["query"] == [ + "INSERT INTO `trace_test`.otel_traces FORMAT JSONEachRow" + ] assert request.headers["content-encoding"] == "gzip"