From ca05eca2d332ae925448f3392bd3146d6fca3068 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Thu, 1 Oct 2026 12:29:09 -0700 Subject: [PATCH 001/203] feat(vertex-ai): add vertex_ai/xai/grok-4.7 pricing (#44059) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 20 +++++++++++++++++++ model_prices_and_context_window.json | 20 +++++++++++++++++++ 2 files changed, 40 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 44b5cb0f59f..3392aa4d868 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -79427,5 +79427,25 @@ "supported_endpoints": [ "/v1/audio/speech" ] + }, + "vertex_ai/xai/grok-4.7": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_200k_tokens": 1e-06, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, + "litellm_provider": "vertex_ai", + "max_input_tokens": 524288, + "max_output_tokens": 524288, + "max_tokens": 524288, + "mode": "chat", + "output_cost_per_token": 6e-06, + "output_cost_per_token_above_200k_tokens": 1.2e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true } } diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 44b5cb0f59f..3392aa4d868 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -79427,5 +79427,25 @@ "supported_endpoints": [ "/v1/audio/speech" ] + }, + "vertex_ai/xai/grok-4.7": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_200k_tokens": 1e-06, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, + "litellm_provider": "vertex_ai", + "max_input_tokens": 524288, + "max_output_tokens": 524288, + "max_tokens": 524288, + "mode": "chat", + "output_cost_per_token": 6e-06, + "output_cost_per_token_above_200k_tokens": 1.2e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true } } From 5e5882244adb81c44cf440d4652f24604abd1291 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 1 Oct 2026 12:34:01 -0700 Subject: [PATCH 002/203] feat(ui): drop the Beta badge from the Cost Optimization nav item (#43967) Co-authored-by: Krrish Dholakia Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ui/litellm-dashboard/src/components/leftnav.test.tsx | 6 +++--- ui/litellm-dashboard/src/components/leftnav.tsx | 6 +----- 2 files changed, 4 insertions(+), 8 deletions(-) diff --git a/ui/litellm-dashboard/src/components/leftnav.test.tsx b/ui/litellm-dashboard/src/components/leftnav.test.tsx index df94772400e..10c5a2fc7a7 100644 --- a/ui/litellm-dashboard/src/components/leftnav.test.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.test.tsx @@ -607,13 +607,13 @@ describe("Sidebar (leftnav)", () => { expect(label).toHaveClass("group-data-[collapsed=true]/sidebar:hidden"); }); - it("shows Cost Optimization with a Beta badge and no feature-flag gate", () => { + it("shows Cost Optimization without a Beta badge and no feature-flag gate", () => { const { container } = renderWithProviders(); const costOptimization = container.querySelector('a[href*="cost-optimization"]'); expect(costOptimization).not.toBeNull(); - expect(costOptimization!).toHaveTextContent(/Cost Optimization/); - expect(costOptimization!).toHaveTextContent(/Beta/); + expect(costOptimization!).toHaveTextContent("Cost Optimization"); + expect(costOptimization!).not.toHaveTextContent("Beta"); expect(container.querySelector('a[href*="projects"]')).toBeNull(); }); diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index 43e2ae17c24..2679478a414 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -240,11 +240,7 @@ const menuGroups: MenuGroup[] = [ page: "cost-optimization", icon: , roles: [...all_admin_roles, ...internalUserRoles], - label: ( - - Cost Optimization - - ), + label: "Cost Optimization", }, { key: "logs", page: "logs", label: "Logs", icon: }, { From 0da00d4b2ee756aff75d24c5f4aed0c93ff5b1f4 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Thu, 1 Oct 2026 12:54:26 -0700 Subject: [PATCH 003/203] test(e2e): bill Sail windows that synchronous calls can still use (#44058) Sail now rejects completion_window "flex" on synchronous requests with a 400 saying flex is only for background responses or Batch work. The chat flex case and the responses flex case have failed on every scheduled litellm-e2e run in builds 337, 340 and 341. The chat cases keep balanced and auto, and the responses case sends a caller metadata.completion_window of balanced, so both still prove the window reaches Sail and the bill uses that window's distinct rates --- tests/e2e/coverage_registry/llm_conversational.yaml | 4 ++-- tests/e2e/llm_translation/test_sail_e2e.py | 10 ++++++---- 2 files changed, 8 insertions(+), 6 deletions(-) diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index 61f3be34a43..919884b66f0 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -101,10 +101,10 @@ - {id: llm.messages.together_ai.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together over /v1/messages streaming"} - {id: llm.messages.together_ai.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together tool calls over /v1/messages"} - {id: llm.messages.together_ai.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together tool result round trip over /v1/messages"} -- {id: llm.chat_completions.sail.service_tier.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: chat_completions, route: sail, capability: service_tier, streaming: nonstream, assertions: [cost_logged], source: "llm_translation/test_sail_e2e.py", rationale: "service_tier flex, balanced and auto map to Sail completion windows and bill the matching price columns"} +- {id: llm.chat_completions.sail.service_tier.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: chat_completions, route: sail, capability: service_tier, streaming: nonstream, assertions: [cost_logged], source: "llm_translation/test_sail_e2e.py", rationale: "service_tier balanced and auto map to Sail completion windows and bill the matching price columns; Sail serves flex only to background responses and Batch"} - {id: llm.chat_completions.sail.service_tier.nonstream.rejects_unknown_tier, module: llm, tier: P1, subject_endpoint: chat_completions, route: sail, capability: service_tier, streaming: nonstream, assertions: [rejects_unknown_tier], source: "llm_translation/test_sail_e2e.py", rationale: "A service_tier Sail has no completion window for is a 400 without drop_params"} - {id: llm.chat_completions.sail.service_tier.nonstream.drops_unknown_tier_and_bills_asap, module: llm, tier: P1, subject_endpoint: chat_completions, route: sail, capability: service_tier, streaming: nonstream, assertions: [drops_unknown_tier_and_bills_asap], source: "llm_translation/test_sail_e2e.py", rationale: "An unknown service_tier under drop_params is dropped and billed at asap in both the cost header and spend log"} -- {id: llm.responses.sail.service_tier.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: responses, route: sail, capability: service_tier, streaming: nonstream, assertions: [cost_logged], source: "llm_translation/test_sail_e2e.py", rationale: "A caller metadata.completion_window of flex on /v1/responses bills Sail flex rates"} +- {id: llm.responses.sail.service_tier.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: responses, route: sail, capability: service_tier, streaming: nonstream, assertions: [cost_logged], source: "llm_translation/test_sail_e2e.py", rationale: "A caller metadata.completion_window of balanced on /v1/responses bills Sail balanced rates"} - {id: llm.messages.sail.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: sail, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_sail_e2e.py", rationale: "Sail over /v1/messages"} - {id: llm.chat_completions.anthropic.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic over /chat/completions: cost header and spend row agree"} - {id: llm.chat_completions.anthropic.multi_turn.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic tool result round trip over /chat/completions"} diff --git a/tests/e2e/llm_translation/test_sail_e2e.py b/tests/e2e/llm_translation/test_sail_e2e.py index cf662afea90..7267052e12c 100644 --- a/tests/e2e/llm_translation/test_sail_e2e.py +++ b/tests/e2e/llm_translation/test_sail_e2e.py @@ -117,7 +117,7 @@ def _assert_spend_row_matches(proxy: ProxyClient, key: str, header_cost: float) class TestSailChatCompletions: @pytest.mark.covers("llm.chat_completions.sail.service_tier.nonstream.cost_logged") @pytest.mark.parametrize( - ("service_tier", "billed_tier"), [("flex", "flex"), ("balanced", "balanced"), ("auto", "base")] + ("service_tier", "billed_tier"), [("balanced", "balanced"), ("auto", "base")] ) def test_service_tier_bills_the_matching_completion_window( self, @@ -176,7 +176,7 @@ class TestSailChatCompletions: class TestSailResponses: @pytest.mark.covers("llm.responses.sail.service_tier.nonstream.cost_logged") - def test_flex_completion_window_bills_flex_rates( + def test_caller_completion_window_bills_its_rates( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: model, key = _register(proxy, resources) @@ -185,7 +185,7 @@ class TestSailResponses: model=model, input=f"{PROMPT} {unique_marker()}", max_output_tokens=MAX_TOKENS, - metadata={"completion_window": "flex"}, + metadata={"completion_window": "balanced"}, extra_body=NO_PROXY_CACHE, ) usage: Final = raw.parse().usage @@ -196,7 +196,9 @@ class TestSailResponses: completion=usage.output_tokens, ) - header_cost: Final = _assert_billed_at("flex", tokens, response_header(raw.headers, "x-litellm-response-cost")) + header_cost: Final = _assert_billed_at( + "balanced", tokens, response_header(raw.headers, "x-litellm-response-cost") + ) _assert_spend_row_matches(proxy, key, header_cost) From 65a316fc9209c3d75a167c8b6ef5b9f4bab030ba Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Thu, 1 Oct 2026 13:00:54 -0700 Subject: [PATCH 004/203] ci(circleci): test Redis behavior against local Redis and print short tracebacks (#44062) * ci(circleci): test Redis behavior against local Redis and print short tracebacks The redis_caching_unit_tests job ran three legacy files against the shared remote Redis. The DualCache and batch-read logic that never needed a server now lives in tests/unit/caching/test_dual_cache.py with a mocked RedisCache, and the behavior that does need one (the increment-with-floor Lua script, read-through, deletes, batch reads) moved to tests/integration, which starts a local Redis. test_returned_settings only read REDIS_PORT and is replaced by a unit test of Router.get_settings CircleCI pytest runs now use --tb=short so failure output stays readable in the test results tab * ci(integration): print short tracebacks from run.py and allow Redis in the sdk shard --- .circleci/config.yml | 133 +++------ .circleci/scripts/run_integration.sh | 2 +- tests/integration/README.md | 2 +- .../test_redis_increment_with_floor.py | 29 +- tests/integration/run.py | 1 + .../integration/sdk/test_dual_cache_redis.py | 94 ++++++ tests/local_testing/test_dual_cache.py | 274 ------------------ .../test_redis_batch_optimizations.py | 123 -------- tests/local_testing/test_router_utils.py | 67 ----- tests/unit/caching/test_dual_cache.py | 97 +++++++ tests/unit/test_router_get_settings.py | 26 ++ 11 files changed, 267 insertions(+), 581 deletions(-) rename tests/{local_testing => integration/routing}/test_redis_increment_with_floor.py (65%) create mode 100644 tests/integration/sdk/test_dual_cache_redis.py delete mode 100644 tests/local_testing/test_dual_cache.py delete mode 100644 tests/local_testing/test_redis_batch_optimizations.py create mode 100644 tests/unit/test_router_get_settings.py diff --git a/.circleci/config.yml b/.circleci/config.yml index afed7853ac6..1798abe9de5 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -408,7 +408,7 @@ jobs: - run: name: Run Windows-specific test command: | - uv run --no-sync python -m pytest tests/windows_tests/ -v + uv run --no-sync python -m pytest --tb=short tests/windows_tests/ -v windows_release_wheel: executor: @@ -551,7 +551,7 @@ jobs: echo "$TEST_FILES" | circleci tests run \ --split-by=timings \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --cov=./litellm --cov=./enterprise/litellm_enterprise \ --cov-report=xml \ @@ -625,7 +625,7 @@ jobs: echo "$TEST_FILES" | circleci tests run \ --split-by=timings \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --cov=./litellm --cov=./enterprise/litellm_enterprise \ --cov-report=xml \ @@ -697,7 +697,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/local_testing/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --junitxml=test-results/junit.xml \ --durations=5 \ @@ -752,7 +752,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/proxy_admin_ui_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -815,7 +815,7 @@ jobs: echo "$TEST_FILES" | circleci tests run \ --split-by=timings \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ -k 'router' \ -n 4 \ @@ -859,7 +859,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/router_unit_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -904,7 +904,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/local_testing/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --junitxml=test-results/junit.xml \ --durations=5 \ @@ -948,7 +948,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/llm_translation/**/test_*.py" | grep -v "^tests/llm_translation/realtime/") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --junitxml=test-results/junit.xml \ --durations=20 \ @@ -986,7 +986,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/llm_translation/realtime/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -1031,7 +1031,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/agent_tests/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv -s \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -1075,7 +1075,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/guardrails_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -1121,7 +1121,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/unified_google_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv -s \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -1176,7 +1176,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/llm_responses_api_testing/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --junitxml=test-results/junit.xml \ --durations=5 \ @@ -1210,7 +1210,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/ocr_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -1254,7 +1254,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/search_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -1298,7 +1298,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/batches_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv -s \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -1342,7 +1342,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/litellm_utils_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv -s \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -1387,7 +1387,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/pass_through_unit_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -1432,7 +1432,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/image_gen_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --junitxml=test-results/junit.xml \ --durations=5 \ @@ -1466,7 +1466,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/logging_callback_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ -n 4 \ @@ -1511,7 +1511,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/audio_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv -s \ --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ --junitxml=test-results/junit.xml \ @@ -1531,61 +1531,6 @@ jobs: paths: - audio_coverage.xml - audio_coverage - redis_caching_unit_tests: - docker: - - *python312_image - working_directory: ~/project - - steps: - - checkout - - skip_if_unrelated_changes - - setup_google_dns - - restore_cache: - keys: - - v1-uv-cache-{{ checksum "uv.lock" }} - - install_uv - - install_rust - - run: - name: Install Dependencies - command: | - uv sync --frozen --all-groups --all-extras --python 3.12 - - save_cache: - paths: - - ~/.cache/uv - key: v1-uv-cache-{{ checksum "uv.lock" }} - # Run pytest and generate JUnit XML report - - run: - name: Run tests - command: | - mkdir -p test-results - TEST_FILES=$(printf "%s\n" \ - tests/local_testing/test_dual_cache.py \ - tests/local_testing/test_redis_batch_optimizations.py \ - tests/local_testing/test_redis_increment_with_floor.py \ - tests/local_testing/test_router_utils.py) - echo "$TEST_FILES" | circleci tests run \ - --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ - -vv -s \ - --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \ - --junitxml=test-results/junit.xml \ - --durations=5 -n 2 \ - --reruns 2 --reruns-delay 1" - no_output_timeout: 20m - - run: - name: Rename the coverage files - command: | - mv coverage.xml redis_caching_coverage.xml - mv .coverage redis_caching_coverage - - # Store test results - - store_test_results: - path: test-results - - persist_to_workspace: - root: . - paths: - - redis_caching_coverage.xml - - redis_caching_coverage installing_litellm_on_python: docker: - *python312_image @@ -1605,7 +1550,7 @@ jobs: - run: name: Run tests command: | - uv run --no-sync python -m pytest -vv tests/local_testing/test_basic_python_version.py -k "not legacy_resolver" + uv run --no-sync python -m pytest --tb=short -vv tests/local_testing/test_basic_python_version.py -k "not legacy_resolver" installing_litellm_on_python_3_13: docker: @@ -1629,7 +1574,7 @@ jobs: - run: name: Run tests command: | - uv run --no-sync python -m pytest -v tests/local_testing/test_basic_python_version.py -k "not legacy_resolver" + uv run --no-sync python -m pytest --tb=short -v tests/local_testing/test_basic_python_version.py -k "not legacy_resolver" installing_litellm_on_python_v2_migration_resolver: docker: @@ -1660,7 +1605,7 @@ jobs: - run: name: Run both migration resolvers against Postgres command: | - uv run --no-sync python -m pytest -vv \ + uv run --no-sync python -m pytest --tb=short -vv \ tests/local_testing/test_basic_python_version.py::test_litellm_proxy_server_config_no_general_settings \ tests/local_testing/test_basic_python_version.py::test_litellm_proxy_server_config_no_general_settings_legacy_resolver @@ -1829,7 +1774,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/basic_proxy_startup_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --junitxml=test-results/junit-2.xml \ --durations=5" @@ -1926,7 +1871,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -s -v \ --junitxml=test-results/junit.xml \ -n 4 \ @@ -2013,7 +1958,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/openai_endpoints_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -s -vv \ --junitxml=test-results/junit.xml \ --durations=5" @@ -2096,7 +2041,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/otel_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --junitxml=test-results/junit.xml \ --durations=5" @@ -2148,7 +2093,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/basic_proxy_startup_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --junitxml=test-results/junit-2.xml \ --durations=5" @@ -2229,7 +2174,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/spend_tracking_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --junitxml=test-results/junit.xml \ --durations=5" @@ -2334,7 +2279,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/multi_instance_e2e_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --junitxml=test-results/junit.xml \ --durations=5" @@ -2406,7 +2351,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/store_model_in_db_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --junitxml=test-results/junit.xml \ --durations=5" @@ -2491,7 +2436,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/basic_proxy_startup_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv \ --junitxml=test-results/junit-2.xml \ --durations=5" @@ -2588,7 +2533,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/pass_through_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -v \ --junitxml=test-results/junit.xml \ --durations=5" @@ -2659,7 +2604,7 @@ jobs: TEST_FILES=$(circleci tests glob "tests/proxy_e2e_anthropic_messages_tests/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \ + --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ -vv -s \ --junitxml=test-results/junit.xml \ --durations=5" @@ -2689,7 +2634,7 @@ jobs: - run: name: Combine Coverage command: | - uv tool run --from 'coverage[toml]==7.10.6' coverage combine realtime_translation_coverage ocr_coverage search_coverage logging_coverage audio_coverage local_testing_part1_coverage local_testing_part2_coverage pass_through_unit_tests_coverage batches_coverage guardrails_coverage redis_caching_coverage agent_coverage google_generate_content_endpoint_coverage litellm_utils_coverage router_unit_tests_coverage auth_ui_unit_tests_coverage + uv tool run --from 'coverage[toml]==7.10.6' coverage combine realtime_translation_coverage ocr_coverage search_coverage logging_coverage audio_coverage local_testing_part1_coverage local_testing_part2_coverage pass_through_unit_tests_coverage batches_coverage guardrails_coverage agent_coverage google_generate_content_endpoint_coverage litellm_utils_coverage router_unit_tests_coverage auth_ui_unit_tests_coverage uv tool run --from 'coverage[toml]==7.10.6' coverage xml - codecov/upload: file: ./coverage.xml @@ -3189,7 +3134,7 @@ jobs: name: Test provider capture and replay harness command: | mkdir -p test-results/provider-replay-harness - uv run --no-sync pytest -q --noconftest -o addopts= -o pythonpath=tests/e2e -p no:rerunfailures \ + uv run --no-sync pytest --tb=short -q --noconftest -o addopts= -o pythonpath=tests/e2e -p no:rerunfailures \ --junitxml=test-results/provider-replay-harness/junit.xml \ tests/e2e/test_provider_edge.py tests/e2e/test_fixture_bundle.py \ tests/e2e/test_fixture_canonical.py tests/e2e/test_fixture_mode.py \ @@ -3492,7 +3437,6 @@ workflows: - image_gen_testing - logging_testing - audio_testing - - redis_caching_unit_tests - upload-coverage: requires: - realtime_translation_testing @@ -3507,7 +3451,6 @@ workflows: - image_gen_testing - logging_testing - audio_testing - - redis_caching_unit_tests - langfuse_logging_unit_tests - local_testing_part1 - local_testing_part2 diff --git a/.circleci/scripts/run_integration.sh b/.circleci/scripts/run_integration.sh index 26240487c48..ba24e66ba1c 100644 --- a/.circleci/scripts/run_integration.sh +++ b/.circleci/scripts/run_integration.sh @@ -191,7 +191,7 @@ if [ "$suite" = management ] || [ "$suite" = mcp ]; then fi if [ "$suite" = providers ]; then - INTEGRATION_RUN_ID="$integration_identity" .venv/bin/python -m pytest --noconftest -o addopts= \ + INTEGRATION_RUN_ID="$integration_identity" .venv/bin/python -m pytest --tb=short --noconftest -o addopts= \ --strict-markers --strict-config -p no:pytest-retry -p no:rerunfailures --timeout=30 \ tests/e2e/test_provider_edge.py::TestReplayMode::test_content_drift_returns_the_miss_status_naming_both_keys \ tests/e2e/test_provider_edge.py::TestReplayMode::test_exhausted_key_returns_the_miss_status \ diff --git a/tests/integration/README.md b/tests/integration/README.md index c559e7545e0..2204cde11e3 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -30,7 +30,7 @@ Streaming checks send real HTTP transfer chunks, including one-byte partitions, The `messages_endpoint/` directory holds `/v1/messages` endpoint contracts: native-provider backends under `providers/` (`anthropic`, `bedrock`, `gemini`) and the translation bridges (`responses_bridge`, `chat_bridge`) at the top level. It runs in the providers shard; `run.py` selects test files recursively under each scheduled directory -The sdk shard exercises the SDK's own HTTP clients against local protocol peers with no gateway in the path, so a case here fails only when the client library or its wire behavior changes. The HTTP/2 case runs a hypercorn TLS peer offering h2 and http/1.1 over ALPN, drives the sync and async httpx handlers at it with `LITELLM_HTTP2` off and on, and asserts the version both the client and the peer observed on the wire. Put a test here only when it needs no proxy, database or Redis; a case that reaches the gateway belongs in one of the other shards +The sdk shard exercises the SDK's own HTTP clients against local protocol peers with no gateway in the path, so a case here fails only when the client library or its wire behavior changes. The HTTP/2 case runs a hypercorn TLS peer offering h2 and http/1.1 over ALPN, drives the sync and async httpx handlers at it with `LITELLM_HTTP2` off and on, and asserts the version both the client and the peer observed on the wire. Put a test here only when it needs no proxy or database. CircleCI starts a local Redis for this shard like the others, so SDK-side caching cases that need a real Redis server belong here too; a case that reaches the gateway belongs in one of the other shards The extensions shard uses the built-in generic callback and guardrail transports. It checks callback correlation and credential exclusion, guardrail rewriting and denial, retained OpenAI consumers and A2A wire versions. CircleCI runs it on parallel nodes, and each node starts its own database, Redis, upstream and proxy and runs its share of the group's files serially, split by recorded timings with `circleci tests split`. Tests keep the isolation of a serial run; they still must not assume a particular set of sibling files. `run.py --list` prints a group's files and `run.py ...` runs a subset of them diff --git a/tests/local_testing/test_redis_increment_with_floor.py b/tests/integration/routing/test_redis_increment_with_floor.py similarity index 65% rename from tests/local_testing/test_redis_increment_with_floor.py rename to tests/integration/routing/test_redis_increment_with_floor.py index e358d5f31e0..d5535d30715 100644 --- a/tests/local_testing/test_redis_increment_with_floor.py +++ b/tests/integration/routing/test_redis_increment_with_floor.py @@ -1,15 +1,9 @@ -"""Least-busy routing keeps its in-flight counters in Redis, and the clamp at zero plus the -create-once TTL both live inside a Lua script. Nothing but a real Redis runs that script, so -these are the only tests that fail when the script itself is wrong.""" - import os import uuid +from collections.abc import Iterator from typing import Final import pytest -from dotenv import load_dotenv - -load_dotenv() from litellm.caching.redis_cache import RedisCache @@ -17,14 +11,14 @@ TTL: Final = 600 @pytest.fixture -def counter(): - cache: Final = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - key: Final = f"lit7039-{uuid.uuid4()}" +def counter() -> Iterator[tuple[RedisCache, str, str]]: + cache: Final = RedisCache(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) + key: Final = f"increment-with-floor-{uuid.uuid4()}" yield cache, key, cache.check_and_fix_namespace(key=key) cache.delete_cache(key) -def test_a_counter_adds_every_increment_and_reads_back_what_it_holds(counter): +def test_a_counter_adds_every_increment_and_reads_back_what_it_holds(counter: tuple[RedisCache, str, str]) -> None: cache, key, _ = counter assert cache.increment_with_floor(key, 3, TTL) == 3 @@ -32,10 +26,7 @@ def test_a_counter_adds_every_increment_and_reads_back_what_it_holds(counter): assert cache.batch_get_counts([key]) == (5,) -def test_a_decrement_past_zero_leaves_the_counter_at_zero(counter): - """A worker whose counter expired mid-request decrements a key that is no longer there. - Without the clamp that deployment reads negative, and least-busy pins every later request - on it until the count climbs back to zero.""" +def test_a_decrement_past_zero_leaves_the_counter_at_zero(counter: tuple[RedisCache, str, str]) -> None: cache, key, _ = counter assert cache.increment_with_floor(key, 1, TTL) == 1 @@ -43,9 +34,7 @@ def test_a_decrement_past_zero_leaves_the_counter_at_zero(counter): assert cache.batch_get_counts([key]) == (0,) -def test_traffic_never_pushes_a_counters_expiry_back_out(counter): - """The TTL is what releases a count whose worker died mid-request. Rewriting it on every - touch would keep that stuck count alive for as long as the group takes traffic.""" +def test_traffic_never_pushes_a_counters_expiry_back_out(counter: tuple[RedisCache, str, str]) -> None: cache, key, namespaced_key = counter cache.increment_with_floor(key, 1, TTL) @@ -57,7 +46,7 @@ def test_traffic_never_pushes_a_counters_expiry_back_out(counter): assert cache.redis_client.ttl(namespaced_key) <= 30 -def test_clamping_to_zero_keeps_the_expiry_it_already_had(counter): +def test_clamping_to_zero_keeps_the_expiry_it_already_had(counter: tuple[RedisCache, str, str]) -> None: cache, key, namespaced_key = counter cache.increment_with_floor(key, 1, TTL) @@ -68,7 +57,7 @@ def test_clamping_to_zero_keeps_the_expiry_it_already_had(counter): @pytest.mark.asyncio -async def test_the_async_counter_behaves_the_same_way(counter): +async def test_the_async_counter_behaves_the_same_way(counter: tuple[RedisCache, str, str]) -> None: cache, key, namespaced_key = counter assert await cache.async_increment_with_floor(key, 2, TTL) == 2 diff --git a/tests/integration/run.py b/tests/integration/run.py index 19bce35f542..30c1352f048 100644 --- a/tests/integration/run.py +++ b/tests/integration/run.py @@ -72,6 +72,7 @@ def main() -> int: "no:rerunfailures", "--timeout=90", "--durations=15", + "--tb=short", f"--hypothesis-seed={options.seed}", f"--integration-order-seed={options.order_seed}", f"--junitxml={output / 'junit.xml'}", diff --git a/tests/integration/sdk/test_dual_cache_redis.py b/tests/integration/sdk/test_dual_cache_redis.py new file mode 100644 index 00000000000..ddd01480d36 --- /dev/null +++ b/tests/integration/sdk/test_dual_cache_redis.py @@ -0,0 +1,94 @@ +import asyncio +import os +import uuid +from typing import Final +from unittest.mock import patch + +import pytest + +from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.caching.redis_cache import RedisCache + + +def _redis_cache() -> RedisCache: + return RedisCache(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) + + +@pytest.mark.asyncio +async def test_a_value_only_in_redis_is_read_once_from_redis_then_from_memory() -> None: + redis_cache: Final = _redis_cache() + dual_cache: Final = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis_cache) + sync_key: Final = f"redis-only-sync-{uuid.uuid4()}" + async_key: Final = f"redis-only-async-{uuid.uuid4()}" + redis_cache.set_cache(sync_key, {"v": "sync"}) + await redis_cache.async_set_cache(async_key, {"v": "async"}) + + assert dual_cache.get_cache(sync_key) == {"v": "sync"} + assert await dual_cache.async_get_cache(async_key) == {"v": "async"} + + with ( + patch.object(redis_cache, "get_cache") as sync_redis_read, + patch.object(redis_cache, "async_get_cache") as async_redis_read, + ): + assert dual_cache.get_cache(sync_key) == {"v": "sync"} + assert await dual_cache.async_get_cache(async_key) == {"v": "async"} + sync_redis_read.assert_not_called() + async_redis_read.assert_not_called() + + +@pytest.mark.asyncio +async def test_a_deleted_key_is_gone_from_both_memory_and_redis() -> None: + redis_cache: Final = _redis_cache() + dual_cache: Final = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis_cache) + sync_key: Final = f"deleted-sync-{uuid.uuid4()}" + async_key: Final = f"deleted-async-{uuid.uuid4()}" + dual_cache.set_cache(sync_key, {"v": "sync"}) + await dual_cache.async_set_cache(async_key, {"v": "async"}) + + dual_cache.delete_cache(sync_key) + await dual_cache.async_delete_cache(async_key) + + assert dual_cache.get_cache(sync_key) is None + assert await dual_cache.async_get_cache(async_key) is None + assert redis_cache.get_cache(sync_key) is None + assert await redis_cache.async_get_cache(async_key) is None + + +@pytest.mark.asyncio +async def test_a_batch_read_without_an_in_memory_cache_reads_redis() -> None: + redis_cache: Final = _redis_cache() + dual_cache: Final = DualCache(in_memory_cache=None, redis_cache=redis_cache) + key: Final = f"no-memory-{uuid.uuid4()}" + await redis_cache.async_set_cache(key, {"v": "from-redis"}) + + assert await dual_cache.async_batch_get_cache([key]) == [{"v": "from-redis"}] + + +@pytest.mark.asyncio +async def test_sync_and_async_batch_reads_share_one_redis_without_sync_reads_going_async() -> None: + redis_cache: Final = _redis_cache() + dual_cache: Final = DualCache(redis_cache=redis_cache) + run_id: Final = uuid.uuid4().hex + sync_keys: Final = [f"sync_{run_id}_{index}" for index in range(5)] + async_keys: Final = [f"async_{run_id}_{index}" for index in range(5)] + in_loop_keys: Final = [f"in_loop_{run_id}_{index}" for index in range(3)] + survivor_key: Final = f"survivor_{run_id}" + expected: Final = {key: {"key": key} for key in [*sync_keys, *async_keys, *in_loop_keys, survivor_key]} + await asyncio.gather(*(redis_cache.async_set_cache(key, value, ttl=60) for key, value in expected.items())) + + concurrent_results: Final = await asyncio.gather( + *(asyncio.to_thread(dual_cache.batch_get_cache, keys=[key]) for key in sync_keys), + *(dual_cache.async_batch_get_cache(keys=[key]) for key in async_keys), + ) + assert list(concurrent_results) == [[expected[key]] for key in [*sync_keys, *async_keys]] + + with patch.object( + redis_cache, + "async_batch_get_cache", + side_effect=AssertionError("sync batch reads must not call async Redis"), + ): + in_loop_results: Final = [dual_cache.batch_get_cache(keys=[key]) for key in in_loop_keys] + + assert in_loop_results == [[expected[key]] for key in in_loop_keys] + assert await dual_cache.async_batch_get_cache(keys=[survivor_key]) == [expected[survivor_key]] diff --git a/tests/local_testing/test_dual_cache.py b/tests/local_testing/test_dual_cache.py deleted file mode 100644 index 43b10a9557a..00000000000 --- a/tests/local_testing/test_dual_cache.py +++ /dev/null @@ -1,274 +0,0 @@ -import os -import time -import traceback -from litellm._uuid import uuid - -from dotenv import load_dotenv - -load_dotenv() - -import asyncio -import hashlib -import random - -import pytest - -import litellm -from litellm import aembedding, completion, embedding -from litellm.caching.caching import Cache - -from unittest.mock import AsyncMock, patch, MagicMock, call -import datetime -from datetime import timedelta -from litellm.caching import * - - -@pytest.mark.parametrize("is_async", [True, False]) -@pytest.mark.asyncio -async def test_dual_cache_get_set(is_async): - """Test that DualCache reads from in-memory cache first for both sync and async operations""" - in_memory = InMemoryCache() - redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=redis_cache) - - # Test basic set/get - test_key = f"test_key_{str(uuid.uuid4())}" - test_value = {"test": "value"} - - if is_async: - await dual_cache.async_set_cache(test_key, test_value) - mock_method = "async_get_cache" - else: - dual_cache.set_cache(test_key, test_value) - mock_method = "get_cache" - - # Mock Redis get to ensure we're not calling it - # this should only read in memory since we just set test_key - with patch.object(redis_cache, mock_method) as mock_redis_get: - if is_async: - result = await dual_cache.async_get_cache(test_key) - else: - result = dual_cache.get_cache(test_key) - - assert result == test_value - mock_redis_get.assert_not_called() # Verify Redis wasn't accessed - - -@pytest.mark.parametrize("is_async", [True, False]) -@pytest.mark.asyncio -async def test_dual_cache_local_only(is_async): - """Test that when local_only=True, only in-memory cache is used""" - in_memory = InMemoryCache() - redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=redis_cache) - - test_key = f"test_key_{str(uuid.uuid4())}" - test_value = {"test": "value"} - - # Mock Redis methods to ensure they're not called - redis_set_method = "async_set_cache" if is_async else "set_cache" - redis_get_method = "async_get_cache" if is_async else "get_cache" - - with ( - patch.object(redis_cache, redis_set_method) as mock_redis_set, - patch.object(redis_cache, redis_get_method) as mock_redis_get, - ): - - # Set value with local_only=True - if is_async: - await dual_cache.async_set_cache(test_key, test_value, local_only=True) - result = await dual_cache.async_get_cache(test_key, local_only=True) - else: - dual_cache.set_cache(test_key, test_value, local_only=True) - result = dual_cache.get_cache(test_key, local_only=True) - - assert result == test_value - mock_redis_set.assert_not_called() # Verify Redis set wasn't called - mock_redis_get.assert_not_called() # Verify Redis get wasn't called - - -@pytest.mark.parametrize("is_async", [True, False]) -@pytest.mark.asyncio -async def test_dual_cache_value_not_in_memory(is_async): - """Test that DualCache falls back to Redis when value isn't in memory, - and subsequent requests use in-memory cache""" - - in_memory = InMemoryCache() - redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=redis_cache) - - test_key = f"test_key_{str(uuid.uuid4())}" - test_value = {"test": "value"} - - # First, set value only in Redis - if is_async: - await redis_cache.async_set_cache(test_key, test_value) - else: - redis_cache.set_cache(test_key, test_value) - - # First request - should fall back to Redis and populate in-memory - if is_async: - result = await dual_cache.async_get_cache(test_key) - else: - result = dual_cache.get_cache(test_key) - - assert result == test_value - - # Second request - should now use in-memory cache - with patch.object( - redis_cache, "async_get_cache" if is_async else "get_cache" - ) as mock_redis_get: - if is_async: - result = await dual_cache.async_get_cache(test_key) - else: - result = dual_cache.get_cache(test_key) - - assert result == test_value - mock_redis_get.assert_not_called() # Verify Redis wasn't accessed second time - - -@pytest.mark.parametrize("is_async", [True, False]) -@pytest.mark.asyncio -async def test_dual_cache_batch_operations(is_async): - """Test batch get/set operations use in-memory cache correctly""" - in_memory = InMemoryCache() - redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=redis_cache) - - test_keys = [f"test_key_{str(uuid.uuid4())}" for _ in range(3)] - test_values = [{"test": f"value_{i}"} for i in range(3)] - cache_list = list(zip(test_keys, test_values)) - - # Set values - if is_async: - await dual_cache.async_set_cache_pipeline(cache_list) - else: - for key, value in cache_list: - dual_cache.set_cache(key, value) - - # Verify in-memory cache is used for subsequent reads - with patch.object( - redis_cache, "async_batch_get_cache" if is_async else "batch_get_cache" - ) as mock_redis_get: - if is_async: - results = await dual_cache.async_batch_get_cache(test_keys) - else: - results = dual_cache.batch_get_cache(test_keys, parent_otel_span=None) - - assert results == test_values - mock_redis_get.assert_not_called() - - -@pytest.mark.parametrize("is_async", [True, False]) -@pytest.mark.asyncio -async def test_dual_cache_increment(is_async): - """Test increment operations only use in memory when local_only=True""" - in_memory = InMemoryCache() - redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=redis_cache) - - test_key = f"counter_{str(uuid.uuid4())}" - increment_value = 1 - - # increment should use in-memory cache - with patch.object( - redis_cache, "async_increment" if is_async else "increment_cache" - ) as mock_redis_increment: - if is_async: - result = await dual_cache.async_increment_cache( - test_key, - increment_value, - local_only=True, - parent_otel_span=None, - ) - else: - result = dual_cache.increment_cache( - test_key, increment_value, local_only=True - ) - - assert result == increment_value - mock_redis_increment.assert_not_called() - - -@pytest.mark.asyncio -async def test_dual_cache_sadd(): - """Test set add operations use in-memory cache for reads""" - in_memory = InMemoryCache() - redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=redis_cache) - - test_key = f"set_{str(uuid.uuid4())}" - test_values = ["value1", "value2", "value3"] - - # Add values to set - await dual_cache.async_set_cache_sadd(test_key, test_values) - - # Verify in-memory cache is used for subsequent operations - with patch.object(redis_cache, "async_get_cache") as mock_redis_get: - result = await dual_cache.async_get_cache(test_key) - assert set(result) == set(test_values) - mock_redis_get.assert_not_called() - - -@pytest.mark.parametrize("is_async", [True, False]) -@pytest.mark.asyncio -async def test_dual_cache_delete(is_async): - """Test delete operations remove from both caches""" - in_memory = InMemoryCache() - redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=redis_cache) - - test_key = f"test_key_{str(uuid.uuid4())}" - test_value = {"test": "value"} - - # Set value - if is_async: - await dual_cache.async_set_cache(test_key, test_value) - else: - dual_cache.set_cache(test_key, test_value) - - # Delete value - if is_async: - await dual_cache.async_delete_cache(test_key) - else: - dual_cache.delete_cache(test_key) - - # Verify value is deleted from both caches - if is_async: - result = await dual_cache.async_get_cache(test_key) - else: - result = dual_cache.get_cache(test_key) - - assert result is None - - -@pytest.mark.asyncio -async def test_dual_cache_concurrent_sync_and_async_redis_reads(): - """Sync and async batch reads share one Redis backend in one process, and sync reads never open an async connection""" - redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - dual_cache = DualCache(redis_cache=redis_cache) - - run_id = str(uuid.uuid4()) - sync_keys = [f"sync_{run_id}_{index}" for index in range(5)] - async_keys = [f"async_{run_id}_{index}" for index in range(5)] - in_loop_keys = [f"in_loop_{run_id}_{index}" for index in range(3)] - survivor_key = f"survivor_{run_id}" - expected = {key: {"key": key} for key in [*sync_keys, *async_keys, *in_loop_keys, survivor_key]} - for key, value in expected.items(): - await redis_cache.async_set_cache(key, value, ttl=60) - - concurrent_results = await asyncio.gather( - *(asyncio.to_thread(dual_cache.batch_get_cache, keys=[key]) for key in sync_keys), - *(dual_cache.async_batch_get_cache(keys=[key]) for key in async_keys), - ) - assert list(concurrent_results) == [[expected[key]] for key in [*sync_keys, *async_keys]] - - with patch.object( - redis_cache, - "async_batch_get_cache", - side_effect=AssertionError("sync batch reads must not call async Redis"), - ): - in_loop_results = [dual_cache.batch_get_cache(keys=[key]) for key in in_loop_keys] - - assert in_loop_results == [[expected[key]] for key in in_loop_keys] - assert await dual_cache.async_batch_get_cache(keys=[survivor_key]) == [expected[survivor_key]] diff --git a/tests/local_testing/test_redis_batch_optimizations.py b/tests/local_testing/test_redis_batch_optimizations.py deleted file mode 100644 index d49939cff1a..00000000000 --- a/tests/local_testing/test_redis_batch_optimizations.py +++ /dev/null @@ -1,123 +0,0 @@ -""" -Tests for Redis batch caching optimizations (commit 3f52e8c) - -Verifies: - -1. Batch cache size increased from 100 → 1000 (minimum 1k) -2. Repeated Redis queries for cache misses are throttled -""" - -import os -import time -from unittest.mock import AsyncMock, patch - -import pytest -from dotenv import load_dotenv - -load_dotenv() - -import uuid -from litellm.caching.dual_cache import DualCache -from litellm.caching.in_memory_cache import InMemoryCache -from litellm.caching.redis_cache import RedisCache -from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE - - -@pytest.fixture -def cache_setup(): - """Create cache instances for testing""" - in_memory = InMemoryCache() - redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - dual_cache = DualCache( - in_memory_cache=in_memory, - redis_cache=redis_cache, - default_max_redis_batch_cache_size=DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE, - ) - return dual_cache, in_memory, redis_cache - - -@pytest.mark.asyncio -async def test_batch_cache_size_is_1000_minimum(cache_setup): - """Verify batch cache size is set to 1000 (never below 1k)""" - dual_cache, _, _ = cache_setup - - # Critical: batch cache size must be at least DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE - assert ( - dual_cache.last_redis_batch_access_time.max_size - >= DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE - ) - - -@pytest.mark.asyncio -async def test_throttling_prevents_duplicate_redis_calls(cache_setup): - """Test throttling prevents repeated Redis queries for cache misses""" - dual_cache, _, redis_cache = cache_setup - - test_keys = [f"miss_{str(uuid.uuid4())}" for _ in range(3)] - - # Set short expiry for testing - dual_cache.redis_batch_cache_expiry = 0.1 # 100ms - - with patch.object( - redis_cache, "async_batch_get_cache", new_callable=AsyncMock - ) as mock_redis: - mock_redis.return_value = {key: None for key in test_keys} - - # First call hits Redis (no throttle data exists) - await dual_cache.async_batch_get_cache(test_keys) - assert mock_redis.call_count == 1 - - # Second call immediately - throttled (within expiry window) - await dual_cache.async_batch_get_cache(test_keys) - assert mock_redis.call_count == 1 - - # Verify all keys tracked in throttle cache - for key in test_keys: - assert key in dual_cache.last_redis_batch_access_time - - # Wait for expiry time to pass - time.sleep(0.15) - - # Third call after expiry - call_count increases to 2 - await dual_cache.async_batch_get_cache(test_keys) - assert mock_redis.call_count == 2 - - -@pytest.mark.asyncio -async def test_basic_functionality_not_broken(cache_setup): - """Ensure basic cache functionality still works after optimizations""" - dual_cache, _, _ = cache_setup - - # Test basic set/get works - test_key = f"functional_test_{str(uuid.uuid4())}" - test_value = {"test": "data"} - - await dual_cache.async_set_cache(test_key, test_value) - result = await dual_cache.async_get_cache(test_key) - - assert result == test_value - - -@pytest.mark.asyncio -async def test_batch_get_with_no_in_memory_cache(): - """Test that batch get works when in_memory_cache is None""" - redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT")) - - # Create DualCache with no in-memory cache - dual_cache = DualCache( - in_memory_cache=None, # This is the edge case we're testing - redis_cache=redis_cache, - ) - - # Set some test data directly in Redis - test_key = f"no_memory_test_{str(uuid.uuid4())}" - test_value = {"test": "data_without_memory_cache"} - - await redis_cache.async_set_cache(test_key, test_value) - - # Should not crash when fetching from Redis without in-memory cache - result = await dual_cache.async_batch_get_cache([test_key]) - - assert result is not None - assert len(result) == 1 - assert result[0] == test_value diff --git a/tests/local_testing/test_router_utils.py b/tests/local_testing/test_router_utils.py index 635bda55144..aa617b09731 100644 --- a/tests/local_testing/test_router_utils.py +++ b/tests/local_testing/test_router_utils.py @@ -18,73 +18,6 @@ from unittest.mock import patch, MagicMock, AsyncMock load_dotenv() -def test_returned_settings(): - # this tests if the router raises an exception when invalid params are set - # in this test both deployments have bad keys - Keep this test. It validates if the router raises the most recent exception - litellm.set_verbose = True - import openai - - try: - print("testing if router raises an exception") - model_list = [ - { - "model_name": "gpt-3.5-turbo", # openai model name - "litellm_params": { # params for litellm completion/embedding call - "model": "azure/gpt-4.1-mini", - "api_key": "bad-key", - "api_version": os.getenv("AZURE_API_VERSION"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - }, - "tpm": 240000, - "rpm": 1800, - }, - { - "model_name": "gpt-3.5-turbo", # openai model name - "litellm_params": { # - "model": "gpt-3.5-turbo", - "api_key": "bad-key", - }, - "tpm": 240000, - "rpm": 1800, - }, - ] - router = Router( - model_list=model_list, - redis_host=os.getenv("REDIS_HOST"), - redis_password=os.getenv("REDIS_PASSWORD"), - redis_port=int(os.getenv("REDIS_PORT")), - routing_strategy="latency-based-routing", - routing_strategy_args={"ttl": 10}, - set_verbose=False, - num_retries=3, - retry_after=5, - allowed_fails=1, - cooldown_time=30, - ) # type: ignore - - settings = router.get_settings() - print(settings) - - """ - routing_strategy: "simple-shuffle" - routing_strategy_args: {"ttl": 10} # Average the last 10 calls to compute avg latency per model - allowed_fails: 1 - num_retries: 3 - retry_after: 5 # seconds to wait before retrying a failed request - cooldown_time: 30 # seconds to cooldown a deployment after failure - """ - assert settings["routing_strategy"] == "latency-based-routing" - assert settings["routing_strategy_args"]["ttl"] == 10 - assert settings["allowed_fails"] == 1 - assert settings["num_retries"] == 3 - assert settings["retry_after"] == 5 - assert settings["cooldown_time"] == 30 - - except Exception: - print(traceback.format_exc()) - pytest.fail("An error occurred - " + traceback.format_exc()) - - from litellm.types.utils import CallTypes diff --git a/tests/unit/caching/test_dual_cache.py b/tests/unit/caching/test_dual_cache.py index 521fda31b58..46600e0bf60 100644 --- a/tests/unit/caching/test_dual_cache.py +++ b/tests/unit/caching/test_dual_cache.py @@ -925,3 +925,100 @@ async def test_shared_batch_read_keeps_a_caches_own_tier_failure_to_itself_like_ assert shared == separate == [None, None, [3]] assert redis.async_batch_get_cache.await_args_list[0].args[0] == ["b1", "c1"] + + +def _write_through_dual_cache() -> tuple[DualCache, MagicMock]: + redis_cache: Final = MagicMock(spec=RedisCache) + return DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis_cache), redis_cache + + +@pytest.mark.asyncio +async def test_a_written_value_is_read_back_from_memory_without_a_redis_read(): + dual_cache, redis_cache = _write_through_dual_cache() + + dual_cache.set_cache("sync-key", {"v": 1}) + await dual_cache.async_set_cache("async-key", {"v": 2}) + + assert dual_cache.get_cache("sync-key") == {"v": 1} + assert await dual_cache.async_get_cache("async-key") == {"v": 2} + redis_cache.set_cache.assert_called_once() + redis_cache.async_set_cache.assert_awaited_once() + redis_cache.get_cache.assert_not_called() + redis_cache.async_get_cache.assert_not_called() + + +@pytest.mark.asyncio +async def test_local_only_reads_and_writes_never_reach_redis(): + dual_cache, redis_cache = _write_through_dual_cache() + + dual_cache.set_cache("sync-key", "sync", local_only=True) + await dual_cache.async_set_cache("async-key", "async", local_only=True) + + assert dual_cache.get_cache("sync-key", local_only=True) == "sync" + assert await dual_cache.async_get_cache("async-key", local_only=True) == "async" + assert dual_cache.get_cache("missing", local_only=True) is None + assert await dual_cache.async_get_cache("missing", local_only=True) is None + redis_cache.set_cache.assert_not_called() + redis_cache.async_set_cache.assert_not_called() + redis_cache.get_cache.assert_not_called() + redis_cache.async_get_cache.assert_not_called() + + +@pytest.mark.asyncio +async def test_batch_reads_of_written_keys_are_served_from_memory(): + dual_cache, redis_cache = _write_through_dual_cache() + entries: Final = (("a", {"v": "a"}), ("b", {"v": "b"}), ("c", {"v": "c"})) + + await dual_cache.async_set_cache_pipeline(entries) + dual_cache.set_cache("d", {"v": "d"}) + + assert await dual_cache.async_batch_get_cache(["a", "b", "c"]) == [{"v": "a"}, {"v": "b"}, {"v": "c"}] + assert dual_cache.batch_get_cache(["d"], parent_otel_span=None) == [{"v": "d"}] + redis_cache.async_set_cache_pipeline.assert_awaited_once() + redis_cache.async_batch_get_cache.assert_not_called() + redis_cache.batch_get_cache.assert_not_called() + + +@pytest.mark.asyncio +async def test_local_only_increments_count_in_memory_without_touching_redis(): + dual_cache, redis_cache = _write_through_dual_cache() + + assert dual_cache.increment_cache("sync-counter", 2, local_only=True) == 2 + assert dual_cache.increment_cache("sync-counter", 3, local_only=True) == 5 + assert await dual_cache.async_increment_cache("async-counter", 4, local_only=True) == 4 + redis_cache.increment_cache.assert_not_called() + redis_cache.async_increment.assert_not_called() + + +@pytest.mark.asyncio +async def test_set_members_added_through_the_dual_cache_are_read_from_memory(): + dual_cache, redis_cache = _write_through_dual_cache() + + await dual_cache.async_set_cache_sadd("members", ["value1", "value2", "value3"]) + + assert set(await dual_cache.async_get_cache("members")) == {"value1", "value2", "value3"} + redis_cache.async_set_cache_sadd.assert_awaited_once() + redis_cache.async_get_cache.assert_not_called() + + +def test_the_batch_read_throttle_tracks_at_least_the_default_number_of_keys(): + assert DualCache().last_redis_batch_access_time.max_size >= DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE + + +@pytest.mark.asyncio +async def test_async_batch_reads_of_missing_keys_hit_redis_once_per_expiry_window(): + redis_cache: Final = MagicMock(spec=RedisCache) + keys: Final = ["miss-a", "miss-b", "miss-c"] + redis_cache.async_batch_get_cache = AsyncMock(return_value=dict.fromkeys(keys)) + dual_cache: Final = DualCache( + in_memory_cache=InMemoryCache(), redis_cache=redis_cache, default_redis_batch_cache_expiry=60 + ) + + await dual_cache.async_batch_get_cache(keys) + await dual_cache.async_batch_get_cache(keys) + assert redis_cache.async_batch_get_cache.await_count == 1 + assert all(key in dual_cache.last_redis_batch_access_time for key in keys) + + dual_cache.last_redis_batch_access_time.update({key: time.time() - 61 for key in keys}) + await dual_cache.async_batch_get_cache(keys) + assert redis_cache.async_batch_get_cache.await_count == 2 diff --git a/tests/unit/test_router_get_settings.py b/tests/unit/test_router_get_settings.py new file mode 100644 index 00000000000..a4675715490 --- /dev/null +++ b/tests/unit/test_router_get_settings.py @@ -0,0 +1,26 @@ +from typing import Final + +from litellm import Router + + +def test_get_settings_returns_the_routing_and_retry_settings_the_router_was_built_with(): + router: Final = Router( + model_list=[ + {"model_name": "gpt-4.1-mini", "litellm_params": {"model": "openai/gpt-4.1-mini", "api_key": "fake-key"}} + ], + routing_strategy="latency-based-routing", + routing_strategy_args={"ttl": 10}, + num_retries=3, + retry_after=5, + allowed_fails=1, + cooldown_time=30, + ) + + settings: Final = router.get_settings() + + assert settings["routing_strategy"] == "latency-based-routing" + assert settings["routing_strategy_args"]["ttl"] == 10 + assert settings["allowed_fails"] == 1 + assert settings["num_retries"] == 3 + assert settings["retry_after"] == 5 + assert settings["cooldown_time"] == 30 From ba8cd1ee3159f94281ffb6e261be7e6a56290a84 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 1 Oct 2026 13:06:29 -0700 Subject: [PATCH 005/203] fix(router): keep silent_model out of embedding provider requests (#44064) Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/types/litellm_params.py | 1 + tests/unit/test_router_silent_experiment.py | 64 +++++++++++++++++++++ tests/unit/types/test_litellm_params.py | 1 + 3 files changed, 66 insertions(+) diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index 20214078852..51f6671e9d6 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -151,6 +151,7 @@ class DeploymentOptions: order: int | None = None tag_regex: Sequence[str] | None = None max_file_size_mb: float | None = None + silent_model: str | Sequence[str] | None = None @dataclass(frozen=True, slots=True, kw_only=True) diff --git a/tests/unit/test_router_silent_experiment.py b/tests/unit/test_router_silent_experiment.py index 722a76fa7ef..e184164d009 100644 --- a/tests/unit/test_router_silent_experiment.py +++ b/tests/unit/test_router_silent_experiment.py @@ -1,11 +1,14 @@ import asyncio +import json import time from collections.abc import Callable, Mapping from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +import respx import litellm from litellm.integrations.custom_logger import CustomLogger @@ -603,6 +606,67 @@ def test_silent_experiment_sends_shadow_request_attributed_to_the_silent_model(r assert primary_metadata == {"model_group": "primary-model"} + + +_EMBEDDING_API_BASE: Final = "https://embeddings.example.test/v1" + + +def _strict_embedding_route(respx_mock: respx.MockRouter) -> respx.Route: + return respx_mock.post(f"{_EMBEDDING_API_BASE}/embeddings").mock( + return_value=httpx.Response( + 200, + json={ + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}], + "model": "embed-model", + "usage": {"prompt_tokens": 2, "total_tokens": 2}, + }, + ) + ) + + +def _embedding_router_with_silent_model() -> Router: + return Router( + model_list=[ + { + "model_name": "embed-primary", + "litellm_params": { + "model": "openai/embed-model", + "api_base": _EMBEDDING_API_BASE, + "api_key": "fake-key", + "silent_model": "embed-shadow", + }, + } + ] + ) + + +def test_embedding_with_silent_model_sends_provider_body_without_it(respx_mock: respx.MockRouter) -> None: + route: Final = _strict_embedding_route(respx_mock) + + response: Final = _embedding_router_with_silent_model().embedding( + model="embed-primary", input=["black dresses"], input_type="query" + ) + + request_body: Final = json.loads(route.calls.last.request.read()) + assert request_body == {"model": "embed-model", "input": ["black dresses"], "input_type": "query"} + assert response.data[0]["embedding"] == [0.1, 0.2] + + +@pytest.mark.asyncio +async def test_aembedding_with_silent_model_sends_provider_body_without_it( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + route: Final = _strict_embedding_route(respx_mock) + + response: Final = await _embedding_router_with_silent_model().aembedding( + model="embed-primary", input=["black dresses"], input_type="query" + ) + + request_body: Final = json.loads(route.calls.last.request.read()) + assert request_body == {"model": "embed-model", "input": ["black dresses"], "input_type": "query"} + assert response.data[0]["embedding"] == [0.1, 0.2] @pytest.mark.parametrize("run_silent_experiment", SILENT_EXPERIMENT_RUNNERS) def test_silent_experiment_does_not_launch_from_a_shadow_request(run_silent_experiment): router = Router(model_list=_streaming_model_list(["shadow-a"])) diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index 16ae6b963d0..e3bbae39468 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -129,6 +129,7 @@ OPTION_NAMES: Final = ( "order", "tag_regex", "max_file_size_mb", + "silent_model", "auto_router_config_path", "auto_router_config", "auto_router_default_model", From ac8c5aa4b926ad07ac717b70ba8d7d2c6e8460b5 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 1 Oct 2026 13:26:16 -0700 Subject: [PATCH 006/203] fix(cost-map): add perplexity, openrouter, voyage and nebius models and fix registry metadata (#43907) * fix(cost-map): add nebius qwen3.8-27b and correct nebius context limits Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost-map): add and correct provider deprecation dates for deepseek, gemini and azure models Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost-map): add perplexity, openrouter and voyage models and correct gemini, nebius and perplexity metadata Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 204 +++++++++++++++--- model_prices_and_context_window.json | 204 +++++++++++++++--- 2 files changed, 356 insertions(+), 52 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 3392aa4d868..40fdf083cf8 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -5516,7 +5516,7 @@ "supports_web_search": false }, "azure/gpt-4.1-nano": { - "deprecation_date": "2027-04-14", + "deprecation_date": "2026-10-14", "cache_read_input_token_cost": 2.5e-08, "input_cost_per_token": 1e-07, "input_cost_per_token_batches": 5e-08, @@ -5550,7 +5550,7 @@ "supports_vision": true }, "azure/gpt-4.1-nano-2025-04-14": { - "deprecation_date": "2027-04-14", + "deprecation_date": "2026-10-14", "cache_read_input_token_cost": 2.5e-08, "input_cost_per_token": 1e-07, "input_cost_per_token_batches": 5e-08, @@ -6133,7 +6133,7 @@ "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, "cache_read_input_token_cost": 4e-07, - "deprecation_date": "2027-07-31", + "deprecation_date": "2027-06-25", "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, "input_cost_per_token": 4e-06, @@ -6168,7 +6168,7 @@ "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, "cache_read_input_token_cost": 6e-08, - "deprecation_date": "2027-07-31", + "deprecation_date": "2027-06-25", "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, "input_cost_per_token": 6e-07, @@ -6346,7 +6346,7 @@ "supports_tool_choice": true }, "azure/gpt-4o-transcribe": { - "deprecation_date": "2026-12-31", + "deprecation_date": "2026-10-15", "input_cost_per_audio_token": 2.5e-06, "input_cost_per_token": 2.5e-06, "litellm_provider": "azure", @@ -10972,7 +10972,7 @@ "supports_web_search": false }, "azure/us/gpt-4.1-nano-2025-04-14": { - "deprecation_date": "2027-04-14", + "deprecation_date": "2026-10-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, @@ -16146,6 +16146,7 @@ }, "deepseek-chat": { "cache_read_input_token_cost": 2.8e-08, + "deprecation_date": "2026-07-24", "input_cost_per_token": 2.8e-07, "litellm_provider": "deepseek", "max_input_tokens": 131072, @@ -16167,6 +16168,7 @@ }, "deepseek-reasoner": { "cache_read_input_token_cost": 2.8e-08, + "deprecation_date": "2026-07-24", "input_cost_per_token": 2.8e-07, "litellm_provider": "deepseek", "max_input_tokens": 131072, @@ -22085,6 +22087,7 @@ "deepseek/deepseek-chat": { "cache_creation_input_token_cost": 0.0, "cache_read_input_token_cost": 2.8e-08, + "deprecation_date": "2026-07-24", "input_cost_per_token": 2.8e-07, "input_cost_per_token_cache_hit": 2.8e-08, "litellm_provider": "deepseek", @@ -22139,6 +22142,7 @@ }, "deepseek/deepseek-reasoner": { "cache_read_input_token_cost": 2.8e-08, + "deprecation_date": "2026-07-24", "input_cost_per_token": 2.8e-07, "input_cost_per_token_cache_hit": 2.8e-08, "litellm_provider": "deepseek", @@ -29469,9 +29473,11 @@ "input_cost_per_token_batches": 6.25e-07, "input_cost_per_token_flex": 6.25e-07, "output_cost_per_token_batches": 5e-06, - "output_cost_per_token_flex": 5e-06 + "output_cost_per_token_flex": 5e-06, + "supports_url_context": true }, "gemini/gemini-2.5-computer-use-preview-10-2025": { + "deprecation_date": "2026-07-28", "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, "litellm_provider": "gemini", @@ -39903,6 +39909,7 @@ "nebius/deepseek-ai/DeepSeek-V4-Pro-0813": { "input_cost_per_token": 1.32e-06, "litellm_provider": "nebius", + "max_input_tokens": 979000, "mode": "chat", "output_cost_per_token": 3.96e-06, "source": "https://tokenfactory.nebius.com/models/catalog/text2text/deepseek-ai%2FDeepSeek-V4-Pro-0813", @@ -39912,7 +39919,7 @@ "nebius/deepseek-ai/DeepSeek-V4.1-Flash": { "input_cost_per_token": 3e-07, "litellm_provider": "nebius", - "max_input_tokens": 1048576, + "max_input_tokens": 1048000, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", @@ -40164,6 +40171,16 @@ "supports_reasoning": true, "source": "https://tokenfactory.nebius.com/models/catalog/text2text/Qwen%2FQwen3.5-397B-A17B" }, + "nebius/Qwen/Qwen3.8-27B": { + "input_cost_per_token": 4.5e-07, + "litellm_provider": "nebius", + "max_input_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 3e-06, + "source": "https://tokenfactory.nebius.com/models/catalog/text2text/Qwen%2FQwen3.8-27B", + "supports_function_calling": true, + "supports_reasoning": true + }, "nebius/zai-org/GLM-5.1": { "max_tokens": 202752, "max_input_tokens": 202752, @@ -40191,8 +40208,8 @@ "nebius/zai-org/GLM-5.3": { "input_cost_per_token": 1.4e-06, "litellm_provider": "nebius", - "max_input_tokens": 1048576, - "max_tokens": 1048576, + "max_input_tokens": 1024000, + "max_tokens": 1024000, "mode": "chat", "output_cost_per_token": 4.4e-06, "source": "https://tokenfactory.nebius.com/models/catalog/text2text/zai-org%2FGLM-5.3", @@ -40209,7 +40226,8 @@ "mode": "chat", "supports_function_calling": true, "supports_reasoning": true, - "source": "https://tokenfactory.nebius.com/models/catalog/text2text/zai-org%2FGLM-5.3-Flash" + "source": "https://tokenfactory.nebius.com/models/catalog/text2text/zai-org%2FGLM-5.3-Flash", + "supports_vision": true }, "nebius/BAAI/bge-en-icl": { "max_tokens": 32768, @@ -51467,6 +51485,26 @@ "mode": "rerank", "output_cost_per_token": 0.0 }, + "voyage/rerank-1": { + "input_cost_per_token": 5e-08, + "litellm_provider": "voyage", + "max_input_tokens": 8000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "rerank", + "output_cost_per_token": 0.0, + "source": "https://docs.voyageai.com/docs/pricing" + }, + "voyage/rerank-lite-1": { + "input_cost_per_token": 2e-08, + "litellm_provider": "voyage", + "max_input_tokens": 4000, + "max_output_tokens": 4000, + "max_tokens": 4000, + "mode": "rerank", + "output_cost_per_token": 0.0, + "source": "https://docs.voyageai.com/docs/pricing" + }, "voyage/rerank-2.5": { "input_cost_per_token": 5e-08, "litellm_provider": "voyage", @@ -51603,6 +51641,16 @@ "mode": "embedding", "output_cost_per_token": 0.0 }, + "voyage/voyage-large-2-instruct": { + "input_cost_per_token": 1.2e-07, + "litellm_provider": "voyage", + "max_input_tokens": 16000, + "max_tokens": 16000, + "mode": "embedding", + "output_cost_per_token": 0.0, + "output_vector_size": 1024, + "source": "https://docs.voyageai.com/docs/pricing" + }, "voyage/voyage-law-2": { "input_cost_per_token": 1.2e-07, "litellm_provider": "voyage", @@ -57215,7 +57263,8 @@ "supports_function_calling": true, "supports_response_schema": false, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_reasoning": true }, "gemini/gemini-3.1-flash-live-preview": { "input_cost_per_audio_token": 3e-06, @@ -57253,7 +57302,8 @@ "rpm": 10, "gemini_audio_only_live": true, "input_cost_per_second": 8.33333333333e-05, - "supports_response_schema": false + "supports_response_schema": false, + "supports_reasoning": true }, "gemini/gemini-3.1-flash-tts-preview": { "input_cost_per_token": 1e-06, @@ -57298,7 +57348,8 @@ "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" - ] + ], + "supports_prompt_caching": true }, "gemini/gemini-3.8-flash-lite-tts": { "cache_read_input_token_cost": 1.25e-07, @@ -57322,7 +57373,8 @@ "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" - ] + ], + "supports_prompt_caching": true }, "gemini-2.5-flash-preview-tts": { "input_cost_per_token": 5e-07, @@ -66081,7 +66133,7 @@ "gemini/lyria-3.5": { "input_cost_per_token": 0, "litellm_provider": "gemini", - "max_input_tokens": 1048576, + "max_input_tokens": 131072, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", @@ -66189,12 +66241,12 @@ "mode": "responses", "supports_web_search": true, "supports_function_calling": true, - "input_cost_per_token": 5e-06, - "output_cost_per_token": 3e-05, - "cache_read_input_token_cost": 5e-07, - "input_cost_per_token_above_272k_tokens": 1e-05, - "output_cost_per_token_above_272k_tokens": 4.5e-05, - "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "input_cost_per_token": 4e-06, + "output_cost_per_token": 2e-05, + "cache_read_input_token_cost": 4e-07, + "input_cost_per_token_above_272k_tokens": 8e-06, + "output_cost_per_token_above_272k_tokens": 3e-05, + "cache_read_input_token_cost_above_272k_tokens": 8e-07, "source": "https://docs.perplexity.ai/docs/agent-api/models" }, "perplexity/openai/gpt-5.6-terra": { @@ -70601,7 +70653,7 @@ "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" }, "azure/eu/gpt-4.1-nano": { - "deprecation_date": "2027-04-14", + "deprecation_date": "2026-10-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, @@ -70980,7 +71032,8 @@ "supports_function_calling": true, "supports_response_schema": false, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_reasoning": true }, "gemini/gemini-3.8-live-extended-thinking": { "input_cost_per_audio_token": 3e-06, @@ -71001,7 +71054,8 @@ "supports_function_calling": true, "supports_response_schema": false, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_reasoning": true }, "azure/us/codex-mini": { "deprecation_date": "2026-11-15", @@ -71048,7 +71102,7 @@ "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" }, "azure/us/gpt-4.1-nano": { - "deprecation_date": "2027-04-14", + "deprecation_date": "2026-10-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, @@ -73780,6 +73834,36 @@ "supports_vision": false, "supports_web_search": false }, + "openrouter/apodex/apodex-1.1-mini:free": { + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "openrouter", + "max_input_tokens": 262144, + "max_output_tokens": 235929, + "max_tokens": 235929, + "mode": "chat", + "source": "https://openrouter.ai/api/v1/models", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "openrouter/unbiased/pareto-26.10-preview": { + "input_cost_per_token": 8e-07, + "output_cost_per_token": 3.2e-06, + "cache_read_input_token_cost": 3e-08, + "litellm_provider": "openrouter", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "source": "https://openrouter.ai/api/v1/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_tool_choice": true, + "supports_vision": true + }, "openrouter/dots-studio/dots-3-note-preview:free": { "deprecation_date": "2026-12-31", "input_cost_per_token": 0.0, @@ -78890,6 +78974,74 @@ "cache_read_input_token_cost": 2e-07, "source": "https://docs.perplexity.ai/docs/agent-api/models" }, + "perplexity/anthropic/claude-fable-5-1": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 1e-05, + "output_cost_per_token": 5e-05, + "cache_read_input_token_cost": 2.5e-07, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/anthropic/claude-opus-5-5": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 4e-06, + "output_cost_per_token": 2e-05, + "cache_read_input_token_cost": 2e-07, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/openai/gpt-6.1-sol": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 2e-06, + "output_cost_per_token": 1e-05, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token_above_272k_tokens": 4e-06, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/openai/gpt-6-sol": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 2e-06, + "output_cost_per_token": 1e-05, + "cache_read_input_token_cost": 2e-07, + "input_cost_per_token_above_272k_tokens": 4e-06, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/openai/gpt-6-luna": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 1e-07, + "output_cost_per_token": 5e-07, + "cache_read_input_token_cost": 1e-08, + "input_cost_per_token_above_272k_tokens": 2e-07, + "output_cost_per_token_above_272k_tokens": 7.5e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-08, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/google/gemini-3.8-flash": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 7.5e-07, + "output_cost_per_token": 3.75e-06, + "cache_read_input_token_cost": 7.5e-08, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/xai/grok-4.7": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 2e-06, + "output_cost_per_token": 6e-06, + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token_above_200k_tokens": 4e-06, + "output_cost_per_token_above_200k_tokens": 1.2e-05, + "cache_read_input_token_cost_above_200k_tokens": 1e-06, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, "us-gov.anthropic.claude-sonnet-5-5": { "bedrock_converse_supports_strict_tools": false, "bedrock_output_config_effort_ceiling": "xhigh", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 3392aa4d868..40fdf083cf8 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -5516,7 +5516,7 @@ "supports_web_search": false }, "azure/gpt-4.1-nano": { - "deprecation_date": "2027-04-14", + "deprecation_date": "2026-10-14", "cache_read_input_token_cost": 2.5e-08, "input_cost_per_token": 1e-07, "input_cost_per_token_batches": 5e-08, @@ -5550,7 +5550,7 @@ "supports_vision": true }, "azure/gpt-4.1-nano-2025-04-14": { - "deprecation_date": "2027-04-14", + "deprecation_date": "2026-10-14", "cache_read_input_token_cost": 2.5e-08, "input_cost_per_token": 1e-07, "input_cost_per_token_batches": 5e-08, @@ -6133,7 +6133,7 @@ "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, "cache_read_input_token_cost": 4e-07, - "deprecation_date": "2027-07-31", + "deprecation_date": "2027-06-25", "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, "input_cost_per_token": 4e-06, @@ -6168,7 +6168,7 @@ "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, "cache_read_input_token_cost": 6e-08, - "deprecation_date": "2027-07-31", + "deprecation_date": "2027-06-25", "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, "input_cost_per_token": 6e-07, @@ -6346,7 +6346,7 @@ "supports_tool_choice": true }, "azure/gpt-4o-transcribe": { - "deprecation_date": "2026-12-31", + "deprecation_date": "2026-10-15", "input_cost_per_audio_token": 2.5e-06, "input_cost_per_token": 2.5e-06, "litellm_provider": "azure", @@ -10972,7 +10972,7 @@ "supports_web_search": false }, "azure/us/gpt-4.1-nano-2025-04-14": { - "deprecation_date": "2027-04-14", + "deprecation_date": "2026-10-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, @@ -16146,6 +16146,7 @@ }, "deepseek-chat": { "cache_read_input_token_cost": 2.8e-08, + "deprecation_date": "2026-07-24", "input_cost_per_token": 2.8e-07, "litellm_provider": "deepseek", "max_input_tokens": 131072, @@ -16167,6 +16168,7 @@ }, "deepseek-reasoner": { "cache_read_input_token_cost": 2.8e-08, + "deprecation_date": "2026-07-24", "input_cost_per_token": 2.8e-07, "litellm_provider": "deepseek", "max_input_tokens": 131072, @@ -22085,6 +22087,7 @@ "deepseek/deepseek-chat": { "cache_creation_input_token_cost": 0.0, "cache_read_input_token_cost": 2.8e-08, + "deprecation_date": "2026-07-24", "input_cost_per_token": 2.8e-07, "input_cost_per_token_cache_hit": 2.8e-08, "litellm_provider": "deepseek", @@ -22139,6 +22142,7 @@ }, "deepseek/deepseek-reasoner": { "cache_read_input_token_cost": 2.8e-08, + "deprecation_date": "2026-07-24", "input_cost_per_token": 2.8e-07, "input_cost_per_token_cache_hit": 2.8e-08, "litellm_provider": "deepseek", @@ -29469,9 +29473,11 @@ "input_cost_per_token_batches": 6.25e-07, "input_cost_per_token_flex": 6.25e-07, "output_cost_per_token_batches": 5e-06, - "output_cost_per_token_flex": 5e-06 + "output_cost_per_token_flex": 5e-06, + "supports_url_context": true }, "gemini/gemini-2.5-computer-use-preview-10-2025": { + "deprecation_date": "2026-07-28", "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, "litellm_provider": "gemini", @@ -39903,6 +39909,7 @@ "nebius/deepseek-ai/DeepSeek-V4-Pro-0813": { "input_cost_per_token": 1.32e-06, "litellm_provider": "nebius", + "max_input_tokens": 979000, "mode": "chat", "output_cost_per_token": 3.96e-06, "source": "https://tokenfactory.nebius.com/models/catalog/text2text/deepseek-ai%2FDeepSeek-V4-Pro-0813", @@ -39912,7 +39919,7 @@ "nebius/deepseek-ai/DeepSeek-V4.1-Flash": { "input_cost_per_token": 3e-07, "litellm_provider": "nebius", - "max_input_tokens": 1048576, + "max_input_tokens": 1048000, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", @@ -40164,6 +40171,16 @@ "supports_reasoning": true, "source": "https://tokenfactory.nebius.com/models/catalog/text2text/Qwen%2FQwen3.5-397B-A17B" }, + "nebius/Qwen/Qwen3.8-27B": { + "input_cost_per_token": 4.5e-07, + "litellm_provider": "nebius", + "max_input_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 3e-06, + "source": "https://tokenfactory.nebius.com/models/catalog/text2text/Qwen%2FQwen3.8-27B", + "supports_function_calling": true, + "supports_reasoning": true + }, "nebius/zai-org/GLM-5.1": { "max_tokens": 202752, "max_input_tokens": 202752, @@ -40191,8 +40208,8 @@ "nebius/zai-org/GLM-5.3": { "input_cost_per_token": 1.4e-06, "litellm_provider": "nebius", - "max_input_tokens": 1048576, - "max_tokens": 1048576, + "max_input_tokens": 1024000, + "max_tokens": 1024000, "mode": "chat", "output_cost_per_token": 4.4e-06, "source": "https://tokenfactory.nebius.com/models/catalog/text2text/zai-org%2FGLM-5.3", @@ -40209,7 +40226,8 @@ "mode": "chat", "supports_function_calling": true, "supports_reasoning": true, - "source": "https://tokenfactory.nebius.com/models/catalog/text2text/zai-org%2FGLM-5.3-Flash" + "source": "https://tokenfactory.nebius.com/models/catalog/text2text/zai-org%2FGLM-5.3-Flash", + "supports_vision": true }, "nebius/BAAI/bge-en-icl": { "max_tokens": 32768, @@ -51467,6 +51485,26 @@ "mode": "rerank", "output_cost_per_token": 0.0 }, + "voyage/rerank-1": { + "input_cost_per_token": 5e-08, + "litellm_provider": "voyage", + "max_input_tokens": 8000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "rerank", + "output_cost_per_token": 0.0, + "source": "https://docs.voyageai.com/docs/pricing" + }, + "voyage/rerank-lite-1": { + "input_cost_per_token": 2e-08, + "litellm_provider": "voyage", + "max_input_tokens": 4000, + "max_output_tokens": 4000, + "max_tokens": 4000, + "mode": "rerank", + "output_cost_per_token": 0.0, + "source": "https://docs.voyageai.com/docs/pricing" + }, "voyage/rerank-2.5": { "input_cost_per_token": 5e-08, "litellm_provider": "voyage", @@ -51603,6 +51641,16 @@ "mode": "embedding", "output_cost_per_token": 0.0 }, + "voyage/voyage-large-2-instruct": { + "input_cost_per_token": 1.2e-07, + "litellm_provider": "voyage", + "max_input_tokens": 16000, + "max_tokens": 16000, + "mode": "embedding", + "output_cost_per_token": 0.0, + "output_vector_size": 1024, + "source": "https://docs.voyageai.com/docs/pricing" + }, "voyage/voyage-law-2": { "input_cost_per_token": 1.2e-07, "litellm_provider": "voyage", @@ -57215,7 +57263,8 @@ "supports_function_calling": true, "supports_response_schema": false, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_reasoning": true }, "gemini/gemini-3.1-flash-live-preview": { "input_cost_per_audio_token": 3e-06, @@ -57253,7 +57302,8 @@ "rpm": 10, "gemini_audio_only_live": true, "input_cost_per_second": 8.33333333333e-05, - "supports_response_schema": false + "supports_response_schema": false, + "supports_reasoning": true }, "gemini/gemini-3.1-flash-tts-preview": { "input_cost_per_token": 1e-06, @@ -57298,7 +57348,8 @@ "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" - ] + ], + "supports_prompt_caching": true }, "gemini/gemini-3.8-flash-lite-tts": { "cache_read_input_token_cost": 1.25e-07, @@ -57322,7 +57373,8 @@ "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" - ] + ], + "supports_prompt_caching": true }, "gemini-2.5-flash-preview-tts": { "input_cost_per_token": 5e-07, @@ -66081,7 +66133,7 @@ "gemini/lyria-3.5": { "input_cost_per_token": 0, "litellm_provider": "gemini", - "max_input_tokens": 1048576, + "max_input_tokens": 131072, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", @@ -66189,12 +66241,12 @@ "mode": "responses", "supports_web_search": true, "supports_function_calling": true, - "input_cost_per_token": 5e-06, - "output_cost_per_token": 3e-05, - "cache_read_input_token_cost": 5e-07, - "input_cost_per_token_above_272k_tokens": 1e-05, - "output_cost_per_token_above_272k_tokens": 4.5e-05, - "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "input_cost_per_token": 4e-06, + "output_cost_per_token": 2e-05, + "cache_read_input_token_cost": 4e-07, + "input_cost_per_token_above_272k_tokens": 8e-06, + "output_cost_per_token_above_272k_tokens": 3e-05, + "cache_read_input_token_cost_above_272k_tokens": 8e-07, "source": "https://docs.perplexity.ai/docs/agent-api/models" }, "perplexity/openai/gpt-5.6-terra": { @@ -70601,7 +70653,7 @@ "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" }, "azure/eu/gpt-4.1-nano": { - "deprecation_date": "2027-04-14", + "deprecation_date": "2026-10-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, @@ -70980,7 +71032,8 @@ "supports_function_calling": true, "supports_response_schema": false, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_reasoning": true }, "gemini/gemini-3.8-live-extended-thinking": { "input_cost_per_audio_token": 3e-06, @@ -71001,7 +71054,8 @@ "supports_function_calling": true, "supports_response_schema": false, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_reasoning": true }, "azure/us/codex-mini": { "deprecation_date": "2026-11-15", @@ -71048,7 +71102,7 @@ "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" }, "azure/us/gpt-4.1-nano": { - "deprecation_date": "2027-04-14", + "deprecation_date": "2026-10-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, @@ -73780,6 +73834,36 @@ "supports_vision": false, "supports_web_search": false }, + "openrouter/apodex/apodex-1.1-mini:free": { + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "openrouter", + "max_input_tokens": 262144, + "max_output_tokens": 235929, + "max_tokens": 235929, + "mode": "chat", + "source": "https://openrouter.ai/api/v1/models", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "openrouter/unbiased/pareto-26.10-preview": { + "input_cost_per_token": 8e-07, + "output_cost_per_token": 3.2e-06, + "cache_read_input_token_cost": 3e-08, + "litellm_provider": "openrouter", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "source": "https://openrouter.ai/api/v1/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_tool_choice": true, + "supports_vision": true + }, "openrouter/dots-studio/dots-3-note-preview:free": { "deprecation_date": "2026-12-31", "input_cost_per_token": 0.0, @@ -78890,6 +78974,74 @@ "cache_read_input_token_cost": 2e-07, "source": "https://docs.perplexity.ai/docs/agent-api/models" }, + "perplexity/anthropic/claude-fable-5-1": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 1e-05, + "output_cost_per_token": 5e-05, + "cache_read_input_token_cost": 2.5e-07, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/anthropic/claude-opus-5-5": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 4e-06, + "output_cost_per_token": 2e-05, + "cache_read_input_token_cost": 2e-07, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/openai/gpt-6.1-sol": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 2e-06, + "output_cost_per_token": 1e-05, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token_above_272k_tokens": 4e-06, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/openai/gpt-6-sol": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 2e-06, + "output_cost_per_token": 1e-05, + "cache_read_input_token_cost": 2e-07, + "input_cost_per_token_above_272k_tokens": 4e-06, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/openai/gpt-6-luna": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 1e-07, + "output_cost_per_token": 5e-07, + "cache_read_input_token_cost": 1e-08, + "input_cost_per_token_above_272k_tokens": 2e-07, + "output_cost_per_token_above_272k_tokens": 7.5e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-08, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/google/gemini-3.8-flash": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 7.5e-07, + "output_cost_per_token": 3.75e-06, + "cache_read_input_token_cost": 7.5e-08, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "perplexity/xai/grok-4.7": { + "litellm_provider": "perplexity", + "mode": "responses", + "input_cost_per_token": 2e-06, + "output_cost_per_token": 6e-06, + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token_above_200k_tokens": 4e-06, + "output_cost_per_token_above_200k_tokens": 1.2e-05, + "cache_read_input_token_cost_above_200k_tokens": 1e-06, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, "us-gov.anthropic.claude-sonnet-5-5": { "bedrock_converse_supports_strict_tools": false, "bedrock_output_config_effort_ceiling": "xhigh", From 008fcb4fe3be331b8766d2e8ba31686319f0f8e4 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 1 Oct 2026 20:34:57 +0000 Subject: [PATCH 007/203] feat(tool-policies): show the user who owns the key that discovered a tool (#43892) * feat(tool-policies): show the user who owns the key that discovered a tool GET /v1/tool/list and GET /v1/tool/{tool_name} resolve the discovering key's owner from the verification token and user tables at response time and return it as a nullable user field. The Tool Policies page adds a User column that shows alias, then email, then ID, with the same cell the Virtual Keys page uses. Keys without an owner, deleted owners, and rows without a key hash show no user, and a database failure in the owner lookup keeps the tools listed with user null Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(tool-policies): bound the owner lookup with chunked membership queries The key-by-token and user-by-id lookups behind the tool rows' user field put every distinct key hash into one IN list. BaseRepository gains find_many_in, which runs the repository's chunked membership query and converts the rows like find_many does, and the owner lookup uses it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(tool-policies): cover the owner column across the tool routes and the dashboard Integration cells for the direct, detail and filtered tool routes, owners without alias or email, deleted owners and keys, keyless and unknown-key historical rows, more keys than one membership chunk, repeated reads, two-worker reads during discovery and a failed owner lookup. A Playwright cell drives the bundled Tool Policies page against the live proxy and follows the owner link Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: ryan Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_lazy_openapi_snapshot.json | 45 ++ litellm/proxy/db/tool_registry_writer.py | 59 ++- litellm/repositories/base_repository.py | 7 +- litellm/types/tool_management.py | 7 + tests/e2e/ui/fixtures/pages.ts | 1 + .../tests/integrationCritical/expected.json | 3 +- .../toolPoliciesUserColumn.spec.ts | 186 +++++++ tests/integration/_support/tool_rows.py | 19 + .../management/test_tool_policy_user.py | 486 ++++++++++++++++++ .../proxy/db/test_tool_registry_writer.py | 75 +++ tests/unit/repositories/test_repositories.py | 13 + .../ToolPoliciesTableColumns.test.tsx | 32 +- .../ToolPolicies/ToolPoliciesTableColumns.tsx | 28 +- .../src/components/networking.tsx | 5 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 10 + 15 files changed, 968 insertions(+), 8 deletions(-) create mode 100644 tests/e2e/ui/tests/integrationCritical/toolPoliciesUserColumn.spec.ts create mode 100644 tests/integration/_support/tool_rows.py create mode 100644 tests/integration/management/test_tool_policy_user.py diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 32559b98aea..663a8e0d6d8 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -52122,6 +52122,16 @@ ], "title": "Updated By" }, + "user": { + "anyOf": [ + { + "$ref": "#/components/schemas/ToolDiscoveryUser" + }, + { + "type": "null" + } + ] + }, "user_agent": { "anyOf": [ { @@ -52160,6 +52170,41 @@ "title": "ToolDetailResponse", "type": "object" }, + "ToolDiscoveryUser": { + "properties": { + "user_alias": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "User Alias" + }, + "user_email": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "User Email" + }, + "user_id": { + "title": "User Id", + "type": "string" + } + }, + "required": [ + "user_id" + ], + "title": "ToolDiscoveryUser", + "type": "object" + }, "ToolListResponse": { "properties": { "tools": { diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index cd0aa75b859..cef90eb89c2 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -8,6 +8,7 @@ Admins use the management endpoints to read and update input_policy / output_pol import uuid from collections.abc import Mapping, Sequence from datetime import datetime, timezone +from types import MappingProxyType from typing import TYPE_CHECKING, Final, Protocol from pydantic import TypeAdapter @@ -18,8 +19,11 @@ from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.repositories.object_permission_repository import ObjectPermissionRepository from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ToolRepository +from litellm.repositories.user_repository import UserRepository +from litellm.repositories.verification_token_repository import VerificationTokenRepository from litellm.types.tool_management import ( LiteLLM_ToolTableRow, + ToolDiscoveryUser, ToolPolicyOverrideRow, ) @@ -155,18 +159,65 @@ async def batch_upsert_tools( verbose_proxy_logger.error("tool_registry_writer batch_upsert_tools error: %s", e) +_NO_OWNERS: Final[Mapping[str, ToolDiscoveryUser]] = MappingProxyType({}) + + +async def _key_owners(prisma_client: "PrismaClient", key_hashes: frozenset[str]) -> Mapping[str, ToolDiscoveryUser]: + """Map each key hash to the user that owns the key, skipping keys without an owner or an unknown owner.""" + if not key_hashes: + return _NO_OWNERS + keys: Final = await VerificationTokenRepository(prisma_client).find_many_in("token", sorted(key_hashes)) + owner_ids: Final = frozenset(key.user_id for key in keys if key.user_id) + if not owner_ids: + return _NO_OWNERS + users: Final = await UserRepository(prisma_client).find_many_in("user_id", sorted(owner_ids)) + users_by_id: Final = MappingProxyType( + { + user.user_id: ToolDiscoveryUser( + user_id=user.user_id, user_email=user.user_email, user_alias=user.user_alias + ) + for user in users + } + ) + return MappingProxyType( + {key.token: users_by_id[key.user_id] for key in keys if key.token and key.user_id in users_by_id} + ) + + +async def _key_owners_or_none( + prisma_client: "PrismaClient", key_hashes: frozenset[str] +) -> Mapping[str, ToolDiscoveryUser]: + from prisma.errors import PrismaError + + try: + return await _key_owners(prisma_client, key_hashes) + except PrismaError as e: + verbose_proxy_logger.error("tool_registry_writer owner lookup error: %s", e) + return _NO_OWNERS + + +async def _with_owners( + prisma_client: "PrismaClient", tools: Sequence[LiteLLM_ToolTableRow] +) -> tuple[LiteLLM_ToolTableRow, ...]: + """Attach to each tool the user owning the key that discovered it; tools stay listed when that lookup fails.""" + owners: Final = await _key_owners_or_none( + prisma_client, frozenset(tool.key_hash for tool in tools if tool.key_hash) + ) + return tuple(tool.model_copy(update=MappingProxyType({"user": owners.get(tool.key_hash or "")})) for tool in tools) + + async def list_tools( prisma_client: "PrismaClient", input_policy: str | None = None, ) -> list[LiteLLM_ToolTableRow]: - """Return all tools, optionally filtered by input_policy.""" + """Return all tools, optionally filtered by input_policy, each with the user owning the key that discovered it.""" try: where: Final[Mapping[str, str]] = {"input_policy": input_policy} if input_policy is not None else {} rows: Final = await _tool_table_actions(prisma_client).find_many( where=where, order={"created_at": "desc"}, ) - return [_row_to_model(row) for row in rows] + return list(await _with_owners(prisma_client, tuple(_row_to_model(row) for row in rows))) except Exception as e: verbose_proxy_logger.error("tool_registry_writer list_tools error: %s", e) return [] @@ -176,14 +227,14 @@ async def get_tool( prisma_client: "PrismaClient", tool_name: str, ) -> LiteLLM_ToolTableRow | None: - """Return a single tool row by tool_name.""" + """Return a single tool row by tool_name, with the user owning the key that discovered it.""" try: row: Final = await _tool_table_actions(prisma_client).find_unique( where={"tool_name": tool_name}, ) if row is None: return None - return _row_to_model(row) + return (await _with_owners(prisma_client, (_row_to_model(row),)))[0] except Exception as e: verbose_proxy_logger.error("tool_registry_writer get_tool error: %s", e) return None diff --git a/litellm/repositories/base_repository.py b/litellm/repositories/base_repository.py index 065842b39e2..81fba770b70 100644 --- a/litellm/repositories/base_repository.py +++ b/litellm/repositories/base_repository.py @@ -3,11 +3,12 @@ Base repository class with common functionality. """ from abc import ABC, abstractmethod -from collections.abc import Iterable, Mapping, Sequence +from collections.abc import Hashable, Iterable, Mapping, Sequence from typing import Any, Final, Generic, Protocol, TypeVar, runtime_checkable from pydantic import BaseModel +from litellm.repositories.chunked_in import find_many_in from litellm.repositories.prisma_protocols import TableActions T = TypeVar("T", bound=BaseModel) @@ -92,6 +93,10 @@ class BaseRepository(ABC, Generic[T]): ) return self._to_model_list(records) + async def find_many_in(self, field: str, values: Iterable[Hashable]) -> list[T]: + """Records whose `field` is one of `values`, queried in chunks that stay under the bind-parameter cap.""" + return self._to_model_list(await find_many_in(self.table, field, values)) + async def create(self, data: Mapping[str, object]) -> T: """Create a new record.""" record: Final = await self.table.create(data=data) diff --git a/litellm/types/tool_management.py b/litellm/types/tool_management.py index 6fc19250ae9..13553dbecc6 100644 --- a/litellm/types/tool_management.py +++ b/litellm/types/tool_management.py @@ -13,6 +13,12 @@ ToolInputPolicy = Literal["trusted", "untrusted", "blocked"] ToolOutputPolicy = Literal["trusted", "untrusted"] +class ToolDiscoveryUser(BaseModel): + user_id: str + user_email: str | None = None + user_alias: str | None = None + + class LiteLLM_ToolTableRow(BaseModel): tool_id: str tool_name: str @@ -25,6 +31,7 @@ class LiteLLM_ToolTableRow(BaseModel): team_id: str | None = None key_alias: str | None = None user_agent: str | None = None + user: ToolDiscoveryUser | None = None last_used_at: datetime | None = None created_at: datetime | None = None updated_at: datetime | None = None diff --git a/tests/e2e/ui/fixtures/pages.ts b/tests/e2e/ui/fixtures/pages.ts index ba5887f3113..8210334c166 100644 --- a/tests/e2e/ui/fixtures/pages.ts +++ b/tests/e2e/ui/fixtures/pages.ts @@ -26,6 +26,7 @@ export enum Page { Logs = "logs", McpServers = "mcp-servers", SearchTools = "search-tools", + ToolPolicies = "tool-policies", TagManagement = "tag-management", VectorStores = "vector-stores", NewUsage = "new_usage", diff --git a/tests/e2e/ui/tests/integrationCritical/expected.json b/tests/e2e/ui/tests/integrationCritical/expected.json index 1614b188188..c6ee6051cd4 100644 --- a/tests/e2e/ui/tests/integrationCritical/expected.json +++ b/tests/e2e/ui/tests/integrationCritical/expected.json @@ -8,5 +8,6 @@ "tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::a server without per-user variables shows no credential row", "tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::clearing credentials for a server deleted underneath the modal reports the failure without losing the page", "tests/e2e/ui/tests/integrationCritical/costOptimizationModelGroups.spec.ts::cache leakage by model merges a deployment's resolved and requested model names into its model group", - "tests/e2e/ui/tests/integrationCritical/logsDrawerCredentialCanary.spec.ts::the Logs drawer renders the stored request without the deployment api_key" + "tests/e2e/ui/tests/integrationCritical/logsDrawerCredentialCanary.spec.ts::the Logs drawer renders the stored request without the deployment api_key", + "tests/e2e/ui/tests/integrationCritical/toolPoliciesUserColumn.spec.ts::the Tool Policies page names the user behind the key that discovered a tool" ] diff --git a/tests/e2e/ui/tests/integrationCritical/toolPoliciesUserColumn.spec.ts b/tests/e2e/ui/tests/integrationCritical/toolPoliciesUserColumn.spec.ts new file mode 100644 index 00000000000..c65c8774d89 --- /dev/null +++ b/tests/e2e/ui/tests/integrationCritical/toolPoliciesUserColumn.spec.ts @@ -0,0 +1,186 @@ +import { test, expect, type APIRequestContext } from "@playwright/test"; +import { randomUUID } from "node:crypto"; +import { execFileSync } from "node:child_process"; +import * as path from "node:path"; +import { Page } from "../../fixtures/pages"; +import { dismissFeedbackPopup, navigateToPage } from "../../helpers/navigation"; + +/** + * The Tool Policies table gets a User column: the owner of the key that discovered the tool, shown + * as alias (then email, then id) linking to the user's page, and a plain dash when the key has no + * owner. Both rows are produced the way a customer produces them, a chat completion carrying a + * tool through the proxy, so the column is read from the same registry the proxy writes. + */ +const unhex = (): string => randomUUID().replaceAll("-", ""); + +const toolCall = (model: string, toolName: string) => ({ + model, + messages: [{ role: "user", content: "tool policy user column" }], + tools: [ + { + type: "function", + function: { + name: toolName, + description: "integration tool", + parameters: { type: "object", properties: {} }, + }, + }, + ], +}); + +test("the Tool Policies page names the user behind the key that discovered a tool", async ({ + page, + request, +}) => { + const master = process.env.LITELLM_MASTER_KEY ?? "sk-integration-master"; + const upstream = ( + process.env.INTEGRATION_UPSTREAM_URL ?? "http://127.0.0.1:8190" + ).replace(/\/+$/, ""); + const auth = { Authorization: `Bearer ${master}` }; + const marker = unhex(); + const alias = `ui-owner-${marker}`; + const model = `ui-tool-policies-${marker}`; + const ownedTool = `ui_owned_tool_${marker}`; + const unownedTool = `ui_unowned_tool_${marker}`; + const support = (...args: string[]) => + execFileSync( + process.env.INTEGRATION_PYTHON ?? "python", + [ + path.resolve( + __dirname, + "../../../../integration/_support/tool_rows.py", + ), + ...args, + ], + { encoding: "utf8", timeout: 10_000, killSignal: "SIGKILL" }, + ); + + const post = async (api: APIRequestContext, route: string, data: object) => { + const response = await api.post(route, { headers: auth, data }); + expect(response.status(), `POST ${route}: ${await response.text()}`).toBe( + 200, + ); + return response.json(); + }; + + let modelId = ""; + let userId = ""; + const keys: string[] = []; + try { + modelId = ( + await post(request, "/model/new", { + model_name: model, + litellm_params: { + model: `openai/${model}`, + api_key: "sk-upstream", + api_base: `${upstream}/v1`, + }, + }) + ).model_id; + userId = ( + await post(request, "/user/new", { + user_id: `ui-user-${marker}`, + user_alias: alias, + user_email: `${alias}@integration.example`, + auto_create_key: false, + }) + ).user_id; + const ownedKey = ( + await post(request, "/key/generate", { user_id: userId, models: [model] }) + ).key; + const unownedKey = ( + await post(request, "/key/generate", { models: [model] }) + ).key; + keys.push(ownedKey, unownedKey); + for (const [key, toolName] of [ + [ownedKey, ownedTool], + [unownedKey, unownedTool], + ]) { + const response = await request.post("/v1/chat/completions", { + headers: { Authorization: `Bearer ${key}` }, + data: toolCall(model, toolName), + }); + expect(response.status(), await response.text()).toBe(200); + } + await expect + .poll( + async () => { + const response = await request.get("/v1/tool/list", { + headers: auth, + }); + if (response.status() !== 200) return []; + const names = ( + (await response.json()).tools as { tool_name: string }[] + ).map((tool) => tool.tool_name); + return [ownedTool, unownedTool].filter((name) => + names.includes(name), + ); + }, + { + timeout: 70_000, + message: "the discovered tools never reached the registry", + }, + ) + .toEqual([ownedTool, unownedTool]); + + await page.goto("/ui/login"); + await page.getByPlaceholder("Enter your username").fill("admin"); + await page.getByPlaceholder("Enter your password").fill(master); + await page.getByRole("button", { name: "Login", exact: true }).click(); + await expect(page).toHaveURL( + (url) => + url.pathname.startsWith("/ui") && !url.pathname.includes("login"), + ); + await navigateToPage(page, Page.ToolPolicies); + await dismissFeedbackPopup(page); + + const table = page.locator("table").filter({ visible: true }).first(); + const headers = table.getByRole("columnheader"); + await expect(headers.filter({ hasText: /^User$/ })).toHaveCount(1, { + timeout: 20_000, + }); + const headerTexts = (await headers.allInnerTexts()).map((text) => + text.trim(), + ); + const userColumn = headerTexts.indexOf("User"); + expect(userColumn, `columns: ${headerTexts.join(", ")}`).toBeGreaterThan( + -1, + ); + + const search = page + .getByTestId("datatable-search") + .filter({ visible: true }); + await expect(search).toBeVisible({ timeout: 20_000 }); + await search.fill(unownedTool); + const unownedRow = table + .locator("tbody tr") + .filter({ hasText: unownedTool }); + await expect(unownedRow).toHaveCount(1, { timeout: 30_000 }); + const unownedCell = unownedRow.getByRole("cell").nth(userColumn); + await expect(unownedCell).toHaveText("-"); + await expect(unownedCell.getByRole("link")).toHaveCount(0); + + await search.fill(ownedTool); + const ownedRow = table.locator("tbody tr").filter({ hasText: ownedTool }); + await expect(ownedRow).toHaveCount(1, { timeout: 30_000 }); + const ownerLink = ownedRow + .getByRole("cell") + .nth(userColumn) + .getByRole("link", { name: alias, exact: true }); + await expect(ownerLink).toBeVisible(); + expect(await ownerLink.getAttribute("href")).toContain( + `user=${encodeURIComponent(userId)}`, + ); + await ownerLink.click(); + await expect(page).toHaveURL( + (url) => + url.searchParams.get("user") === userId || + url.pathname.includes(userId), + ); + } finally { + support("clear", ownedTool, unownedTool); + if (keys.length) await post(request, "/key/delete", { keys }); + if (userId) await post(request, "/user/delete", { user_ids: [userId] }); + if (modelId) await post(request, "/model/delete", { id: modelId }); + } +}); diff --git a/tests/integration/_support/tool_rows.py b/tests/integration/_support/tool_rows.py new file mode 100644 index 00000000000..bb460022242 --- /dev/null +++ b/tests/integration/_support/tool_rows.py @@ -0,0 +1,19 @@ +import json +import sys +from typing import Final, LiteralString + +from integration._support.database import write_rows + +CLEAR_QUERY: Final[LiteralString] = 'DELETE FROM "LiteLLM_ToolTable" WHERE tool_name = %s' + + +def clear(tool_names: tuple[str, ...]) -> None: + for tool_name in tool_names: + write_rows(CLEAR_QUERY, (tool_name,)) + + +if __name__ == "__main__": + if sys.argv[1] != "clear": + raise SystemExit(f"unknown command: {sys.argv[1]}") + clear(tuple(sys.argv[2:])) + sys.stdout.write(json.dumps({"cleared": sys.argv[2:]}) + "\n") diff --git a/tests/integration/management/test_tool_policy_user.py b/tests/integration/management/test_tool_policy_user.py new file mode 100644 index 00000000000..00b4605ee2e --- /dev/null +++ b/tests/integration/management/test_tool_policy_user.py @@ -0,0 +1,486 @@ +import json +import time +import uuid +from collections.abc import Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack +from hashlib import sha256 +from pathlib import Path +from typing import Final, NamedTuple + +import jwt +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows, scratch_database, write_rows +from integration._support.process import owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, wire_server +from jwt.algorithms import RSAAlgorithm +from pydantic import JsonValue + +from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE + +AUDIENCE: Final = "litellm-integration" +KEY_ID: Final = "integration-signing-key" +CLIENT_CLAIM: Final = "client_id" + + +def _tool_call_request(model: str, tool_name: str) -> dict[str, JsonValue]: + return { + "model": model, + "messages": [{"role": "user", "content": "tool policy user control"}], + "tools": [ + { + "type": "function", + "function": { + "name": tool_name, + "description": "integration tool", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + } + + +def _forget_tool(tool_name: str) -> None: + write_rows('DELETE FROM "LiteLLM_ToolTable" WHERE tool_name = %s', (tool_name,)) + + +def _discovered_tool(gateway: Gateway, tool_name: str) -> dict[str, JsonValue]: + def rows() -> list[dict[str, JsonValue]]: + tools: Final = gateway.get("/v1/tool/list")["tools"] + assert isinstance(tools, list) + return [object_value(tool) for tool in tools if object_value(tool)["tool_name"] == tool_name] + + return eventually(rows, lambda found: len(found) == 1, seconds=70)[0] + + +def test_tool_list_reports_the_user_that_owns_the_discovering_key(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + alias: Final = "integration-alias-" + uuid.uuid4().hex + user: Final = scenario.user(user_alias=alias, user_email=f"{alias}@integration.example") + key: Final = scenario.key(user_id=user, models=[model]) + tool_name: Final = "integration_tool_" + uuid.uuid4().hex + scenario.cleanups.callback(_forget_tool, tool_name) + response: Final = gateway.request("POST", "/v1/chat/completions", _tool_call_request(model, tool_name), key=key) + assert response.status_code == 200, response.text + tool: Final = _discovered_tool(gateway, tool_name) + assert tool["key_hash"] == sha256(key.encode()).hexdigest(), tool + assert tool["user"] == {"user_id": user, "user_email": f"{alias}@integration.example", "user_alias": alias}, ( + tool + ) + + +def test_tool_list_reports_no_user_for_a_key_without_an_owner(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = scenario.key(models=[model]) + tool_name: Final = "integration_tool_" + uuid.uuid4().hex + scenario.cleanups.callback(_forget_tool, tool_name) + response: Final = gateway.request("POST", "/v1/chat/completions", _tool_call_request(model, tool_name), key=key) + assert response.status_code == 200, response.text + tool: Final = _discovered_tool(gateway, tool_name) + assert tool["key_hash"] == sha256(key.encode()).hexdigest(), tool + assert tool["user"] is None, tool + + +JWT_SETTINGS: Final[Mapping[str, JsonValue]] = { + "enable_jwt_auth": True, + "litellm_jwtauth": { + "user_id_jwt_field": "sub", + "user_email_jwt_field": "email", + "user_id_upsert": True, + "virtual_key_claim_field": CLIENT_CLAIM, + "unregistered_jwt_client_behavior": "auto_register", + }, +} + + +def _proxy_config( + directory: Path, model: str, upstream_url: str, general_settings: Mapping[str, JsonValue] = JWT_SETTINGS +) -> Path: + config: Final = directory / "tool_policy_user_config.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": model, + "litellm_params": { + "model": "openai/" + model, + "api_base": upstream_url + "/v1", + "api_key": "sk-upstream", + }, + } + ], + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + "store_model_in_db": True, + "proxy_batch_write_at": 1, + "proxy_batch_polling_interval": 1, + **general_settings, + }, + "router_settings": {"disable_cooldowns": True}, + } + ) + ) + return config + + +def _signed_token(private_key: rsa.RSAPrivateKey, user_id: str, email: str, client_id: str) -> str: + now: Final = int(time.time()) + return jwt.encode( + {"sub": user_id, "email": email, CLIENT_CLAIM: client_id, "aud": AUDIENCE, "iat": now, "exp": now + 300}, + private_key, + algorithm="RS256", + headers={"kid": KEY_ID}, + ) + + +def _forget_auto_registered_client(client_id: str, user_id: str) -> None: + write_rows( + 'DELETE FROM "LiteLLM_VerificationToken" WHERE token IN ' + '(SELECT token FROM "LiteLLM_JWTKeyMapping" WHERE jwt_claim_value = %s)', + (client_id,), + ) + write_rows('DELETE FROM "LiteLLM_JWTKeyMapping" WHERE jwt_claim_value = %s', (client_id,)) + write_rows('DELETE FROM "LiteLLM_UserTable" WHERE user_id = %s', (user_id,)) + + +def test_tool_list_reports_the_jwt_user_behind_an_auto_registered_key(gateway: Gateway, tmp_path: Path) -> None: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = json.loads(RSAAlgorithm.to_jwk(private_key.public_key())) + jwks: Final = json.dumps({"keys": [{**public_jwk, "kid": KEY_ID, "use": "sig", "alg": "RS256"}]}).encode() + + def respond(request: Request) -> Reply: + assert request.target == "/jwks", request + return Reply(body=jwks) + + model: Final = "integration-jwt-" + uuid.uuid4().hex + with wire_server(respond) as issuer: + config: Final = _proxy_config(tmp_path, model, gateway.upstream_url) + overrides: Final = {"JWT_PUBLIC_KEY_URL": issuer.url + "/jwks", "JWT_AUDIENCE": AUDIENCE} + with owned_proxy(gateway, tmp_path, overrides, config=config) as candidate, candidate.scenario() as scenario: + user: Final = "integration-jwt-user-" + uuid.uuid4().hex + email: Final = f"{user}@integration.example" + client_id: Final = "integration-client-" + uuid.uuid4().hex + tool_name: Final = "integration_tool_" + uuid.uuid4().hex + scenario.cleanups.callback(_forget_tool, tool_name) + scenario.cleanups.callback(_forget_auto_registered_client, client_id, user) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _tool_call_request(model, tool_name), + key=_signed_token(private_key, user, email, client_id), + ) + assert response.status_code == 200, response.text + mapped: Final = read_rows( + 'SELECT token FROM "LiteLLM_JWTKeyMapping" WHERE jwt_claim_name = %s AND jwt_claim_value = %s', + (CLIENT_CLAIM, client_id), + ) + assert len(mapped) == 1, mapped + assert read_rows( + 'SELECT user_id FROM "LiteLLM_VerificationToken" WHERE token = %s', (mapped[0]["token"],) + ) == [{"user_id": user}] + tool: Final = _discovered_tool(candidate, tool_name) + assert tool["key_hash"] == mapped[0]["token"], tool + assert tool["user"] == {"user_id": user, "user_email": email, "user_alias": None}, tool + + +def _owner(user_id: str, email: str | None, alias: str | None) -> dict[str, JsonValue]: + return {"user_id": user_id, "user_email": email, "user_alias": alias} + + +def _discover(gateway: Gateway, cleanups: ExitStack, model: str, key: str) -> str: + tool_name: Final = "integration_tool_" + uuid.uuid4().hex + cleanups.callback(_forget_tool, tool_name) + response: Final = gateway.request("POST", "/v1/chat/completions", _tool_call_request(model, tool_name), key=key) + assert response.status_code == 200, response.text + return tool_name + + +class Owned(NamedTuple): + tool_name: str + model: str + key: str + owner: dict[str, JsonValue] + + +def _owned_tool(gateway: Gateway, scenario: Scenario, alias: str | None = None) -> Owned: + """A discovered tool, the model and key that discovered it, and the owner the tool routes must report.""" + model: Final = scenario.model() + email: Final = f"{uuid.uuid4().hex}@integration.example" + fields: Final[Mapping[str, JsonValue]] = {"user_alias": alias} if alias else {} + user: Final = scenario.user(user_email=email, **fields) + key: Final = scenario.key(user_id=user, models=[model]) + return Owned(_discover(gateway, scenario.cleanups, model, key), model, key, _owner(user, email, alias)) + + +def _single(gateway: Gateway, tool_name: str) -> dict[str, JsonValue]: + return gateway.get(f"/v1/tool/{tool_name}") + + +def _detail_tool(gateway: Gateway, tool_name: str) -> dict[str, JsonValue]: + return object_value(gateway.get(f"/v1/tool/{tool_name}/detail")["tool"]) + + +def _listed_tools(gateway: Gateway, prefix: str, params: Mapping[str, str] | None = None) -> list[dict[str, JsonValue]]: + tools: Final = gateway.get("/v1/tool/list", params)["tools"] + assert isinstance(tools, list) + return [object_value(tool) for tool in tools if str(object_value(tool)["tool_name"]).startswith(prefix)] + + +def test_tool_get_reports_the_owner_and_null_for_an_unowned_key(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + owned, model, _, owner = _owned_tool(gateway, scenario, alias="alias-" + uuid.uuid4().hex) + unowned: Final = _discover(gateway, scenario.cleanups, model, scenario.key(models=[model])) + assert _discovered_tool(gateway, owned)["user"] == owner + _discovered_tool(gateway, unowned) + assert _single(gateway, owned)["user"] == owner + assert _single(gateway, unowned)["user"] is None + + +def test_tool_detail_carries_the_owner_inside_the_tool(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + tool_name, _, _, owner = _owned_tool(gateway, scenario, alias="alias-" + uuid.uuid4().hex) + assert _discovered_tool(gateway, tool_name)["user"] == owner + assert _detail_tool(gateway, tool_name)["user"] == owner + + +def test_filtered_tool_list_keeps_the_owner(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + tool_name, _, _, owner = _owned_tool(gateway, scenario) + listed: Final = _discovered_tool(gateway, tool_name) + assert listed["input_policy"] == "untrusted", listed + filtered: Final = _listed_tools(gateway, tool_name, {"input_policy": "untrusted"}) + assert [tool["user"] for tool in filtered] == [owner], filtered + assert _listed_tools(gateway, tool_name, {"input_policy": "blocked"}) == [] + + +def test_two_tools_discovered_by_the_same_key_share_the_owner(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + first, model, key, owner = _owned_tool(gateway, scenario) + second: Final = _discover(gateway, scenario.cleanups, model, key) + assert [_discovered_tool(gateway, name)["user"] for name in (first, second)] == [owner, owner] + + +def test_owner_without_alias_or_email_reports_only_the_user_id(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + user: Final = scenario.user() + key: Final = scenario.key(user_id=user, models=[model]) + tool_name: Final = _discover(gateway, scenario.cleanups, model, key) + assert _discovered_tool(gateway, tool_name)["user"] == _owner(user, None, None) + + +def test_missing_tool_is_404_on_get_and_detail(gateway: Gateway) -> None: + missing: Final = "integration_missing_" + uuid.uuid4().hex + for path in (f"/v1/tool/{missing}", f"/v1/tool/{missing}/detail"): + response: Final = gateway.request("GET", path) + assert response.status_code == 404, response.text + assert response.json() == {"detail": f"Tool '{missing}' not found"} + + +def test_non_admin_keys_are_rejected_on_every_tool_read_route(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + tool_name, _, _, owner = _owned_tool(gateway, scenario) + assert _discovered_tool(gateway, tool_name)["user"] == owner + internal: Final = scenario.key(user_id=scenario.user(user_role="internal_user")) + plain: Final = scenario.key() + for key in (internal, plain): + for path in ("/v1/tool/list", f"/v1/tool/{tool_name}", f"/v1/tool/{tool_name}/detail"): + response: Final = gateway.request("GET", path, key=key) + assert response.status_code == 401, (path, response.text) + assert string_value(owner["user_email"]) not in response.text, response.text + + +def test_unauthenticated_tool_reads_are_rejected(gateway: Gateway) -> None: + for path in ("/v1/tool/list", "/v1/tool/some_tool", "/v1/tool/some_tool/detail"): + response: Final = gateway.client.get(path) + assert response.status_code == 401, (path, response.text) + assert "No api key passed in" in response.text, response.text + + +def test_deleting_the_owner_keeps_the_tool_row_without_a_user(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + user: Final = uuid.uuid4().hex + gateway.post("/user/new", {"user_id": user, "auto_create_key": False}) + key: Final = string_value(gateway.post("/key/generate", {"user_id": user, "models": [model]})["key"]) + tool_name: Final = _discover(gateway, scenario.cleanups, model, key) + assert _discovered_tool(gateway, tool_name)["user"] == _owner(user, None, None) + deleted: Final = gateway.request("POST", "/user/delete", {"user_ids": [user]}) + assert deleted.status_code == 200, deleted.text + assert read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE user_id = %s', (user,)) == [] + tool: Final = _discovered_tool(gateway, tool_name) + assert tool["user"] is None, tool + assert tool["key_hash"] == sha256(key.encode()).hexdigest() + assert _single(gateway, tool_name)["user"] is None + + +def test_deleting_the_key_keeps_the_tool_row_without_a_user(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + user: Final = scenario.user() + key: Final = string_value(gateway.post("/key/generate", {"user_id": user, "models": [model]})["key"]) + tool_name: Final = _discover(gateway, scenario.cleanups, model, key) + assert _discovered_tool(gateway, tool_name)["user"] == _owner(user, None, None) + scenario.delete_key(key) + tool: Final = _discovered_tool(gateway, tool_name) + assert tool["user"] is None, tool + assert tool["key_hash"] == sha256(key.encode()).hexdigest() + + +def test_tool_row_without_a_key_hash_is_listed_without_a_user(gateway: Gateway) -> None: + tool_name: Final = "integration_tool_" + uuid.uuid4().hex + with ExitStack() as cleanups: + cleanups.callback(_forget_tool, tool_name) + write_rows( + 'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name) VALUES (gen_random_uuid()::text, %s)', (tool_name,) + ) + tool: Final = _discovered_tool(gateway, tool_name) + assert tool["key_hash"] is None, tool + assert tool["user"] is None, tool + assert _single(gateway, tool_name)["user"] is None + + +def test_tool_row_with_an_unknown_key_hash_is_listed_without_a_user(gateway: Gateway) -> None: + tool_name: Final = "integration_tool_" + uuid.uuid4().hex + key_hash: Final = "integration-unknown-" + uuid.uuid4().hex + with ExitStack() as cleanups: + cleanups.callback(_forget_tool, tool_name) + write_rows( + 'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, key_hash) VALUES (gen_random_uuid()::text, %s, %s)', + (tool_name, key_hash), + ) + tool: Final = _discovered_tool(gateway, tool_name) + assert tool["key_hash"] == key_hash, tool + assert tool["user"] is None, tool + + +def _forget_prefixed(prefix: str) -> None: + write_rows('DELETE FROM "LiteLLM_ToolTable" WHERE tool_name LIKE %s', (prefix + "%",)) + write_rows('DELETE FROM "LiteLLM_VerificationToken" WHERE token LIKE %s', (prefix + "%",)) + + +def test_owner_lookup_spans_more_keys_than_one_chunk(gateway: Gateway) -> None: + prefix: Final = "integration_chunk_" + uuid.uuid4().hex + "_" + count: Final = IN_LIST_CHUNK_SIZE + 1 + with gateway.scenario() as scenario: + user: Final = scenario.user(user_alias="chunk-owner-" + uuid.uuid4().hex) + scenario.cleanups.callback(_forget_prefixed, prefix) + write_rows( + 'INSERT INTO "LiteLLM_VerificationToken" (token, user_id) ' + "SELECT %s || g, %s FROM generate_series(1, %s::int) AS g", + (prefix, user, str(count)), + ) + write_rows( + 'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, key_hash) ' + "SELECT gen_random_uuid()::text, %s || g, %s || g FROM generate_series(1, %s::int) AS g", + (prefix, prefix, str(count)), + ) + listed: Final = _listed_tools(gateway, prefix) + assert len(listed) == count, len(listed) + owners: Final = {json.dumps(tool["user"], sort_keys=True) for tool in listed} + assert len(owners) == 1, owners + assert object_value(listed[0]["user"])["user_id"] == user, listed[0] + + +def test_repeated_tool_list_reads_are_identical_and_leave_rows_unchanged(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + tool_name, _, _, _ = _owned_tool(gateway, scenario) + first: Final = _discovered_tool(gateway, tool_name) + before: Final = read_rows( + 'SELECT tool_name, key_hash, call_count, updated_at::text FROM "LiteLLM_ToolTable" WHERE tool_name = %s', + (tool_name,), + ) + second: Final = _discovered_tool(gateway, tool_name) + after: Final = read_rows( + 'SELECT tool_name, key_hash, call_count, updated_at::text FROM "LiteLLM_ToolTable" WHERE tool_name = %s', + (tool_name,), + ) + assert first == second, (first, second) + assert before == after and len(before) == 1, (before, after) + + +def test_tool_list_total_matches_the_rows_in_postgres(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + tool_name, _, _, _ = _owned_tool(gateway, scenario) + _discovered_tool(gateway, tool_name) + body: Final = gateway.get("/v1/tool/list") + tools: Final = body["tools"] + assert isinstance(tools, list) + names: Final = sorted(str(object_value(tool)["tool_name"]) for tool in tools) + stored: Final = sorted( + str(row["tool_name"]) for row in read_rows('SELECT tool_name FROM "LiteLLM_ToolTable"', ()) + ) + assert body["total"] == len(tools) == len(stored), body["total"] + assert names == stored + + +def test_concurrent_tool_reads_on_two_workers_stay_consistent_during_discovery( + gateway: Gateway, tmp_path: Path +) -> None: + model: Final = "integration-workers-" + uuid.uuid4().hex + config: Final = _proxy_config(tmp_path, model, gateway.upstream_url, {}) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + with candidate.scenario() as scenario: + alias: Final = "burst-owner-" + uuid.uuid4().hex + user: Final = scenario.user(user_alias=alias) + key: Final = scenario.key(user_id=user, models=[model]) + steady: Final = _discover(candidate, scenario.cleanups, model, key) + assert _discovered_tool(candidate, steady)["user"] == _owner(user, None, alias) + paths: Final = tuple( + ("/v1/tool/list", f"/v1/tool/{steady}", f"/v1/tool/{steady}/detail")[index % 3] for index in range(40) + ) + + def read(index: int) -> tuple[int, dict[str, JsonValue], str | None]: + burst: Final = _discover(candidate, scenario.cleanups, model, key) if index == 20 else None + response: Final = candidate.request("GET", paths[index]) + assert response.status_code == 200, (paths[index], response.text) + return index, JSON_OBJECT.validate_json(response.content), burst + + with ThreadPoolExecutor(max_workers=16) as pool: + results: Final = tuple(pool.map(read, range(40))) + for index, body, _ in results: + tool: Final = ( + next(object_value(t) for t in body["tools"] if object_value(t)["tool_name"] == steady) + if paths[index].endswith("/list") + else object_value(body["tool"]) + if paths[index].endswith("/detail") + else body + ) + assert tool["user"] == _owner(user, None, alias), (paths[index], tool) + burst: Final = next(name for _, _, name in results if name) + assert _discovered_tool(candidate, burst)["user"] == _owner(user, None, alias) + + +def test_owner_lookup_failure_keeps_tools_listed_without_a_user(gateway: Gateway, tmp_path: Path) -> None: + model: Final = "integration-fault-" + uuid.uuid4().hex + config: Final = _proxy_config(tmp_path, model, gateway.upstream_url, {}) + with ( + scratch_database() as database_url, + owned_proxy(gateway, tmp_path, {"DATABASE_URL": database_url}, config=config) as candidate, + ): + alias: Final = "fault-owner-" + uuid.uuid4().hex + user: Final = string_value( + candidate.post("/user/new", {"user_alias": alias, "auto_create_key": False})["user_id"] + ) + key: Final = string_value(candidate.post("/key/generate", {"user_id": user, "models": [model]})["key"]) + with ExitStack() as cleanups: + tool_name: Final = _discover(candidate, cleanups, model, key) + cleanups.pop_all() + assert _discovered_tool(candidate, tool_name)["user"] == _owner(user, None, alias) + write_rows('ALTER TABLE "LiteLLM_UserTable" RENAME TO "LiteLLM_UserTable_away"', (), database_url=database_url) + try: + degraded: Final = _discovered_tool(candidate, tool_name) + assert degraded["user"] is None, degraded + assert degraded["key_hash"] == sha256(key.encode()).hexdigest(), degraded + assert _single(candidate, tool_name)["user"] is None + finally: + write_rows( + 'ALTER TABLE "LiteLLM_UserTable_away" RENAME TO "LiteLLM_UserTable"', (), database_url=database_url + ) + assert _discovered_tool(candidate, tool_name)["user"] == _owner(user, None, alias) diff --git a/tests/unit/proxy/db/test_tool_registry_writer.py b/tests/unit/proxy/db/test_tool_registry_writer.py index 6318e4422cf..c9df665741d 100644 --- a/tests/unit/proxy/db/test_tool_registry_writer.py +++ b/tests/unit/proxy/db/test_tool_registry_writer.py @@ -7,6 +7,7 @@ from datetime import datetime, timezone from unittest.mock import AsyncMock, MagicMock import pytest +from prisma.errors import PrismaError from litellm.proxy.db.tool_registry_writer import ( @@ -54,6 +55,8 @@ def _make_prisma( upsert_return=None, find_many_rows=None, find_unique_row=None, + key_rows=(), + user_rows=(), ): """Return a mock prisma_client with litellm_tooltable.upsert, find_many, find_unique.""" prisma = MagicMock() @@ -63,6 +66,10 @@ def _make_prisma( return_value=find_many_rows if find_many_rows is not None else [] ) prisma.db.litellm_tooltable.find_unique = AsyncMock(return_value=find_unique_row) + prisma.db.litellm_verificationtoken = MagicMock() + prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=list(key_rows)) + prisma.db.litellm_usertable = MagicMock() + prisma.db.litellm_usertable.find_many = AsyncMock(return_value=list(user_rows)) return prisma @@ -133,6 +140,56 @@ async def test_list_tools_no_filter(): assert call_kw["order"] == {"created_at": "desc"} +@pytest.mark.asyncio +async def test_list_tools_attaches_the_owner_of_the_discovering_key(): + owned = _mock_row(tool_id="id1", tool_name="owned_tool", key_hash="hash-owned") + orphan = _mock_row(tool_id="id2", tool_name="orphan_tool", key_hash="hash-orphan") + unknown_owner = _mock_row(tool_id="id3", tool_name="unknown_owner_tool", key_hash="hash-unknown-owner") + keyless = _mock_row(tool_id="id4", tool_name="keyless_tool", key_hash=None) + prisma = _make_prisma( + find_many_rows=[owned, orphan, unknown_owner, keyless], + key_rows=[ + {"token": "hash-owned", "user_id": "user-1"}, + {"token": "hash-orphan", "user_id": None}, + {"token": "hash-unknown-owner", "user_id": "user-gone"}, + ], + user_rows=[{"user_id": "user-1", "user_email": "one@example.com", "user_alias": "One"}], + ) + result = await list_tools(prisma) + assert [tool.model_dump(include={"tool_name", "user"}) for tool in result] == [ + { + "tool_name": "owned_tool", + "user": {"user_id": "user-1", "user_email": "one@example.com", "user_alias": "One"}, + }, + {"tool_name": "orphan_tool", "user": None}, + {"tool_name": "unknown_owner_tool", "user": None}, + {"tool_name": "keyless_tool", "user": None}, + ] + key_where = prisma.db.litellm_verificationtoken.find_many.call_args.kwargs["where"] + assert key_where == {"token": {"in": ["hash-orphan", "hash-owned", "hash-unknown-owner"]}} + user_where = prisma.db.litellm_usertable.find_many.call_args.kwargs["where"] + assert user_where == {"user_id": {"in": ["user-1", "user-gone"]}} + + +@pytest.mark.asyncio +async def test_list_tools_keeps_tools_without_owners_when_the_owner_lookup_fails(): + prisma = _make_prisma(find_many_rows=[_mock_row(tool_name="my_tool", key_hash="hash-owned")]) + prisma.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=PrismaError("verification token table down")) + result = await list_tools(prisma) + assert [tool.model_dump(include={"tool_name", "user"}) for tool in result] == [ + {"tool_name": "my_tool", "user": None} + ] + + +@pytest.mark.asyncio +async def test_list_tools_skips_owner_lookup_when_no_tool_has_a_key_hash(): + prisma = _make_prisma(find_many_rows=[_mock_row(key_hash=None)]) + result = await list_tools(prisma) + assert [tool.user for tool in result] == [None] + prisma.db.litellm_verificationtoken.find_many.assert_not_awaited() + prisma.db.litellm_usertable.find_many.assert_not_awaited() + + @pytest.mark.asyncio async def test_list_tools_with_input_policy_filter(): row = _mock_row( @@ -163,6 +220,24 @@ async def test_get_tool_found(): ) +@pytest.mark.asyncio +async def test_get_tool_attaches_the_owner_of_the_discovering_key(): + row = _mock_row(tool_name="my_tool", key_hash="hash-owned") + prisma = _make_prisma( + find_unique_row=row, + key_rows=[{"token": "hash-owned", "user_id": "user-1"}], + user_rows=[{"user_id": "user-1", "user_email": "one@example.com", "user_alias": "One"}], + ) + result = await get_tool(prisma, "my_tool") + assert result is not None + assert result.model_dump(include={"tool_name", "user"}) == { + "tool_name": "my_tool", + "user": {"user_id": "user-1", "user_email": "one@example.com", "user_alias": "One"}, + } + key_where = prisma.db.litellm_verificationtoken.find_many.call_args.kwargs["where"] + assert key_where == {"token": {"in": ["hash-owned"]}} + + @pytest.mark.asyncio async def test_get_tool_not_found(): prisma = _make_prisma(find_unique_row=None) diff --git a/tests/unit/repositories/test_repositories.py b/tests/unit/repositories/test_repositories.py index bd0f194b326..bae6db9ee88 100644 --- a/tests/unit/repositories/test_repositories.py +++ b/tests/unit/repositories/test_repositories.py @@ -196,6 +196,19 @@ class TestBaseRepository: budgets = await repo.find_many(where={"budget_id": "b1"}, skip=0, take=10, order={"budget_id": "asc"}) assert len(budgets) == 1 + @pytest.mark.asyncio + async def test_find_many_in_returns_models_from_every_chunk(self, prisma_client): + budget_ids: Final = tuple(f"b{i}" for i in range(IN_LIST_CHUNK_SIZE + 1)) + + async def find_many(where: dict[str, Any]) -> list[MockRecord]: + return [MockRecord({"budget_id": budget_id, "max_budget": 1.0}) for budget_id in where["budget_id"]["in"]] + + prisma_client.db.litellm_budgettable.find_many = AsyncMock(side_effect=find_many) + budgets = await BudgetRepository(prisma_client).find_many_in("budget_id", budget_ids) + assert [budget.budget_id for budget in budgets] == list(budget_ids) + assert all(isinstance(budget, LiteLLM_BudgetTable) for budget in budgets) + assert prisma_client.db.litellm_budgettable.find_many.await_count == 2 + def test_record_to_dict_branches(self): from litellm.repositories.base_repository import record_to_dict diff --git a/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.test.tsx b/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.test.tsx index bb4cd8a463a..4cb6da90cd1 100644 --- a/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.test.tsx +++ b/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.test.tsx @@ -1,10 +1,12 @@ -import { render, screen } from "@testing-library/react"; +import { render, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { describe, it, expect, vi } from "vitest"; import { flexRender, getCoreRowModel, useReactTable, type ColumnDef } from "@tanstack/react-table"; import { getToolPoliciesTableColumns } from "./ToolPoliciesTableColumns"; import type { ToolRow } from "@/components/networking"; +vi.mock("next/navigation", () => ({ useRouter: () => ({ push: vi.fn() }) })); + const row: ToolRow = { tool_name: "search_docs", input_policy: "untrusted", @@ -60,10 +62,38 @@ describe("getToolPoliciesTableColumns", () => { "team_id", "key_hash", "key_alias", + "user", "user_agent", ]); }); + it("shows the owning user's alias, linking to their detail page", () => { + renderTable({}, [{ ...row, user: { user_id: "user-1", user_email: "one@example.com", user_alias: "Team One" } }]); + + const link = screen.getByRole("link", { name: "Team One" }); + expect(link).toHaveAttribute("href", expect.stringContaining("user-1")); + expect(screen.queryByText("one@example.com")).not.toBeInTheDocument(); + }); + + it("falls back to the owning user's email, then id, when no alias is set", () => { + renderTable({}, [ + { ...row, tool_name: "by_email", user: { user_id: "user-1", user_email: "one@example.com", user_alias: null } }, + { ...row, tool_name: "by_id", user: { user_id: "user-2", user_email: null, user_alias: null } }, + ]); + + expect(screen.getByRole("link", { name: "one@example.com" })).toBeInTheDocument(); + expect(screen.getByRole("link", { name: "user-2" })).toBeInTheDocument(); + }); + + it("renders a dash without a link when the discovering key has no owner", () => { + renderTable({}, [{ ...row, user: null }]); + + const userIndex = getToolPoliciesTableColumns(defaultDeps).findIndex((c) => c.id === "user"); + const userCell = screen.getAllByRole("cell")[userIndex]; + expect(userCell).toHaveTextContent("-"); + expect(within(userCell).queryByRole("link")).not.toBeInTheDocument(); + }); + it("renders the row's identifying fields", () => { renderTable(); diff --git a/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.tsx b/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.tsx index d3d822759e2..5f38c58b1b4 100644 --- a/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.tsx @@ -4,7 +4,7 @@ import { ColumnDef } from "@tanstack/react-table"; import { ToolRow } from "@/components/networking"; import { DataTableSortHeader } from "@/components/shared/DataTable"; -import { DateCell, IdCell, IdentityCell } from "@/components/shared/table_cells"; +import { DateCell, IdCell, IdentityCell, UserPopoverCell } from "@/components/shared/table_cells"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { PolicySelect } from "./PolicySelect"; @@ -127,6 +127,32 @@ export const getToolPoliciesTableColumns = ({ meta: { title: "Key Name" }, cell: ({ row }) => , }, + { + id: "user", + accessorFn: (row) => row.user?.user_alias ?? row.user?.user_email ?? row.user?.user_id ?? "", + header: () => ( + + + User} /> + + The user who owns the key that discovered this tool. Displays the first available value: User Alias, User + Email, or User ID. + + + + ), + size: 160, + enableSorting: false, + meta: { title: "User" }, + cell: ({ row }) => ( + + ), + }, { id: "user_agent", accessorFn: (row) => row.user_agent ?? "", diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 129b1089e9b..d12d0a3219c 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -7614,6 +7614,11 @@ export interface ToolRow { created_by?: string; updated_by?: string; user_agent?: string; + user?: { + user_id: string; + user_email: string | null; + user_alias: string | null; + } | null; last_used_at?: string; } diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index fcb2805c767..ab164a61ca1 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -34736,6 +34736,7 @@ export interface components { updated_at?: string | null; /** Updated By */ updated_by?: string | null; + user?: components["schemas"]["ToolDiscoveryUser"] | null; /** User Agent */ user_agent?: string | null; }; @@ -45621,6 +45622,15 @@ export interface components { overrides?: components["schemas"]["ToolPolicyOverrideRow"][]; tool: components["schemas"]["LiteLLM_ToolTableRow"]; }; + /** ToolDiscoveryUser */ + ToolDiscoveryUser: { + /** User Alias */ + user_alias?: string | null; + /** User Email */ + user_email?: string | null; + /** User Id */ + user_id: string; + }; /** ToolFunction */ ToolFunction: { /** Defer Loading */ From eb103334eece522164137cfd915749159dcc4b96 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Thu, 1 Oct 2026 13:43:18 -0700 Subject: [PATCH 008/203] feat(proxy): gzip buffered responses for clients that accept it (#44052) * feat(proxy): gzip buffered responses for clients that accept it Large JSON reads like /user/daily/activity/aggregated shipped tens of MB uncompressed. Compress single-message bodies of 500B or more when the client's Accept-Encoding allows gzip (q-values and the wildcard honored). Streamed and etagged responses pass through untouched, every negotiable response carries Vary: Accept-Encoding, and bodies of 1MB or more are compressed in a worker thread * fix(proxy): skip partial and no-transform responses in gzip and always release the held start The gzip gate now also skips 206 Partial Content and Cache-Control: no-transform, since compressing either breaks byte ranges or ignores an explicit ban on transforms. A response start without a headers key no longer raises, and a start the app never follows with a body message is forwarded when the app returns instead of being dropped. Co-Authored-By: Claude Opus 5.5 --------- Co-authored-by: Claude Opus 5.5 --- litellm/proxy/middleware/gzip_middleware.py | 96 ++++++++ litellm/proxy/proxy_server.py | 2 + .../proxy/middleware/test_gzip_middleware.py | 213 ++++++++++++++++++ 3 files changed, 311 insertions(+) create mode 100644 litellm/proxy/middleware/gzip_middleware.py create mode 100644 tests/unit/proxy/middleware/test_gzip_middleware.py diff --git a/litellm/proxy/middleware/gzip_middleware.py b/litellm/proxy/middleware/gzip_middleware.py new file mode 100644 index 00000000000..016fec68312 --- /dev/null +++ b/litellm/proxy/middleware/gzip_middleware.py @@ -0,0 +1,96 @@ +import gzip +from types import MappingProxyType +from typing import Final + +import anyio.to_thread +from starlette.datastructures import Headers, MutableHeaders +from starlette.types import ASGIApp, Message, Receive, Scope, Send + +MINIMUM_SIZE_BYTES: Final = 500 +OFF_LOOP_SIZE_BYTES: Final = 1024 * 1024 +COMPRESS_LEVEL: Final = 6 + + +def _coding_weight(part: str) -> tuple[str, float]: + coding, _, params = part.partition(";") + qvalue: Final = next((p.strip()[2:] for p in params.split(";") if p.strip().lower().startswith("q=")), "1") + try: + return coding.strip().lower(), float(qvalue) + except ValueError: + return coding.strip().lower(), 0.0 + + +def accepts_gzip(accept_encoding: str) -> bool: + weights: Final = MappingProxyType(dict(_coding_weight(part) for part in accept_encoding.split(",") if part.strip())) + return weights.get("gzip", weights.get("x-gzip", weights.get("*", 0.0))) > 0 + + +async def _compress(body: bytes) -> bytes: + if len(body) < OFF_LOOP_SIZE_BYTES: + return gzip.compress(body, compresslevel=COMPRESS_LEVEL) + return await anyio.to_thread.run_sync(gzip.compress, body, COMPRESS_LEVEL) + + +class _BufferedBodyGzipResponder: + """Holds the response start until the first body message shows the body is complete, so streams are never delayed.""" + + def __init__(self, send: Send, gzip_accepted: bool) -> None: + self.send = send + self.gzip_accepted = gzip_accepted + self.held_start: Message | None = None + self.decided = False + + async def __call__(self, message: Message) -> None: + if self.decided: + await self.send(message) + return + if message["type"] == "http.response.start": + self.held_start = message + return + self.decided = True + start: Final = self.held_start + if start is None: + await self.send(message) + return + body: Final[bytes] = message.get("body", b"") + start.setdefault("headers", ()) + headers: Final = MutableHeaders(scope=start) + negotiable: Final = ( + message["type"] == "http.response.body" + and not message.get("more_body", False) + and len(body) >= MINIMUM_SIZE_BYTES + and "content-encoding" not in headers + and "etag" not in headers + and start["status"] != 206 + and "no-transform" not in headers.get("cache-control", "").lower() + ) + if negotiable: + headers.add_vary_header("Accept-Encoding") + if not (negotiable and self.gzip_accepted): + await self.send(start) + await self.send(message) + return + compressed: Final = await _compress(body) + headers["content-encoding"] = "gzip" + headers["content-length"] = str(len(compressed)) + await self.send(start) + await self.send({**message, "body": compressed}) + + async def release_held_start(self) -> None: + if not self.decided and self.held_start is not None: + self.decided = True + await self.send(self.held_start) + + +class GZipBufferedResponseMiddleware: + def __init__(self, app: ASGIApp) -> None: + self.app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return + gzip_accepted: Final = accepts_gzip(Headers(scope=scope).get("accept-encoding", "")) + responder: Final = _BufferedBodyGzipResponder(send, gzip_accepted) + await self.app(scope, receive, responder) + await responder.release_held_start() diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4581441ea3a..ff0dd9df16f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -725,6 +725,7 @@ from litellm.proxy.middleware.admission_control_middleware import ( admission_control_state, get_admission_control_settings, ) +from litellm.proxy.middleware.gzip_middleware import GZipBufferedResponseMiddleware from litellm.proxy.middleware.in_flight_requests_middleware import ( InFlightRequestsMiddleware, ) @@ -2444,6 +2445,7 @@ app.add_middleware(BudgetReservationReleaseMiddleware, release=release_unbound_b app.add_middleware(RedisRequestBatchMiddleware) app.add_middleware(InFlightRequestsMiddleware) app.add_middleware(SecurityHeadersMiddleware) +app.add_middleware(GZipBufferedResponseMiddleware) def mount_swagger_ui(): diff --git a/tests/unit/proxy/middleware/test_gzip_middleware.py b/tests/unit/proxy/middleware/test_gzip_middleware.py new file mode 100644 index 00000000000..271ae46bb89 --- /dev/null +++ b/tests/unit/proxy/middleware/test_gzip_middleware.py @@ -0,0 +1,213 @@ +import asyncio +import gzip +import json +from typing import Final + +import pytest +from starlette.applications import Starlette +from starlette.requests import Request +from starlette.responses import JSONResponse, Response, StreamingResponse +from starlette.routing import Route +from starlette.types import ASGIApp, Message, Receive, Scope, Send + +from litellm.proxy.middleware.gzip_middleware import ( + MINIMUM_SIZE_BYTES, + OFF_LOOP_SIZE_BYTES, + GZipBufferedResponseMiddleware, +) + +LARGE_PAYLOAD = {"rows": [{"date": f"2026-09-{day:02d}", "spend": day * 1.5} for day in range(1, 31)] * 20} +STREAM_CHUNKS = tuple(json.dumps({"part": part, "pad": "x" * MINIMUM_SIZE_BYTES}).encode() for part in range(3)) + + +async def _large_json(request: Request) -> Response: + return JSONResponse(LARGE_PAYLOAD) + + +async def _small_json(request: Request) -> Response: + return JSONResponse({"ok": True}) + + +async def _already_encoded(request: Request) -> Response: + return Response(b"x" * (MINIMUM_SIZE_BYTES * 4), headers={"content-encoding": "br"}) + + +async def _with_etag(request: Request) -> Response: + return Response(b"y" * (MINIMUM_SIZE_BYTES * 4), headers={"etag": '"v1"'}) + + +async def _partial(request: Request) -> Response: + return Response(b"p" * (MINIMUM_SIZE_BYTES * 4), status_code=206, headers={"content-range": "bytes 0-1999/9000"}) + + +async def _no_transform(request: Request) -> Response: + return Response(b"n" * (MINIMUM_SIZE_BYTES * 4), headers={"cache-control": "public, no-transform"}) + + +async def _huge(request: Request) -> Response: + return Response(b"z" * (OFF_LOOP_SIZE_BYTES * 2), media_type="application/json") + + +async def _json_stream(request: Request) -> Response: + async def chunks(): + for chunk in STREAM_CHUNKS: + yield chunk + + return StreamingResponse(chunks(), media_type="application/json") + + +APP = Starlette( + routes=[ + Route("/large", _large_json), + Route("/small", _small_json), + Route("/encoded", _already_encoded), + Route("/stream", _json_stream), + Route("/etag", _with_etag), + Route("/huge", _huge), + Route("/partial", _partial), + Route("/no-transform", _no_transform), + ] +) +APP.add_middleware(GZipBufferedResponseMiddleware) + + +async def _send_messages(path: str, accept_encoding: str | None, app: ASGIApp = APP) -> tuple[Message, ...]: + headers = [(b"accept-encoding", accept_encoding.encode())] if accept_encoding is not None else [] + scope = {"type": "http", "method": "GET", "path": path, "query_string": b"", "headers": headers} + sent: list[Message] = [] # mutable-ok: ASGI send callback collects messages in order + requests: Final = iter(({"type": "http.request", "body": b"", "more_body": False},)) + never_disconnects: Final = asyncio.Event() + + async def receive() -> Message: + request: Final = next(requests, None) + if request is not None: + return request + await never_disconnects.wait() + return {"type": "http.disconnect"} + + async def send(message: Message) -> None: + sent.append(message) + + await app(scope, receive, send) + return tuple(sent) + + +def _headers(messages: tuple[Message, ...]) -> dict[str, str]: + return {k.decode(): v.decode() for k, v in messages[0]["headers"]} + + +def _body(messages: tuple[Message, ...]) -> bytes: + return b"".join(m.get("body", b"") for m in messages[1:]) + + +@pytest.mark.parametrize("accept_encoding", ["gzip, deflate, br", "GZIP", "br;q=1, gzip;q=0.5", "x-gzip", "*"]) +@pytest.mark.asyncio +async def test_large_buffered_json_is_gzipped_and_round_trips(accept_encoding): + messages = await _send_messages("/large", accept_encoding) + headers = _headers(messages) + body = _body(messages) + + assert headers["content-encoding"] == "gzip" + assert headers["vary"] == "Accept-Encoding" + assert int(headers["content-length"]) == len(body) + assert json.loads(gzip.decompress(body)) == LARGE_PAYLOAD + assert len(body) < len(json.dumps(LARGE_PAYLOAD)) + + +@pytest.mark.asyncio +async def test_body_above_off_loop_threshold_round_trips(): + messages = await _send_messages("/huge", "gzip") + + assert _headers(messages)["content-encoding"] == "gzip" + assert gzip.decompress(_body(messages)) == b"z" * (OFF_LOOP_SIZE_BYTES * 2) + + +@pytest.mark.parametrize( + ("path", "accept_encoding", "expected_vary"), + [ + ("/large", None, "Accept-Encoding"), + ("/large", "gzip;q=0", "Accept-Encoding"), + ("/small", "gzip", None), + ("/etag", "gzip", None), + ("/stream", "gzip", None), + ], +) +@pytest.mark.asyncio +async def test_vary_marks_every_negotiable_variant(path, accept_encoding, expected_vary): + messages = await _send_messages(path, accept_encoding) + + assert _headers(messages).get("vary") == expected_vary + + +@pytest.mark.parametrize( + ("path", "accept_encoding", "expected_encoding"), + [ + ("/large", None, None), + ("/large", "identity", None), + ("/large", "gzip;q=0", None), + ("/large", "br, gzip; q=0.0", None), + ("/large", "*;q=0", None), + ("/large", "*, gzip;q=0", None), + ("/large", "gzip;q=invalid", None), + ("/small", "gzip", None), + ("/encoded", "gzip", "br"), + ("/etag", "gzip", None), + ("/stream", "gzip", None), + ("/partial", "gzip", None), + ("/no-transform", "gzip", None), + ], +) +@pytest.mark.asyncio +async def test_response_passes_through_unmodified(path, accept_encoding, expected_encoding): + with_header = await _send_messages(path, accept_encoding) + without_header = await _send_messages(path, None) + + assert _headers(with_header).get("content-encoding") == expected_encoding + assert _body(with_header) == _body(without_header) + + +@pytest.mark.asyncio +async def test_streamed_chunks_are_forwarded_one_by_one(): + messages = await _send_messages("/stream", "gzip") + chunks = tuple(m["body"] for m in messages[1:] if m.get("body")) + + assert [m["type"] for m in messages].count("http.response.start") == 1 + assert chunks == STREAM_CHUNKS + + +@pytest.mark.asyncio +async def test_start_message_without_headers_key_is_still_gzipped(): + body: Final = b"h" * (MINIMUM_SIZE_BYTES * 4) + + async def headerless_app(scope: Scope, receive: Receive, send: Send) -> None: + await send({"type": "http.response.start", "status": 200}) + await send({"type": "http.response.body", "body": body}) + + messages = await _send_messages("/", "gzip", GZipBufferedResponseMiddleware(headerless_app)) + + assert _headers(messages)["content-encoding"] == "gzip" + assert gzip.decompress(_body(messages)) == body + + +@pytest.mark.asyncio +async def test_start_without_a_body_message_is_still_forwarded(): + async def start_only_app(scope: Scope, receive: Receive, send: Send) -> None: + await send({"type": "http.response.start", "status": 204, "headers": [(b"x-done", b"1")]}) + + messages = await _send_messages("/", "gzip", GZipBufferedResponseMiddleware(start_only_app)) + + assert messages == ({"type": "http.response.start", "status": 204, "headers": [(b"x-done", b"1")]},) + + +def test_proxy_app_gzips_large_responses_for_clients_that_accept_it(): + from starlette.testclient import TestClient + + from litellm.proxy.proxy_server import app + + client = TestClient(app) + compressed = client.get("/openapi.json", headers={"accept-encoding": "gzip"}) + identity = client.get("/openapi.json", headers={"accept-encoding": "identity"}) + + assert compressed.headers["content-encoding"] == "gzip" + assert int(compressed.headers["content-length"]) < int(identity.headers["content-length"]) + assert compressed.json() == identity.json() From 163ebccad55cf2036f65feeb2d0f0d1fc9d24e7f Mon Sep 17 00:00:00 2001 From: tin-berri Date: Thu, 1 Oct 2026 13:44:32 -0700 Subject: [PATCH 009/203] fix(auto-router): show actual and baseline spend for historical savings (#44057) The usage card hid actual and baseline spend unless every older session could be rebuilt from SpendLogs within two seconds, which on a real gateway it never was. Each complexity router's actual spend is now its rollup spend and its baseline is spend plus recorded savings, for old and new requests alike, so the benchmarks and session endpoints never scan SpendLogs. Adaptive and quality routers record no savings baseline and stay out of the compared totals; savings_estimated_classifier_cost is kept and covers the same compared requests --- .../proxy/db/autorouter_savings_comparison.py | 147 ------------------ .../auto_router_endpoints.py | 124 ++++----------- .../auto_router_endpoints.py | 25 ++- .../spend/test_autorouter_session_rollup.py | 65 -------- .../test_auto_router_endpoints.py | 60 +++++-- .../AutoRouterBenchmarksTab.test.tsx | 26 ++-- .../_components/AutoRouterBenchmarksTab.tsx | 27 ++-- ...KeyAutoRouterUsageTab.integration.test.tsx | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 32 ++-- 9 files changed, 131 insertions(+), 376 deletions(-) delete mode 100644 litellm/proxy/db/autorouter_savings_comparison.py diff --git a/litellm/proxy/db/autorouter_savings_comparison.py b/litellm/proxy/db/autorouter_savings_comparison.py deleted file mode 100644 index 041496d63f2..00000000000 --- a/litellm/proxy/db/autorouter_savings_comparison.py +++ /dev/null @@ -1,147 +0,0 @@ -from collections.abc import Mapping -from contextlib import AbstractAsyncContextManager -from datetime import timedelta -from math import isclose -from types import MappingProxyType -from typing import TYPE_CHECKING, Final, Protocol, cast - -from pydantic import BaseModel, ConfigDict, TypeAdapter - -from litellm._logging import verbose_proxy_logger -from litellm.constants import MAX_SPENDLOG_ROWS_TO_QUERY -from litellm.proxy.db.autorouter_session_rollup import AUTOROUTER_SESSION_WINDOW_SQL -from litellm.proxy.db.create_views import SupportsRawQueries - -if TYPE_CHECKING: - from litellm.proxy.utils import PrismaClient - - -class SessionSavingsComparison(BaseModel): - model_config = ConfigDict(frozen=True, allow_inf_nan=False) - - router_name: str - router_type: str - turns: int - estimated_turns: int - actual_spend: float - classifier_cost: float | None - saved_spend: float - complete: bool - - def coverage_fields(self, recorded_savings: float, recorded_turns: int) -> Mapping[str, float | int]: - if self.turns != recorded_turns or not self.complete: - return MappingProxyType({}) - if not isclose(self.saved_spend, recorded_savings, rel_tol=1e-9, abs_tol=1e-9): - return MappingProxyType({}) - return MappingProxyType( - { - "savings_estimated_turns": self.estimated_turns, - "savings_estimated_actual_spend": self.actual_spend, - "savings_estimated_saved_spend": self.saved_spend, - } - ) - - -class _ReadTransactions(Protocol): - def tx(self, *, timeout: timedelta, max_wait: timedelta) -> AbstractAsyncContextManager[SupportsRawQueries]: ... - - -_COMPARISONS: Final = TypeAdapter(tuple[SessionSavingsComparison, ...]) - - -async def historical_session_comparisons( - prisma_client: "PrismaClient", - start_date: str, - end_date: str, - api_key: str | None, - user_id: str | None, - session_id: str | None = None, -) -> Mapping[tuple[str, str], SessionSavingsComparison]: - try: - reader: Final = cast(_ReadTransactions, prisma_client.read_db) # cast-ok: untyped Prisma transaction delegate - async with reader.tx(timeout=timedelta(seconds=3), max_wait=timedelta(seconds=1)) as transaction: - await transaction.execute_raw("SET TRANSACTION READ ONLY") - await transaction.execute_raw("SET LOCAL statement_timeout = 2000") - rows: Final = await transaction.query_raw( - HISTORICAL_SESSION_COMPARISONS_SQL, - start_date, - end_date, - api_key, - user_id, - session_id, - ) - comparisons: Final = _COMPARISONS.validate_python(rows or ()) - return MappingProxyType({(row.router_name, row.router_type): row for row in comparisons}) - except Exception: # noqa: BLE001 # missing retained logs must not discard recorded dollar savings - verbose_proxy_logger.warning("Historical auto-router cost comparison unavailable; preserving recorded savings") - return MappingProxyType({}) - - -HISTORICAL_SESSION_COMPARISONS_SQL: Final = f""" -WITH {AUTOROUTER_SESSION_WINDOW_SQL}, scoped AS MATERIALIZED ( - SELECT * FROM windowed WHERE $5::text IS NULL OR session_id = $5::text -), limited_logs AS MATERIALIZED ( - SELECT session.api_key, session.session_id, session.router_name, session.router_type, session.comparison_user_id, - session.classifier_cost_recorded_turns = session.turns AS classifier_cost_tracked, - logs.spend, logs.prompt_tokens + logs.completion_tokens AS tokens, - logs.metadata::jsonb -> 'routing_decision' AS decision, - logs.metadata::jsonb -> 'autorouter_savings' AS savings, - logs.metadata::jsonb -> 'autorouter_savings_estimate' AS estimate - FROM scoped AS session JOIN "LiteLLM_SpendLogs" AS logs - ON logs.api_key = session.api_key - AND CASE WHEN char_length(logs.session_id) > 256 - THEN 'sha256:' || encode(sha256(convert_to(logs.session_id, 'UTF8')), 'hex') - ELSE logs.session_id END = session.session_id - AND (session.comparison_user_id IS NULL OR logs."user" = session.comparison_user_id) - AND logs."startTime" BETWEEN session.first_turn_at AND session.last_turn_at - AND COALESCE(logs.metadata::jsonb #>> '{{routing_decision,router_model_name}}', logs.model_group) - = session.router_name - WHERE session.savings_estimated_turns < session.turns - AND logs.status = 'success' AND COALESCE(logs.metadata::jsonb ->> 'internal_call_origin', '') = '' - LIMIT {MAX_SPENDLOG_ROWS_TO_QUERY + 1} -), facts AS ( - SELECT *, - CASE WHEN jsonb_typeof(decision -> 'classifier_cost') = 'number' - THEN (decision ->> 'classifier_cost')::float8 - WHEN classifier_cost_tracked THEN 0 END AS classifier, - CASE WHEN jsonb_typeof(savings) = 'number' AND ( - estimate IS NULL OR estimate = 'null'::jsonb OR ( - jsonb_typeof(estimate -> 'version') = 'number' AND estimate ->> 'version' IN ('1', '2', '3') - AND estimate ->> 'status' = 'estimated' - ) - ) THEN savings::text::float8 END AS saved - FROM limited_logs -), compared AS ( - SELECT api_key, session_id, router_name, router_type, comparison_user_id, - COUNT(*) AS turns, SUM(spend + COALESCE(classifier, 0)) AS spend, SUM(tokens) AS total_tokens, - COUNT(saved) AS estimated_turns, - COALESCE(SUM(spend + COALESCE(classifier, 0)) FILTER (WHERE saved IS NOT NULL), 0)::float8 AS actual_spend, - CASE WHEN COUNT(saved) = COUNT(classifier) FILTER (WHERE saved IS NOT NULL) - THEN COALESCE(SUM(classifier) FILTER (WHERE saved IS NOT NULL), 0)::float8 - END AS estimated_classifier_cost, - COALESCE(SUM(saved), 0)::float8 AS saved_spend - FROM facts GROUP BY 1, 2, 3, 4, 5 -), reconciled AS ( - SELECT session.*, logs.estimated_turns, logs.actual_spend, logs.estimated_classifier_cost, - COALESCE((SELECT COUNT(*) FROM limited_logs) <= {MAX_SPENDLOG_ROWS_TO_QUERY} - AND logs.turns = session.turns AND logs.total_tokens = session.total_tokens - AND ABS(logs.spend - session.spend) <= GREATEST(1e-9, ABS(session.spend) * 1e-9) - AND ABS(logs.saved_spend - session.saved_spend) <= GREATEST(1e-9, ABS(session.saved_spend) * 1e-9), FALSE - ) AS recovered - FROM scoped AS session LEFT JOIN compared AS logs - ON logs.api_key = session.api_key AND logs.session_id = session.session_id - AND logs.router_name = session.router_name AND logs.router_type = session.router_type - AND logs.comparison_user_id IS NOT DISTINCT FROM session.comparison_user_id -) -SELECT router_name, router_type, - SUM(turns)::bigint AS turns, - SUM(CASE WHEN recovered THEN estimated_turns ELSE savings_estimated_turns END)::bigint AS estimated_turns, - SUM(CASE WHEN recovered THEN actual_spend ELSE savings_estimated_actual_spend END)::float8 AS actual_spend, - CASE WHEN BOOL_AND(CASE WHEN recovered THEN estimated_classifier_cost IS NOT NULL - ELSE savings_estimated_turns = turns AND classifier_cost_recorded_turns = turns END) - THEN SUM(CASE WHEN recovered THEN estimated_classifier_cost ELSE classifier_cost END)::float8 - END AS classifier_cost, - SUM(saved_spend)::float8 AS saved_spend, - BOOL_AND(recovered OR savings_estimated_turns = turns) AS complete -FROM reconciled GROUP BY router_name, router_type -""" diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 3bd02fe8738..8b73b8177f4 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -8,7 +8,6 @@ POST /auto_router/validate_complexity_router_config - Dry-run the complexity-rou from collections.abc import Mapping, Sequence from datetime import datetime, timedelta, timezone from itertools import chain, groupby -from math import isclose from types import MappingProxyType from typing import TYPE_CHECKING, Annotated, Final, Protocol from uuid import uuid4 @@ -32,7 +31,6 @@ from litellm.proxy.auth.auth_checks import ( can_key_call_resolved_model, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.db.autorouter_savings_comparison import historical_session_comparisons from litellm.proxy.db.autorouter_session_rollup import ( AUTOROUTER_BENCHMARKS_SQL, bounded_session_id, @@ -643,7 +641,6 @@ class _SessionAggRow(BaseModel): savings_estimated_actual_spend: float = 0.0 savings_estimated_classifier_cost: float | None = None savings_estimated_saved_spend: float = 0.0 - savings_comparison_complete: bool = True classifier_cost: float classifier_cost_recorded_turns: int session_seconds: float @@ -671,25 +668,35 @@ def _cache_bucket(turns: int, hits: int) -> AutoRouterCacheBucket: def _savings_cohort( - turns: int, estimated_turns: int, actual_spend: float, saved_spend: float, recorded_savings: float + turns: int, estimated_turns: int, spend: float, saved_spend: float ) -> tuple[float | None, float | None]: - if turns > 0 and estimated_turns == 0 and recorded_savings == 0: + if turns > 0 and estimated_turns == 0 and saved_spend == 0: return None, None - if not isclose(saved_spend, recorded_savings, rel_tol=1e-9, abs_tol=1e-9): - return recorded_savings, None - return recorded_savings, actual_spend + recorded_savings + return saved_spend, spend + saved_spend + + +def _compared_row(row: _SessionAggRow) -> _SessionAggRow: + _, baseline_spend = _savings_cohort(row.turns, row.savings_estimated_turns, row.spend, row.saved_spend) + compared: Final = row.router_type == "complexity" and baseline_spend is not None + return row.model_copy( + update={ + "savings_estimated_turns": row.turns if compared else 0, + "savings_estimated_actual_spend": row.spend if compared else 0.0, + "savings_estimated_classifier_cost": ( + row.classifier_cost if row.classifier_cost_recorded_turns == row.turns else None + ) + if compared + else 0.0, + "savings_estimated_saved_spend": row.saved_spend if compared else 0.0, + } + ) def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals: return_misses: Final = row.return_turns - row.return_hits - saved_spend, compared_baseline = _savings_cohort( - row.turns, - row.savings_estimated_turns, - row.savings_estimated_actual_spend, - row.savings_estimated_saved_spend, - row.saved_spend, + saved_spend, baseline_spend = _savings_cohort( + row.turns, row.savings_estimated_turns, row.savings_estimated_actual_spend, row.savings_estimated_saved_spend ) - baseline_spend: Final = compared_baseline if row.savings_comparison_complete else None sessions: Final = row.sessions return AutoRouterBenchmarkTotals( sessions=sessions, @@ -700,7 +707,7 @@ def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals: spend=row.spend, savings_estimated_turns=row.savings_estimated_turns, savings_estimated_actual_spend=row.savings_estimated_actual_spend, - savings_estimated_classifier_cost=row.savings_estimated_classifier_cost if baseline_spend is not None else None, + savings_estimated_classifier_cost=row.savings_estimated_classifier_cost, saved_spend=saved_spend, classifier_cost=row.classifier_cost if row.classifier_cost_recorded_turns == row.turns else None, baseline_spend=baseline_spend, @@ -777,7 +784,6 @@ def _summed_agg_row(rows: Sequence[_SessionAggRow]) -> _SessionAggRow: else None ), savings_estimated_saved_spend=sum(row.savings_estimated_saved_spend for row in rows), - savings_comparison_complete=all(row.savings_comparison_complete for row in rows), classifier_cost=sum(row.classifier_cost for row in rows), classifier_cost_recorded_turns=sum(row.classifier_cost_recorded_turns for row in rows), session_seconds=sum(row.session_seconds for row in rows), @@ -852,8 +858,8 @@ async def get_auto_router_benchmarks( Benchmarks for the auto-router dashboard: session shape, savings against the configured baseline, and prompt-caching behaviour bucketed by what the router did. - Reads session rollups folded once per request at spend-write time, with bounded - retained-log recovery for historical comparisons. A user filter selects only turns attributed to that + Reads session rollups folded once per request at spend-write time, so this endpoint + never scans LiteLLM_SpendLogs. A user filter selects only turns attributed to that internal user when written; older key-only history remains outside user views. A session is in the window when it overlaps it: its last turn is on or after start_date and its first turn is on or before end_date. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is @@ -887,44 +893,7 @@ async def get_auto_router_benchmarks( api_key, user_id, ) - recorded_rows: Final = _SESSION_AGG_ROWS.validate_python(raw_rows or ()) - comparisons: Final = ( - await historical_session_comparisons( - prisma_client, - start_day.isoformat(), - (end_day + timedelta(days=1)).isoformat(), - api_key, - user_id, - ) - if any(row.savings_estimated_turns < row.turns for row in recorded_rows) - else MappingProxyType({}) - ) - covered_rows: Final = tuple( - row.model_copy( - update={ - **comparison.coverage_fields(row.saved_spend, row.turns), - "savings_estimated_classifier_cost": comparison.classifier_cost, - "savings_comparison_complete": comparison.complete and comparison.turns == row.turns, - } - ) - if (comparison := comparisons.get((row.router_name, row.router_type))) - else row.model_copy(update={"savings_comparison_complete": row.savings_estimated_turns == row.turns}) - for row in recorded_rows - ) - rows: Final = tuple( - row.model_copy( - update={ - "savings_comparison_complete": row.savings_comparison_complete - and isclose( - row.saved_spend, - row.savings_estimated_saved_spend, - rel_tol=1e-9, - abs_tol=1e-9, - ), - } - ) - for row in covered_rows - ) + rows: Final = tuple(_compared_row(row) for row in _SESSION_AGG_ROWS.validate_python(raw_rows or ())) groups: Final = ( *(_benchmark_group(row) for row in rows), *_idle_router_groups(llm_router, frozenset((row.router_name, row.router_type) for row in rows)), @@ -962,43 +931,16 @@ async def get_auto_router_session( if prisma_client is None: raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) - recorded: Final = await AutoRouterSessionRepository(prisma_client).find_latest_for_key( + row: Final = await AutoRouterSessionRepository(prisma_client).find_latest_for_key( user_api_key_dict.api_key, bounded_session_id(session_id) ) - if recorded is None: + if row is None: raise HTTPException( status_code=404, detail=f"No auto-routed turns recorded for session {session_id!r} under this key" ) - comparisons: Final = ( - await historical_session_comparisons( - prisma_client, - recorded.first_turn_at.isoformat(), - (recorded.last_turn_at + timedelta(microseconds=1)).isoformat(), - user_api_key_dict.api_key, - None, - bounded_session_id(session_id), - ) - if recorded.savings_estimated_turns < recorded.turns - else MappingProxyType({}) - ) - comparison: Final = comparisons.get((recorded.router_name, recorded.router_type)) - row: Final = ( - recorded.model_copy(update=comparison.coverage_fields(recorded.saved_spend, recorded.turns)) - if comparison - else recorded - ) - saved_spend, compared_baseline = _savings_cohort( - row.turns, - row.savings_estimated_turns, - row.savings_estimated_actual_spend, - row.savings_estimated_saved_spend, - row.saved_spend, - ) - baseline_spend: Final = ( - compared_baseline - if row.savings_estimated_turns == row.turns - or (comparison and comparison.complete and comparison.turns == row.turns) - else None + saved_spend, baseline_spend = _savings_cohort(row.turns, row.savings_estimated_turns, row.spend, row.saved_spend) + _, estimated_baseline_spend = _savings_cohort( + row.turns, row.savings_estimated_turns, row.savings_estimated_actual_spend, row.savings_estimated_saved_spend ) return AutoRouterSessionResponse( session_id=session_id, @@ -1010,8 +952,8 @@ async def get_auto_router_session( savings_estimated_turns=row.savings_estimated_turns, savings_estimated_actual_spend=row.savings_estimated_actual_spend, saved_spend=saved_spend, - baseline_spend=baseline_spend if row.savings_estimated_turns == row.turns else None, - savings_estimated_baseline_spend=baseline_spend, + baseline_spend=baseline_spend, + savings_estimated_baseline_spend=estimated_baseline_spend, baseline_model=row.baseline_model, baseline_models=row.baseline_models, ) diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index 67f7ed424e7..00083e01f54 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -218,23 +218,24 @@ class AutoRouterBenchmarkTotals(BaseModel): "subtotal recording, and zero for an empty window" ) savings_estimated_turns: int = Field( - description="Requests with a matching savings comparison, including historical recorded estimates" + description="Requests compared against the baseline: every request on complexity routers that recorded savings" ) savings_estimated_actual_spend: float = Field( - description="Actual spend, including classifier cost, for covered turns only" + description="Actual spend, including classifier cost, for the compared requests" ) savings_estimated_classifier_cost: float | None = Field( default=None, - description="Classifier cost included in the matching historical and newer savings comparison; " + description="Classifier cost included in the compared actual spend; " "null when classification costs for those requests are unavailable", ) saved_spend: float | None = Field( description="Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates" ) - baseline_spend: float | None = Field(description="Estimated single-model cost for covered turns only") - saved_pct: float | None = Field( - description="Total recorded savings over the matching historical and current baseline; null when costs are unavailable" + baseline_spend: float | None = Field( + description="Estimated single-model cost: compared actual spend plus recorded savings; " + "null when traffic has no recorded savings" ) + saved_pct: float | None = Field(description="Recorded savings over baseline_spend, as a percentage") saved_per_session: float | None = Field(description="Recorded savings per session, including historical estimates") cache: AutoRouterCacheStats @@ -266,20 +267,16 @@ class AutoRouterSessionResponse(BaseModel): turns: int = Field(description="Auto-routed turns the rollup has recorded for this session so far") last_model: str = Field(description="The deployment model the most recent turn was routed to") spend: float = Field(description="What the session's routed traffic actually cost, classifier calls included") - savings_estimated_turns: int = Field( - description="Requests with a matching savings comparison, including historical recorded estimates" - ) + savings_estimated_turns: int = Field(description="Requests whose savings estimate recorded its baseline cost") savings_estimated_actual_spend: float = Field( - description="Actual spend, including classifier cost, for covered turns only" + description="Actual spend, including classifier cost, for requests whose estimate recorded its baseline cost" ) saved_spend: float | None = Field( description="Recorded historical savings plus newer estimates, net of classifier cost" ) - baseline_spend: float | None = Field( - description="Estimated single-model cost; unavailable unless every turn is covered" - ) + baseline_spend: float | None = Field(description="Estimated single-model cost: spend plus recorded savings") savings_estimated_baseline_spend: float | None = Field( - description="Estimated single-model cost for covered turns only" + description="Estimated single-model cost for requests whose estimate recorded its baseline cost" ) baseline_model: str | None = Field( description="The savings baseline recorded by most session turns, including historical turns, recorded turn by " diff --git a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py index effb5ca83ed..2c648f309f6 100644 --- a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py +++ b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py @@ -6,7 +6,6 @@ tests/unit/proxy/db/test_autorouter_session_rollup.py. """ import asyncio -import json import time import uuid from datetime import datetime, timedelta, timezone @@ -25,10 +24,6 @@ from litellm.proxy.db.autorouter_session_rollup import ( flush_autorouter_turn_transactions, ) from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import SpendLogCleanup -from litellm.proxy.db.autorouter_savings_comparison import ( - HISTORICAL_SESSION_COMPARISONS_SQL, - SessionSavingsComparison, -) pytestmark = pytest.mark.asyncio(loop_scope="session") @@ -96,66 +91,6 @@ async def _row(db, key: str, session_id: str = "s1", router: str = "auto-1") -> return rows[0] -@pytest.mark.parametrize("historical_saved, damaged, user_id, split_sessions, current_classifier", [ - (29.5, None, None, False, 0.2), (29.5, None, "owner", False, 0.2), (0.0, None, None, False, 0.2), - (-3.0, None, None, False, 0.2), (29.5, "missing", None, False, 0.2), (29.5, "cost", None, False, 0.2), - (0.0, "missing", None, False, 0.2), (29.5, None, None, True, 0.2), (29.5, None, None, False, 0.0), -]) -async def test_historical_and_new_savings_compare_matching_costs_and_exclude_unknown_requests( - db: Prisma, historical_saved: float, damaged: str | None, user_id: str | None, split_sessions: bool, - current_classifier: float, -) -> None: - async with db.tx() as tx: - for table in ("LiteLLM_AutoRouterSession", "LiteLLM_AutoRouterUserSession", "LiteLLM_SpendLogs"): - await tx.execute_raw(f'CREATE TEMP TABLE "{table}" (LIKE public."{table}" INCLUDING ALL) ON COMMIT DROP') - for name, spend, saved, classifier, estimated in ( - ("historical", 9.0, historical_saved, 0.1, False), - ("current", 1.0, 0.5, current_classifier, True), - ("unknown", 99.0, 0.0, 3.0, False), - ): - session_id: Final = "s2" if split_sessions and name == "current" else "s1" - await _turn(tx, "key", "model", T0, spend=spend, saved=saved, classifier_cost=classifier, - estimated=estimated, session_id=session_id) - metadata: Final = { - "routing_decision": {"router_model_name": "auto-1", **({"classifier_cost": classifier} if classifier else {})}, - "autorouter_savings": saved if name != "unknown" else None, - **({"autorouter_savings_estimate": { - "version": 3, "status": "estimated" if estimated else "unknown", - }} if name != "historical" else {}), - } - await tx.execute_raw('''INSERT INTO "LiteLLM_SpendLogs" - (request_id,api_key,session_id,model,"user","startTime","endTime",call_type, - spend,prompt_tokens,completion_tokens,status,metadata) - VALUES ($1,'key',$5,'model','owner',$2::timestamp,$2::timestamp,'acompletion', - $3::float8,100,0,'success',$4::jsonb) - ''', name, T0.isoformat(), spend - classifier, json.dumps(metadata), session_id) - await tx.execute_raw('''INSERT INTO "LiteLLM_AutoRouterUserSession" - (user_id,api_key,session_id,router_name,router_type,first_turn_at,last_turn_at,last_model, - turns,total_tokens,spend,saved_spend,savings_estimated_turns,savings_estimated_actual_spend, - savings_estimated_saved_spend) - SELECT 'owner',api_key,session_id,router_name,router_type,first_turn_at,last_turn_at,last_model, - turns,total_tokens,spend,saved_spend,savings_estimated_turns,savings_estimated_actual_spend, - savings_estimated_saved_spend FROM "LiteLLM_AutoRouterSession" - ''') - if damaged == "missing": - await tx.execute_raw('DELETE FROM "LiteLLM_SpendLogs" WHERE request_id = \'historical\'') - elif damaged == "cost": - await tx.execute_raw('UPDATE "LiteLLM_SpendLogs" SET spend = 1 WHERE request_id = \'historical\'') - rows: Final = await tx.query_raw( - HISTORICAL_SESSION_COMPARISONS_SQL, "2026-08-01", "2026-08-02", "key", user_id, None, - ) - comparison: Final = SessionSavingsComparison.model_validate(rows[0]) - assert comparison.saved_spend == historical_saved + 0.5 - assert comparison.complete is (damaged is None) - assert comparison.classifier_cost == (pytest.approx(0.1 + current_classifier) if damaged is None else None) - assert comparison.coverage_fields(historical_saved + 0.5, 4) == {} - assert comparison.coverage_fields(historical_saved + 0.5, 3) == ({ - "savings_estimated_turns": 2, - "savings_estimated_actual_spend": 10.0, - "savings_estimated_saved_spend": historical_saved + 0.5, - } if damaged is None else {}) - - async def test_every_turn_lands_in_exactly_one_bucket(db): key = f"k-{uuid.uuid4()}" await _turn(db, key, "A", T0, ttl=300) diff --git a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py index 385b2b1cc5b..2cbba9da8b3 100644 --- a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py @@ -721,26 +721,60 @@ class TestAutoRouterBenchmarks: assert totals.saved_pct == -100.0 assert totals.classifier_cost == 0.4 + @pytest.mark.asyncio @pytest.mark.parametrize("estimated_turns", [0, 4]) - def test_recorded_savings_survive_when_historical_comparison_costs_are_missing(self, estimated_turns: int) -> None: - from litellm.proxy.management_endpoints.auto_router_endpoints import _benchmark_totals - + async def test_historical_savings_without_recorded_baselines_compare_against_all_spend( + self, estimated_turns: int, monkeypatch: pytest.MonkeyPatch + ) -> None: row: Final = self.ROW.model_copy( update={ "savings_estimated_turns": estimated_turns, "savings_estimated_actual_spend": 2.0 if estimated_turns else 0.0, + "savings_estimated_classifier_cost": None, "savings_estimated_saved_spend": -0.5 if estimated_turns else 0.0, } ) - totals: Final = _benchmark_totals(row) - assert totals.spend == 10.0 - assert totals.savings_estimated_turns == estimated_turns - assert totals.saved_spend == 30.0 - assert totals.baseline_spend is None - assert totals.savings_estimated_classifier_cost is None - assert totals.saved_pct is None + response: Final = await self._benchmarks(monkeypatch, rows=[row.model_dump()], model_list=[]) + assert response.groups[0].model_dump(exclude={"router_name", "router_type", "tier_turns"}) == ( + response.totals.model_dump() + ) + totals: Final = response.totals + assert (totals.spend, totals.saved_spend, totals.baseline_spend, totals.saved_pct) == (10.0, 30.0, 40.0, 75.0) + assert (totals.savings_estimated_turns, totals.savings_estimated_actual_spend) == (40, 10.0) + assert totals.savings_estimated_classifier_cost == 0.4 assert totals.saved_per_session == 7.5 + @pytest.mark.asyncio + @pytest.mark.parametrize("router_type, saved", [("adaptive", 0.0), ("quality", 0.0), ("quality", 2.0)]) + async def test_only_complexity_routers_enter_the_compared_totals( + self, router_type: str, saved: float, monkeypatch: pytest.MonkeyPatch + ) -> None: + adaptive: Final = self.ROW.model_copy( + update={ + "router_name": f"{router_type}-auto", + "router_type": router_type, + "turns": 10, + "spend": 3.0, + "saved_spend": saved, + "savings_estimated_turns": 0, + "savings_estimated_actual_spend": 0.0, + "savings_estimated_saved_spend": 0.0, + "classifier_cost": 0.0, + "classifier_cost_recorded_turns": 10, + } + ) + response: Final = await self._benchmarks( + monkeypatch, rows=[self.ROW.model_dump(), adaptive.model_dump()], model_list=[] + ) + unbaselined: Final = response.groups[1] + assert (unbaselined.saved_spend, unbaselined.baseline_spend, unbaselined.saved_pct) == (None, None, None) + assert (unbaselined.savings_estimated_turns, unbaselined.savings_estimated_classifier_cost) == (0, 0.0) + totals: Final = response.totals + assert (totals.turns, totals.spend) == (50, 13.0) + assert (totals.savings_estimated_turns, totals.savings_estimated_actual_spend) == (40, 10.0) + assert (totals.saved_spend, totals.baseline_spend, totals.saved_pct) == (30.0, 40.0, 75.0) + assert totals.savings_estimated_classifier_cost == 0.4 + def test_an_empty_window_folds_to_zeros(self): from litellm.proxy.management_endpoints.auto_router_endpoints import ( _benchmark_totals, @@ -1138,8 +1172,10 @@ class TestAutoRouterSession: "saved_spend": 0.24, "savings_estimated_turns": 3 if estimated else 0, "savings_estimated_actual_spend": 0.14 if estimated else 0.0, - "baseline_spend": pytest.approx(0.38) if turns == 3 else None, - "savings_estimated_baseline_spend": pytest.approx(0.38) if turns == 3 else None, + "baseline_spend": pytest.approx(spend + 0.24), + "savings_estimated_baseline_spend": ( + pytest.approx(0.38 if turns == 3 else 0.10) if estimated else None + ), "baseline_model": "anthropic/claude-opus-5", "baseline_models": {"anthropic/claude-opus-5": 3}, } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx index 5e8533c8b82..b4b34dbfaf3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx @@ -70,6 +70,7 @@ const totals = (overrides: Partial = {}): Totals => ({ spend: 359.86, savings_estimated_turns: overrides.turns ?? 3073, savings_estimated_actual_spend: overrides.spend ?? 359.86, + savings_estimated_classifier_cost: overrides.classifier_cost === undefined ? 6.146 : overrides.classifier_cost, classifier_cost: 6.146, saved_spend: 2174.59, baseline_spend: 2534.45, @@ -104,6 +105,7 @@ const zeroTotals: Totals = { spend: 0, savings_estimated_turns: 0, savings_estimated_actual_spend: 0, + savings_estimated_classifier_cost: 0, classifier_cost: 0, saved_spend: 0, baseline_spend: 0, @@ -159,11 +161,10 @@ describe("AutoRouterBenchmarksTab", () => { it.each([ { estimatedTurns: 0, actual: 0, saved: null, pct: null }, - { estimatedTurns: 0, actual: 0, saved: 30, pct: null }, { estimatedTurns: 10, actual: 2, saved: -0.5, pct: -33.3 }, { estimatedTurns: 10, actual: 2, saved: 0, pct: 0 }, - { estimatedTurns: 40, actual: 10, saved: 30, pct: 75 }, - ])("compares matching old and new requests with savings $saved", ({ estimatedTurns, actual, saved, pct }) => { + { estimatedTurns: 3073, actual: 10, saved: 30, pct: 75 }, + ])("compares the requests on routers that recorded savings $saved", ({ estimatedTurns, actual, saved, pct }) => { const comparison = { spend: actual + 99, savings_estimated_turns: estimatedTurns, @@ -189,18 +190,17 @@ describe("AutoRouterBenchmarksTab", () => { ] : ["Unavailable", "Unavailable", "Unavailable", "Unavailable"], ); - expect(screen.queryByText("Actual spend on covered turns")).not.toBeInTheDocument(); + expect(screen.queryByText(/Matching cost details are unavailable/)).not.toBeInTheDocument(); expect(screen.getByLabelText("question-circle")).toBeInTheDocument(); - if (estimatedTurns) { - expect(screen.getByText(`Savings based on ${estimatedTurns} of 3,073 requests`)).toBeInTheDocument(); - const sign = pct && pct > 0 ? "-" : "+"; - const badge = pct === 0 ? "0%" : `${sign}${Math.abs(pct ?? 0).toFixed(0)}%`; + const partial = estimatedTurns > 0 && estimatedTurns < 3073; + expect(screen.queryByText(/adaptive and quality routers are excluded/) != null).toBe(partial); + if (partial) { + expect(screen.getByText(/Compared on 10 of 3,073 requests/)).toBeInTheDocument(); + } + if (pct != null) { + const sign = pct > 0 ? "-" : "+"; + const badge = pct === 0 ? "0%" : `${sign}${Math.abs(pct).toFixed(0)}%`; expect(screen.getByText(badge)).toBeInTheDocument(); - } else if (saved != null) { - expect(screen.getByText("$30.00")).toBeInTheDocument(); - expect( - screen.getByText("Historical savings are included. Matching cost details are unavailable."), - ).toBeInTheDocument(); } }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx index f532c2e4650..24a97587e32 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx @@ -76,10 +76,8 @@ const SpendRow: React.FC<{ label: string; value: string; hint?: string; subdued? const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => { const stats = view.stats; const cheaper = stats.saved_pct != null && stats.saved_pct >= 0; - const completeCoverage = stats.savings_estimated_turns === stats.turns; - const coveredClassifierCost = - stats.savings_estimated_classifier_cost ?? (completeCoverage ? stats.classifier_cost : null); - const classifierCost = stats.baseline_spend == null ? null : coveredClassifierCost; + const classifierCost = stats.baseline_spend == null ? null : stats.savings_estimated_classifier_cost ?? null; + const comparedAll = stats.savings_estimated_turns === stats.turns; return (
@@ -101,15 +99,10 @@ const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => { )}
- {stats.baseline_spend != null && !completeCoverage && ( + {stats.baseline_spend != null && !comparedAll && (

- Savings based on {stats.savings_estimated_turns.toLocaleString()} of {stats.turns.toLocaleString()}{" "} - requests -

- )} - {stats.saved_spend != null && stats.baseline_spend == null && ( -

- Historical savings are included. Matching cost details are unavailable. + Compared on {stats.savings_estimated_turns.toLocaleString()} of {stats.turns.toLocaleString()} requests; + adaptive and quality routers are excluded

)} @@ -118,7 +111,7 @@ const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => {
= ({ isPending, error, data,

- Savings, actual spend, and baseline compare the same historical and newer requests with recorded estimates, - including zero or negative savings. Requests without estimates are excluded. Savings are net of recorded LLM - classification cost. If historical cost details are unavailable, recorded savings remain visible without a - baseline or percentage. The range counts whole sessions that overlap it, so totals can differ from savings views - that group usage by UTC day. + Actual spend covers every request on complexity routers, including LLM classification cost. Baseline is actual + spend plus recorded savings, so savings can be zero or negative. The range counts whole sessions that overlap + it, so totals can differ from savings views that group usage by UTC day.

diff --git a/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx b/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx index 807bd2f4f15..5c7b9b28789 100644 --- a/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx @@ -34,6 +34,7 @@ const stats = { spend: 1.25, savings_estimated_turns: 4, savings_estimated_actual_spend: 1.25, + savings_estimated_classifier_cost: 0.25, classifier_cost: 0.25, saved_spend: 8.75, baseline_spend: 10, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index ab164a61ca1..c6bb9be41df 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -1263,8 +1263,8 @@ export interface paths { * @description Benchmarks for the auto-router dashboard: session shape, savings against the configured * baseline, and prompt-caching behaviour bucketed by what the router did. * - * Reads session rollups folded once per request at spend-write time, with bounded - * retained-log recovery for historical comparisons. A user filter selects only turns attributed to that + * Reads session rollups folded once per request at spend-write time, so this endpoint + * never scans LiteLLM_SpendLogs. A user filter selects only turns attributed to that * internal user when written; older key-only history remains outside user views. A session * is in the window when it overlaps it: its last turn is on or after start_date and its first turn is on or before * end_date. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is @@ -25464,7 +25464,7 @@ export interface components { avg_turns_per_session: number; /** * Baseline Spend - * @description Estimated single-model cost for covered turns only + * @description Estimated single-model cost: compared actual spend plus recorded savings; null when traffic has no recorded savings */ baseline_spend: number | null; cache: components["schemas"]["AutoRouterCacheStats"]; @@ -25485,7 +25485,7 @@ export interface components { router_type: string; /** * Saved Pct - * @description Total recorded savings over the matching historical and current baseline; null when costs are unavailable + * @description Recorded savings over baseline_spend, as a percentage */ saved_pct: number | null; /** @@ -25500,17 +25500,17 @@ export interface components { saved_spend: number | null; /** * Savings Estimated Actual Spend - * @description Actual spend, including classifier cost, for covered turns only + * @description Actual spend, including classifier cost, for the compared requests */ savings_estimated_actual_spend: number; /** * Savings Estimated Classifier Cost - * @description Classifier cost included in the matching historical and newer savings comparison; null when classification costs for those requests are unavailable + * @description Classifier cost included in the compared actual spend; null when classification costs for those requests are unavailable */ savings_estimated_classifier_cost?: number | null; /** * Savings Estimated Turns - * @description Requests with a matching savings comparison, including historical recorded estimates + * @description Requests compared against the baseline: every request on complexity routers that recorded savings */ savings_estimated_turns: number; /** Sessions */ @@ -25543,7 +25543,7 @@ export interface components { avg_turns_per_session: number; /** * Baseline Spend - * @description Estimated single-model cost for covered turns only + * @description Estimated single-model cost: compared actual spend plus recorded savings; null when traffic has no recorded savings */ baseline_spend: number | null; cache: components["schemas"]["AutoRouterCacheStats"]; @@ -25554,7 +25554,7 @@ export interface components { classifier_cost: number | null; /** * Saved Pct - * @description Total recorded savings over the matching historical and current baseline; null when costs are unavailable + * @description Recorded savings over baseline_spend, as a percentage */ saved_pct: number | null; /** @@ -25569,17 +25569,17 @@ export interface components { saved_spend: number | null; /** * Savings Estimated Actual Spend - * @description Actual spend, including classifier cost, for covered turns only + * @description Actual spend, including classifier cost, for the compared requests */ savings_estimated_actual_spend: number; /** * Savings Estimated Classifier Cost - * @description Classifier cost included in the matching historical and newer savings comparison; null when classification costs for those requests are unavailable + * @description Classifier cost included in the compared actual spend; null when classification costs for those requests are unavailable */ savings_estimated_classifier_cost?: number | null; /** * Savings Estimated Turns - * @description Requests with a matching savings comparison, including historical recorded estimates + * @description Requests compared against the baseline: every request on complexity routers that recorded savings */ savings_estimated_turns: number; /** Sessions */ @@ -25868,7 +25868,7 @@ export interface components { }; /** * Baseline Spend - * @description Estimated single-model cost; unavailable unless every turn is covered + * @description Estimated single-model cost: spend plus recorded savings */ baseline_spend: number | null; /** @@ -25893,17 +25893,17 @@ export interface components { saved_spend: number | null; /** * Savings Estimated Actual Spend - * @description Actual spend, including classifier cost, for covered turns only + * @description Actual spend, including classifier cost, for requests whose estimate recorded its baseline cost */ savings_estimated_actual_spend: number; /** * Savings Estimated Baseline Spend - * @description Estimated single-model cost for covered turns only + * @description Estimated single-model cost for requests whose estimate recorded its baseline cost */ savings_estimated_baseline_spend: number | null; /** * Savings Estimated Turns - * @description Requests with a matching savings comparison, including historical recorded estimates + * @description Requests whose savings estimate recorded its baseline cost */ savings_estimated_turns: number; /** Session Id */ From be67fce26a19669e0696082ec3d6395cbdbcc703 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Thu, 1 Oct 2026 13:45:32 -0700 Subject: [PATCH 010/203] refactor(proxy): inject tracing receiver and access context (#44035) * refactor(proxy): inject tracing receiver and access context * refactor(proxy): own tracing resources through FastAPI lifespan * test(proxy): pass tracing dependency in Lens lifecycle * refactor(proxy): stop tracing logger cooperatively * refactor(proxy): derive tracing permissions in one place * refactor(proxy): compose application lifespan state * refactor(proxy): give Lens tracing storage directly * refactor(tracing): name shared ClickHouse storage explicitly * refactor(tracing): extract shared ClickHouse storage crate * test(proxy): isolate db push timeout from Lens safety check * fix(tracing): drain spend retries during shutdown --- litellm-rust/Cargo.lock | 17 +- litellm-rust/Cargo.toml | 1 + litellm-rust/crates/python-bridge/Cargo.toml | 1 + .../crates/python-bridge/src/routes/traces.rs | 26 +- .../crates/storage-clickhouse/Cargo.toml | 20 + .../crates/storage-clickhouse/README.md | 5 + .../crates/storage-clickhouse/src/error.rs | 29 ++ .../crates/storage-clickhouse/src/insert.rs | 70 ++++ .../crates/storage-clickhouse/src/lib.rs | 127 +++++++ .../crates/storage-clickhouse/src/read.rs | 113 ++++++ .../storage-clickhouse/tests/connection.rs | 34 ++ .../storage-clickhouse/tests/transport.rs | 34 ++ litellm-rust/crates/traces/AGENTS.md | 2 +- litellm-rust/crates/traces/Cargo.toml | 2 +- litellm-rust/crates/traces/src/error.rs | 30 -- litellm-rust/crates/traces/src/insert.rs | 62 +-- litellm-rust/crates/traces/src/lib.rs | 84 +---- litellm-rust/crates/traces/src/sql.rs | 113 +----- litellm-rust/crates/traces/tests/queries.rs | 11 - .../clickhouse/clickhouse_batch_logger.py | 27 +- litellm/integrations/clickhouse/schema.py | 4 +- litellm/proxy/_types.py | 6 + litellm/proxy/lens/endpoints.py | 49 ++- litellm/proxy/proxy_server.py | 156 ++++---- litellm/proxy/tracing_endpoints.py | 86 +++-- litellm/proxy/tracing_runtime.py | 66 ++++ litellm/rust_bridge/traces.py | 2 +- litellm/tracing/AGENTS.md | 4 +- litellm/tracing/receiver.py | 10 +- litellm/tracing/store.py | 6 +- tests/proxy_behavior/lens/test_lifecycle.py | 6 +- .../test_prisma_toolchain.py | 2 +- .../test_clickhouse_batch_logger.py | 65 +++- tests/test_litellm/tracing/test_store.py | 12 +- .../proxy/proxy_server/test_proxy_config.py | 74 ++-- tests/unit/proxy/test_tracing_endpoints.py | 355 ++++++++++++++++-- 36 files changed, 1152 insertions(+), 559 deletions(-) create mode 100644 litellm-rust/crates/storage-clickhouse/Cargo.toml create mode 100644 litellm-rust/crates/storage-clickhouse/README.md create mode 100644 litellm-rust/crates/storage-clickhouse/src/error.rs create mode 100644 litellm-rust/crates/storage-clickhouse/src/insert.rs create mode 100644 litellm-rust/crates/storage-clickhouse/src/lib.rs create mode 100644 litellm-rust/crates/storage-clickhouse/src/read.rs create mode 100644 litellm-rust/crates/storage-clickhouse/tests/connection.rs create mode 100644 litellm-rust/crates/storage-clickhouse/tests/transport.rs delete mode 100644 litellm-rust/crates/traces/tests/queries.rs create mode 100644 litellm/proxy/tracing_runtime.py diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 0ad05d99e76..57e9e803e4a 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -4086,6 +4086,7 @@ dependencies = [ "litellm-secrets", "litellm-secrets-aws", "litellm-secrets-types", + "litellm-storage-clickhouse", "litellm-token-counter", "litellm-traces", "litellm-tracing", @@ -4288,6 +4289,20 @@ dependencies = [ "veil", ] +[[package]] +name = "litellm-storage-clickhouse" +version = "0.1.0" +dependencies = [ + "flate2", + "litellm-http", + "rstest", + "serde", + "serde_json", + "thiserror 2.0.19", + "tokio", + "url", +] + [[package]] name = "litellm-testkit" version = "0.1.0" @@ -4371,6 +4386,7 @@ dependencies = [ "base64 0.22.1", "flate2", "litellm-http", + "litellm-storage-clickhouse", "opentelemetry-proto", "prost", "rstest", @@ -4381,7 +4397,6 @@ dependencies = [ "thiserror 2.0.19", "time", "tokio", - "url", ] [[package]] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 257a47268e4..450253ea768 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -13,6 +13,7 @@ litellm-config = { path = "crates/config" } litellm-router = { path = "crates/router" } litellm-tracing = { path = "crates/tracing" } litellm-traces = { path = "crates/traces" } +litellm-storage-clickhouse = { path = "crates/storage-clickhouse" } litellm-core = { path = "crates/core" } litellm-gateway-mcp = { path = "crates/gateway-mcp" } litellm-gateway = { path = "crates/gateway" } diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 99c95632bb3..a1d1f63d6f3 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -22,6 +22,7 @@ tiktoken = ["litellm-token-counter/tiktoken"] fancy-regex.workspace = true litellm-tracing.workspace = true litellm-traces.workspace = true +litellm-storage-clickhouse.workspace = true litellm-host.workspace = true bytes.workspace = true futures-util.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 2e7a6b178a8..47c924f4842 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -1,7 +1,8 @@ use std::collections::BTreeMap; use litellm_http::ClientVariant; -use litellm_traces::{Connection, Error, InsertTable, Parameter, ReadQuery}; +use litellm_storage_clickhouse::Storage; +use litellm_traces::{Error, InsertTable, Parameter, ReadQuery}; use pyo3::{ exceptions::{PyOverflowError, PyRuntimeError, PyValueError}, prelude::*, @@ -27,9 +28,7 @@ fn map_error(error: Error) -> PyErr { #[pyclass] pub struct NativeTraceStorage { - database: String, - writer: Connection, - reader: Option, + storage: Storage, } #[pymethods] @@ -39,12 +38,7 @@ impl NativeTraceStorage { fn new(database: String, url: &str, reader_url: Option<&str>) -> PyResult { litellm_traces::schema_statements(&database, 1, 1).map_err(map_error)?; Ok(Self { - writer: Connection::writer(url).map_err(map_error)?, - reader: reader_url - .map(|value| Connection::reader(value, &database)) - .transpose() - .map_err(map_error)?, - database, + storage: Storage::new(database, url, reader_url).map_err(map_error)?, }) } @@ -55,8 +49,8 @@ impl NativeTraceStorage { spend_log_retention_days: u32, ) -> PyResult> { let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; - let connection = self.writer.clone(); - let database = self.database.clone(); + let connection = self.storage.writer().clone(); + let database = self.storage.database().to_owned(); crate::execution::run_async( py, async move { @@ -83,8 +77,8 @@ impl NativeTraceStorage { ) -> PyResult> { let table = InsertTable::parse(table).map_err(map_error)?; let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; - let connection = self.writer.clone(); - let database = self.database.clone(); + let connection = self.storage.writer().clone(); + let database = self.storage.database().to_owned(); crate::execution::run_async( py, async move { @@ -104,7 +98,7 @@ impl NativeTraceStorage { >, ) -> PyResult> { let query = litellm_traces::LensQuery::parse(name).map_err(map_error)?; - let connection = self.reader.clone().ok_or_else(|| { + let connection = self.storage.reader().cloned().ok_or_else(|| { PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL") })?; let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; @@ -127,7 +121,7 @@ impl NativeTraceStorage { >, ) -> PyResult> { let query = ReadQuery::parse(query).map_err(map_error)?; - let connection = self.reader.clone().ok_or_else(|| { + let connection = self.storage.reader().cloned().ok_or_else(|| { PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL") })?; let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; diff --git a/litellm-rust/crates/storage-clickhouse/Cargo.toml b/litellm-rust/crates/storage-clickhouse/Cargo.toml new file mode 100644 index 00000000000..f7f85c0dd8d --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/Cargo.toml @@ -0,0 +1,20 @@ +[package] +name = "litellm-storage-clickhouse" +version = "0.1.0" +description = "Shared ClickHouse connection and HTTP storage for LiteLLM features" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +flate2.workspace = true +litellm-http.workspace = true +serde.workspace = true +serde_json.workspace = true +thiserror.workspace = true +url.workspace = true + +[dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } +rstest.workspace = true +tokio.workspace = true diff --git a/litellm-rust/crates/storage-clickhouse/README.md b/litellm-rust/crates/storage-clickhouse/README.md new file mode 100644 index 00000000000..7c4f86e4589 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/README.md @@ -0,0 +1,5 @@ +# ClickHouse storage + +`litellm-storage-clickhouse` exports `Storage`, a shared writer connection and optional reader connection for one ClickHouse database. It also exports bounded HTTP read and insert execution + +The crate has no trace tables, OTLP types, or named trace queries. `litellm-traces` supplies those rules and uses this storage for both trace rows and spend rows diff --git a/litellm-rust/crates/storage-clickhouse/src/error.rs b/litellm-rust/crates/storage-clickhouse/src/error.rs new file mode 100644 index 00000000000..d283fb9021e --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/src/error.rs @@ -0,0 +1,29 @@ +#[derive(Debug, thiserror::Error)] +pub enum Error { + #[error("invalid ClickHouse insert row")] + InvalidRow, + #[error("invalid ClickHouse insert table")] + InvalidTable, + #[error("invalid ClickHouse HTTP URL")] + InvalidUrl, + #[error("database must be a nonempty SQL identifier and retention must be positive")] + InvalidSchema, + #[error("SQL query must not be empty")] + EmptySql, + #[error("unknown ClickHouse read query")] + InvalidQuery, + #[error("ClickHouse query failed with HTTP status {0}")] + QueryFailed(u16), + #[error("ClickHouse insert failed with HTTP status {0}")] + InsertFailed(u16), + #[error("ClickHouse insert exceeds the encoded size limit")] + InsertTooLarge, + #[error("ClickHouse schema setup failed with HTTP status {0}")] + SchemaFailed(u16), + #[error("ClickHouse query exceeded the response size limit")] + ResponseTooLarge, + #[error("ClickHouse returned an invalid or failed JSON query response")] + InvalidResponse, + #[error("ClickHouse query transport failed")] + Transport, +} diff --git a/litellm-rust/crates/storage-clickhouse/src/insert.rs b/litellm-rust/crates/storage-clickhouse/src/insert.rs new file mode 100644 index 00000000000..3be73b086a0 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/src/insert.rs @@ -0,0 +1,70 @@ +use std::{io::Write, time::Duration}; + +use flate2::{Compression, write::GzEncoder}; +use litellm_http::Client; + +use crate::{Connection, Error, valid_identifier}; + +const INSERT_TIMEOUT: Duration = Duration::from_secs(30); + +pub async fn insert_encoded_rows( + client: &Client, + connection: &Connection, + database: &str, + table: &str, + token: &str, + encoded: &str, +) -> Result<(), Error> { + if !valid_identifier(database) { + return Err(Error::InvalidSchema); + } + if !valid_identifier(table) { + return Err(Error::InvalidTable); + } + let mut encoder = GzEncoder::new(Vec::new(), Compression::default()); + encoder + .write_all(encoded.as_bytes()) + .map_err(|_| Error::InvalidRow)?; + let body = encoder.finish().map_err(|_| Error::InvalidRow)?; + let mut url = connection.url().clone(); + let existing_pairs: Vec<(String, String)> = url + .query_pairs() + .filter(|(key, _)| { + !matches!( + key.as_ref(), + "query" + | "async_insert" + | "async_insert_deduplicate" + | "wait_for_async_insert" + | "input_format_skip_unknown_fields" + | "date_time_input_format" + ) + }) + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect(); + url.query_pairs_mut() + .clear() + .extend_pairs(existing_pairs) + .append_pair( + "query", + &format!("INSERT INTO `{database}`.{} FORMAT JSONEachRow", table), + ) + .append_pair("insert_deduplication_token", token) + .append_pair("async_insert", "1") + .append_pair("async_insert_deduplicate", "1") + .append_pair("wait_for_async_insert", "1") + .append_pair("input_format_skip_unknown_fields", "0") + .append_pair("date_time_input_format", "best_effort"); + let response = client + .post(url) + .timeout(INSERT_TIMEOUT) + .header("Content-Encoding", "gzip") + .body(body) + .send() + .await + .map_err(|_| Error::Transport)?; + if !response.status().is_success() { + return Err(Error::InsertFailed(response.status().as_u16())); + } + Ok(()) +} diff --git a/litellm-rust/crates/storage-clickhouse/src/lib.rs b/litellm-rust/crates/storage-clickhouse/src/lib.rs new file mode 100644 index 00000000000..f2b34eddbf8 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/src/lib.rs @@ -0,0 +1,127 @@ +mod error; +mod insert; +mod read; + +pub use error::Error; +pub use insert::insert_encoded_rows; +pub use read::{Parameter, execute_read}; +use url::Url; + +#[derive(Clone)] +pub struct Connection { + url: Url, +} + +impl Connection { + pub fn parse(value: &str) -> Result { + let url = Url::parse(value).map_err(|_| Error::InvalidUrl)?; + if !matches!(url.scheme(), "http" | "https") || url.host().is_none() { + return Err(Error::InvalidUrl); + } + Ok(Self { url }) + } + + pub fn configured( + url: &str, + database: &str, + user: &str, + password: &str, + ) -> Result { + let mut connection = Self::parse(url)?; + connection + .url + .set_username(user) + .map_err(|_| Error::InvalidUrl)?; + connection + .url + .set_password(Some(password)) + .map_err(|_| Error::InvalidUrl)?; + let pairs: Vec<_> = connection + .url + .query_pairs() + .filter(|(key, _)| !matches!(key.as_ref(), "database" | "user" | "password")) + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect(); + connection + .url + .query_pairs_mut() + .clear() + .extend_pairs(pairs) + .append_pair("database", database); + Ok(connection) + } + + pub fn writer(url: &str) -> Result { + let mut connection = Self::parse(url)?; + let pairs: Vec<_> = connection + .url + .query_pairs() + .filter(|(key, _)| !matches!(key.as_ref(), "database" | "readonly" | "query")) + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect(); + connection.url.query_pairs_mut().clear().extend_pairs(pairs); + Ok(connection) + } + + pub fn reader(url: &str, database: &str) -> Result { + let mut connection = Self::parse(url)?; + let pairs: Vec<_> = connection + .url + .query_pairs() + .filter(|(key, _)| key != "database") + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect(); + connection + .url + .query_pairs_mut() + .clear() + .extend_pairs(pairs) + .append_pair("database", database); + Ok(connection) + } + + pub fn url(&self) -> &Url { + &self.url + } +} + +#[derive(Clone)] +pub struct Storage { + database: String, + writer: Connection, + reader: Option, +} + +impl Storage { + pub fn new(database: String, url: &str, reader_url: Option<&str>) -> Result { + if !valid_identifier(&database) { + return Err(Error::InvalidSchema); + } + Ok(Self { + writer: Connection::writer(url)?, + reader: reader_url + .map(|value| Connection::reader(value, &database)) + .transpose()?, + database, + }) + } + + pub fn database(&self) -> &str { + &self.database + } + + pub fn writer(&self) -> &Connection { + &self.writer + } + + pub fn reader(&self) -> Option<&Connection> { + self.reader.as_ref() + } +} + +pub(crate) fn valid_identifier(value: &str) -> bool { + !value.is_empty() + && value + .bytes() + .all(|c| c.is_ascii_alphanumeric() || c == b'_') +} diff --git a/litellm-rust/crates/storage-clickhouse/src/read.rs b/litellm-rust/crates/storage-clickhouse/src/read.rs new file mode 100644 index 00000000000..99c6a5120f3 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/src/read.rs @@ -0,0 +1,113 @@ +use std::{collections::BTreeMap, time::Duration}; + +use litellm_http::Client; +use serde::Deserialize; + +use crate::{Connection, Error}; + +const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024; + +#[derive(Debug, Deserialize)] +#[serde(untagged)] +pub enum Parameter { + Text(String), + Integer(i64), + Strings(Vec), +} + +impl Parameter { + fn encoded(&self) -> String { + match self { + Self::Text(value) => escaped(value), + Self::Integer(value) => value.to_string(), + Self::Strings(values) => format!( + "[{}]", + values + .iter() + .map(|value| format!("'{}'", escaped(value).replace('\'', "\\'"))) + .collect::>() + .join(",") + ), + } + } +} + +fn escaped(value: &str) -> String { + value + .replace('\\', "\\\\") + .replace('\t', "\\t") + .replace('\n', "\\n") + .replace('\r', "\\r") + .replace('\0', "\\0") +} + +pub async fn execute_read( + client: &Client, + connection: &Connection, + sql: &str, + parameters: &BTreeMap, +) -> Result { + if sql.trim().is_empty() { + return Err(Error::EmptySql); + } + + let mut url = connection.url().clone(); + + let existing_pairs: Vec<(String, String)> = url + .query_pairs() + .filter(|(key, _)| { + !key.starts_with("param_") + && !matches!( + key.as_ref(), + "query" + | "readonly" + | "default_format" + | "max_result_rows" + | "result_overflow_mode" + | "max_execution_time" + | "wait_end_of_query" + ) + }) + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect(); + url.query_pairs_mut() + .clear() + .extend_pairs(existing_pairs) + .append_pair("readonly", "1") + .append_pair("max_result_rows", "1000") + .append_pair("result_overflow_mode", "throw") + .append_pair("max_execution_time", "10") + .append_pair("wait_end_of_query", "1") + .append_pair("default_format", "JSON"); + + url.query_pairs_mut().extend_pairs( + parameters + .iter() + .map(|(name, value)| (format!("param_{name}"), value.encoded())), + ); + + let request = client + .post(url) + .timeout(Duration::from_secs(15)) + .body(sql.to_owned()); + let mut response = request.send().await.map_err(|_| Error::Transport)?; + if !response.status().is_success() { + return Err(Error::QueryFailed(response.status().as_u16())); + } + + let mut body = Vec::new(); + while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? { + if body.len() + chunk.len() > MAX_RESPONSE_BYTES { + return Err(Error::ResponseTooLarge); + } + body.extend_from_slice(&chunk); + } + + let json: serde_json::Value = + serde_json::from_slice(&body).map_err(|_| Error::InvalidResponse)?; + if json.get("exception").is_some() || !json.get("data").is_some_and(serde_json::Value::is_array) + { + return Err(Error::InvalidResponse); + } + String::from_utf8(body).map_err(|_| Error::InvalidResponse) +} diff --git a/litellm-rust/crates/storage-clickhouse/tests/connection.rs b/litellm-rust/crates/storage-clickhouse/tests/connection.rs new file mode 100644 index 00000000000..0874b693249 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/tests/connection.rs @@ -0,0 +1,34 @@ +use litellm_storage_clickhouse::{Connection, Storage}; +use rstest::rstest; + +#[rstest] +#[case::http("http://localhost:8123", true)] +#[case::https("https://localhost:8443", true)] +#[case::tcp("tcp://localhost:9000", false)] +#[case::missing_host("http://", false)] +fn accepts_only_clickhouse_http_urls(#[case] value: &str, #[case] expected: bool) { + assert_eq!(Connection::parse(value).is_ok(), expected); +} + +#[rstest] +#[case::writer_only(None, false)] +#[case::separate_reader(Some("http://localhost:8124"), true)] +fn storage_exports_writer_and_optional_reader( + #[case] reader_url: Option<&str>, + #[case] has_reader: bool, +) { + let storage = Storage::new("litellm".to_owned(), "http://localhost:8123", reader_url) + .expect("valid ClickHouse URLs"); + + assert_eq!(storage.database(), "litellm"); + assert_eq!(storage.writer().url().host_str(), Some("localhost")); + assert_eq!(storage.writer().url().port(), Some(8123)); + assert_eq!(storage.reader().is_some(), has_reader); +} + +#[rstest] +#[case::empty("")] +#[case::injection("db; DROP DATABASE default")] +fn storage_rejects_invalid_database(#[case] database: &str) { + assert!(Storage::new(database.to_owned(), "http://localhost:8123", None).is_err()); +} diff --git a/litellm-rust/crates/storage-clickhouse/tests/transport.rs b/litellm-rust/crates/storage-clickhouse/tests/transport.rs new file mode 100644 index 00000000000..0c7ea233522 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/tests/transport.rs @@ -0,0 +1,34 @@ +use std::collections::BTreeMap; + +use litellm_http::Client; +use litellm_storage_clickhouse::{Connection, Error, execute_read, insert_encoded_rows}; +use rstest::rstest; + +#[rstest] +#[case::invalid_database("db; DROP DATABASE default", "spend_logs", true)] +#[case::invalid_table("litellm", "spend_logs; DROP TABLE otel_traces", false)] +#[tokio::test] +async fn insert_rejects_invalid_identifiers( + #[case] database: &str, + #[case] table: &str, + #[case] invalid_database: bool, +) { + let client = Client::no_redirect_for_test(); + let connection = Connection::writer("http://localhost:8123").expect("valid URL"); + let result = insert_encoded_rows(&client, &connection, database, table, "token", "{}").await; + + assert!(matches!(&result, Err(Error::InvalidSchema)) == invalid_database); + assert!(matches!(&result, Err(Error::InvalidTable)) == !invalid_database); +} + +#[rstest] +#[tokio::test] +async fn read_rejects_empty_sql() { + let client = Client::no_redirect_for_test(); + let connection = Connection::reader("http://localhost:8123", "litellm").expect("valid URL"); + + assert!(matches!( + execute_read(&client, &connection, " ", &BTreeMap::new()).await, + Err(Error::EmptySql) + )); +} diff --git a/litellm-rust/crates/traces/AGENTS.md b/litellm-rust/crates/traces/AGENTS.md index a5e2d4be53a..645e88dfae1 100644 --- a/litellm-rust/crates/traces/AGENTS.md +++ b/litellm-rust/crates/traces/AGENTS.md @@ -1,4 +1,4 @@ -- Rust owns OTLP wire decoding, ClickHouse schema, row encoding, named reads, connection validation and transport +- Keep OTLP decoding, trace schema, row encoding and named query selection here. Generic ClickHouse connections and HTTP execution belong in `litellm-storage-clickhouse` - Keep this crate independent of Python; PyO3 conversion and public Python exceptions belong in `python-bridge` - Keep the SQL migrations here as the only ClickHouse schema definition - Use typed query parameters and a dedicated SELECT-only reader with server-side limits diff --git a/litellm-rust/crates/traces/Cargo.toml b/litellm-rust/crates/traces/Cargo.toml index 7d5facaa71e..b7f6e6ae52e 100644 --- a/litellm-rust/crates/traces/Cargo.toml +++ b/litellm-rust/crates/traces/Cargo.toml @@ -12,11 +12,11 @@ opentelemetry-proto = { version = "0.33.0", default-features = false, features = prost = "0.14.4" time = { workspace = true, features = ["formatting"] } litellm-http.workspace = true +litellm-storage-clickhouse.workspace = true sha2.workspace = true serde.workspace = true serde_json.workspace = true thiserror.workspace = true -url.workspace = true [dev-dependencies] litellm-http = { workspace = true, features = ["test-support"] } diff --git a/litellm-rust/crates/traces/src/error.rs b/litellm-rust/crates/traces/src/error.rs index 4a4fdaa00f7..2ccfe0ea8d9 100644 --- a/litellm-rust/crates/traces/src/error.rs +++ b/litellm-rust/crates/traces/src/error.rs @@ -1,33 +1,3 @@ -#[derive(Debug, thiserror::Error)] -pub enum Error { - #[error("invalid ClickHouse insert row")] - InvalidRow, - #[error("invalid ClickHouse insert table")] - InvalidTable, - #[error("invalid ClickHouse HTTP URL")] - InvalidUrl, - #[error("database must be a nonempty SQL identifier and retention must be positive")] - InvalidSchema, - #[error("SQL query must not be empty")] - EmptySql, - #[error("unknown ClickHouse read query")] - InvalidQuery, - #[error("ClickHouse query failed with HTTP status {0}")] - QueryFailed(u16), - #[error("ClickHouse insert failed with HTTP status {0}")] - InsertFailed(u16), - #[error("ClickHouse insert exceeds the encoded size limit")] - InsertTooLarge, - #[error("ClickHouse schema setup failed with HTTP status {0}")] - SchemaFailed(u16), - #[error("ClickHouse query exceeded the response size limit")] - ResponseTooLarge, - #[error("ClickHouse returned an invalid or failed JSON query response")] - InvalidResponse, - #[error("ClickHouse query transport failed")] - Transport, -} - #[derive(Debug, thiserror::Error)] pub enum DecodeError { #[error("invalid OTLP trace payload")] diff --git a/litellm-rust/crates/traces/src/insert.rs b/litellm-rust/crates/traces/src/insert.rs index bbee66f6fa5..6d9a2cab813 100644 --- a/litellm-rust/crates/traces/src/insert.rs +++ b/litellm-rust/crates/traces/src/insert.rs @@ -1,6 +1,5 @@ -use std::{collections::BTreeMap, io::Write, time::Duration}; +use std::collections::BTreeMap; -use flate2::{Compression, write::GzEncoder}; use litellm_http::Client; use serde_json::Value; use sha2::{Digest, Sha256}; @@ -9,7 +8,6 @@ use time::{OffsetDateTime, format_description::well_known::Rfc3339}; use crate::{Connection, Error}; const MAX_INSERT_BYTES: usize = 64 * 1024 * 1024; -const INSERT_TIMEOUT: Duration = Duration::from_secs(30); pub enum InsertTable { OtelTraces, @@ -61,55 +59,15 @@ pub async fn insert_rows( }) .collect(); let encoded = encode_rows_with_limit(rows, MAX_INSERT_BYTES)?; - let mut encoder = GzEncoder::new(Vec::new(), Compression::default()); - encoder - .write_all(encoded.as_bytes()) - .map_err(|_| Error::InvalidRow)?; - let body = encoder.finish().map_err(|_| Error::InvalidRow)?; - let mut url = connection.url().clone(); - let existing_pairs: Vec<(String, String)> = url - .query_pairs() - .filter(|(key, _)| { - !matches!( - key.as_ref(), - "query" - | "async_insert" - | "async_insert_deduplicate" - | "wait_for_async_insert" - | "input_format_skip_unknown_fields" - | "date_time_input_format" - ) - }) - .map(|(key, value)| (key.into_owned(), value.into_owned())) - .collect(); - url.query_pairs_mut() - .clear() - .extend_pairs(existing_pairs) - .append_pair( - "query", - &format!( - "INSERT INTO `{database}`.{} FORMAT JSONEachRow", - table.name() - ), - ) - .append_pair("insert_deduplication_token", &token) - .append_pair("async_insert", "1") - .append_pair("async_insert_deduplicate", "1") - .append_pair("wait_for_async_insert", "1") - .append_pair("input_format_skip_unknown_fields", "0") - .append_pair("date_time_input_format", "best_effort"); - let response = client - .post(url) - .timeout(INSERT_TIMEOUT) - .header("Content-Encoding", "gzip") - .body(body) - .send() - .await - .map_err(|_| Error::Transport)?; - if !response.status().is_success() { - return Err(Error::InsertFailed(response.status().as_u16())); - } - Ok(()) + litellm_storage_clickhouse::insert_encoded_rows( + client, + connection, + database, + table.name(), + &token, + &encoded, + ) + .await } pub fn encode_rows(rows: Vec>) -> Result { diff --git a/litellm-rust/crates/traces/src/lib.rs b/litellm-rust/crates/traces/src/lib.rs index c37602cade4..f5defb36cc2 100644 --- a/litellm-rust/crates/traces/src/lib.rs +++ b/litellm-rust/crates/traces/src/lib.rs @@ -4,87 +4,9 @@ mod otlp; mod schema; mod sql; -pub use error::{DecodeError, Error}; +pub use error::DecodeError; pub use insert::{InsertTable, encode_rows, insert_rows}; +pub use litellm_storage_clickhouse::{Connection, Error, Parameter, execute_read}; pub use otlp::{DecodedSpan, decode_otlp}; pub use schema::{ensure_schema, schema_statements}; -pub use sql::{LensQuery, Parameter, ReadQuery, execute_named_read, execute_read}; -use url::Url; - -#[derive(Clone)] -pub struct Connection { - url: Url, -} - -impl Connection { - pub fn parse(value: &str) -> Result { - let url = Url::parse(value).map_err(|_| Error::InvalidUrl)?; - if !matches!(url.scheme(), "http" | "https") || url.host().is_none() { - return Err(Error::InvalidUrl); - } - Ok(Self { url }) - } - - pub fn configured( - url: &str, - database: &str, - user: &str, - password: &str, - ) -> Result { - let mut connection = Self::parse(url)?; - connection - .url - .set_username(user) - .map_err(|_| Error::InvalidUrl)?; - connection - .url - .set_password(Some(password)) - .map_err(|_| Error::InvalidUrl)?; - let pairs: Vec<_> = connection - .url - .query_pairs() - .filter(|(key, _)| !matches!(key.as_ref(), "database" | "user" | "password")) - .map(|(key, value)| (key.into_owned(), value.into_owned())) - .collect(); - connection - .url - .query_pairs_mut() - .clear() - .extend_pairs(pairs) - .append_pair("database", database); - Ok(connection) - } - - pub fn writer(url: &str) -> Result { - let mut connection = Self::parse(url)?; - let pairs: Vec<_> = connection - .url - .query_pairs() - .filter(|(key, _)| !matches!(key.as_ref(), "database" | "readonly" | "query")) - .map(|(key, value)| (key.into_owned(), value.into_owned())) - .collect(); - connection.url.query_pairs_mut().clear().extend_pairs(pairs); - Ok(connection) - } - - pub fn reader(url: &str, database: &str) -> Result { - let mut connection = Self::parse(url)?; - let pairs: Vec<_> = connection - .url - .query_pairs() - .filter(|(key, _)| key != "database") - .map(|(key, value)| (key.into_owned(), value.into_owned())) - .collect(); - connection - .url - .query_pairs_mut() - .clear() - .extend_pairs(pairs) - .append_pair("database", database); - Ok(connection) - } - - pub fn url(&self) -> &Url { - &self.url - } -} +pub use sql::{LensQuery, ReadQuery, execute_named_read}; diff --git a/litellm-rust/crates/traces/src/sql.rs b/litellm-rust/crates/traces/src/sql.rs index 8346e06cb71..9acb8de0a7a 100644 --- a/litellm-rust/crates/traces/src/sql.rs +++ b/litellm-rust/crates/traces/src/sql.rs @@ -1,12 +1,8 @@ -use std::{collections::BTreeMap, time::Duration}; - -use serde::Deserialize; +use std::collections::BTreeMap; use litellm_http::Client; -use crate::{Connection, Error}; - -const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024; +use crate::{Connection, Error, Parameter, execute_read}; pub enum ReadQuery { ListTraces, @@ -36,111 +32,6 @@ impl ReadQuery { } } -#[derive(Debug, Deserialize)] -#[serde(untagged)] -pub enum Parameter { - Text(String), - Integer(i64), - Strings(Vec), -} - -impl Parameter { - fn encoded(&self) -> String { - match self { - Self::Text(value) => escaped(value), - Self::Integer(value) => value.to_string(), - Self::Strings(values) => format!( - "[{}]", - values - .iter() - .map(|value| format!("'{}'", escaped(value).replace('\'', "\\'"))) - .collect::>() - .join(",") - ), - } - } -} - -fn escaped(value: &str) -> String { - value - .replace('\\', "\\\\") - .replace('\t', "\\t") - .replace('\n', "\\n") - .replace('\r', "\\r") - .replace('\0', "\\0") -} - -pub async fn execute_read( - client: &Client, - connection: &Connection, - sql: &str, - parameters: &BTreeMap, -) -> Result { - if sql.trim().is_empty() { - return Err(Error::EmptySql); - } - - let mut url = connection.url().clone(); - - let existing_pairs: Vec<(String, String)> = url - .query_pairs() - .filter(|(key, _)| { - !key.starts_with("param_") - && !matches!( - key.as_ref(), - "query" - | "readonly" - | "default_format" - | "max_result_rows" - | "result_overflow_mode" - | "max_execution_time" - | "wait_end_of_query" - ) - }) - .map(|(key, value)| (key.into_owned(), value.into_owned())) - .collect(); - url.query_pairs_mut() - .clear() - .extend_pairs(existing_pairs) - .append_pair("readonly", "1") - .append_pair("max_result_rows", "1000") - .append_pair("result_overflow_mode", "throw") - .append_pair("max_execution_time", "10") - .append_pair("wait_end_of_query", "1") - .append_pair("default_format", "JSON"); - - url.query_pairs_mut().extend_pairs( - parameters - .iter() - .map(|(name, value)| (format!("param_{name}"), value.encoded())), - ); - - let request = client - .post(url) - .timeout(Duration::from_secs(15)) - .body(sql.to_owned()); - let mut response = request.send().await.map_err(|_| Error::Transport)?; - if !response.status().is_success() { - return Err(Error::QueryFailed(response.status().as_u16())); - } - - let mut body = Vec::new(); - while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? { - if body.len() + chunk.len() > MAX_RESPONSE_BYTES { - return Err(Error::ResponseTooLarge); - } - body.extend_from_slice(&chunk); - } - - let json: serde_json::Value = - serde_json::from_slice(&body).map_err(|_| Error::InvalidResponse)?; - if json.get("exception").is_some() || !json.get("data").is_some_and(serde_json::Value::is_array) - { - return Err(Error::InvalidResponse); - } - String::from_utf8(body).map_err(|_| Error::InvalidResponse) -} - #[derive(Clone, Copy)] pub enum LensQuery { Sample, diff --git a/litellm-rust/crates/traces/tests/queries.rs b/litellm-rust/crates/traces/tests/queries.rs deleted file mode 100644 index 75dfe0adc19..00000000000 --- a/litellm-rust/crates/traces/tests/queries.rs +++ /dev/null @@ -1,11 +0,0 @@ -use litellm_traces::Connection; -use rstest::rstest; - -#[rstest] -#[case::http("http://localhost:8123", true)] -#[case::https("https://localhost:8443", true)] -#[case::tcp("tcp://localhost:9000", false)] -#[case::missing_host("http://", false)] -fn accepts_only_clickhouse_http_urls(#[case] value: &str, #[case] expected: bool) { - assert_eq!(Connection::parse(value).is_ok(), expected); -} diff --git a/litellm/integrations/clickhouse/clickhouse_batch_logger.py b/litellm/integrations/clickhouse/clickhouse_batch_logger.py index 81601ea2a78..ac782ffebb2 100644 --- a/litellm/integrations/clickhouse/clickhouse_batch_logger.py +++ b/litellm/integrations/clickhouse/clickhouse_batch_logger.py @@ -11,7 +11,8 @@ gzip JSONEachRow insert, either every `CLICKHOUSE_FLUSH_INTERVAL_SECONDS` or as import asyncio import os from collections.abc import Mapping, Sequence -from typing import Any, ClassVar +from contextlib import suppress +from typing import Any, ClassVar, Final from litellm._logging import verbose_logger from litellm.constants import ( @@ -21,11 +22,11 @@ from litellm.constants import ( CLICKHOUSE_MAX_RETRIES, ) from litellm.integrations.custom_batch_logger import CustomBatchLogger -from litellm.rust_bridge.traces import TraceStorage +from litellm.rust_bridge.traces import ClickHouseStorage -def clickhouse_storage_from_env() -> TraceStorage: - return TraceStorage( +def clickhouse_storage_from_env() -> ClickHouseStorage: + return ClickHouseStorage( database=os.getenv("CLICKHOUSE_DATABASE", "litellm"), url=os.getenv("CLICKHOUSE_URL", ""), ) @@ -34,7 +35,7 @@ def clickhouse_storage_from_env() -> TraceStorage: class ClickHouseBatchLogger(CustomBatchLogger): table: ClassVar[str] - def __init__(self, storage: TraceStorage | None = None) -> None: + def __init__(self, storage: ClickHouseStorage | None = None) -> None: self.storage = storage or clickhouse_storage_from_env() self.rows_written = 0 self.rows_dropped = 0 @@ -45,11 +46,27 @@ class ClickHouseBatchLogger(CustomBatchLogger): flush_interval=CLICKHOUSE_FLUSH_INTERVAL_SECONDS, ) self._flush_task: asyncio.Task[None] | None = None + self._stop: Final = asyncio.Event() def start(self) -> None: if self._flush_task is None or self._flush_task.done(): self._flush_task = asyncio.get_running_loop().create_task(self.periodic_flush()) + async def aclose(self) -> None: + self._stop.set() + if self._flush_task is not None: + await self._flush_task + while self.log_queue: + await self.flush_queue() + + async def periodic_flush(self) -> None: + while True: + with suppress(asyncio.TimeoutError): + await asyncio.wait_for(self._stop.wait(), timeout=self.flush_interval) + if self._stop.is_set(): + return + await self.flush_queue() + def is_full(self) -> bool: """Backpressure signal: producers should reject (429) instead of enqueueing.""" return len(self.log_queue) >= CLICKHOUSE_MAX_BUFFERED_ROWS diff --git a/litellm/integrations/clickhouse/schema.py b/litellm/integrations/clickhouse/schema.py index 6bec35c5630..5bf2b21cda5 100644 --- a/litellm/integrations/clickhouse/schema.py +++ b/litellm/integrations/clickhouse/schema.py @@ -1,11 +1,11 @@ from typing import Final -from litellm.rust_bridge.traces import TraceStorage +from litellm.rust_bridge.traces import ClickHouseStorage OTEL_TRACES_TABLE: Final = "otel_traces" AGENT_TRACES_BY_KEY_TABLE: Final = "agent_traces_by_key" SPEND_LOGS_TABLE: Final = "spend_logs" -async def ensure_schema(storage: TraceStorage, trace_retention_days: int, spend_log_retention_days: int) -> None: +async def ensure_schema(storage: ClickHouseStorage, trace_retention_days: int, spend_log_retention_days: int) -> None: await storage.ensure_schema(trace_retention_days, spend_log_retention_days) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 2da31d9f0fd..a471fb6f6f8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -88,11 +88,17 @@ from .types_utils.utils import get_instance_fn, validate_custom_validate_return_ if TYPE_CHECKING: from opentelemetry.trace import Span as _Span + from litellm.tracing import TraceReceiver + Span = _Span | Any else: Span = Any +class ProxyLifespanState(TypedDict): + tracing_receiver: ReadOnly["TraceReceiver | None"] + + class ReconcileOutcome(NamedTuple): """What a model reconcile observed, captured while it still held the reconcile lock. diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index f24715d8265..0349c594adf 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -35,7 +35,7 @@ from litellm.proxy.lens.models import ( WorkerCreated, ) from litellm.proxy.lens.repository import LensRepository, WriterDatabase -from litellm.proxy.lens.sources import SourceReader, parse_execution +from litellm.proxy.lens.sources import SourceReader, Storage, parse_execution from litellm.proxy.lens.state import ( can_access, claim_job, @@ -45,10 +45,12 @@ from litellm.proxy.lens.state import ( replace_job, snapshot_finding, ) +from litellm.proxy.tracing_runtime import provide_storage router: Final = APIRouter(prefix="/lens", tags=["Lens"]) _bearer: Final = HTTPBearer() Auth: TypeAlias = Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)] +StorageDep: TypeAlias = Annotated[Storage | None, Depends(provide_storage)] def repository() -> LensRepository: @@ -59,10 +61,13 @@ def repository() -> LensRepository: return LensRepository(WriterDatabase(writer_wrapper(prisma_client.db))) -def source_reader() -> SourceReader: - from litellm.proxy.tracing_endpoints import get_receiver - - return SourceReader(get_receiver().store.storage) +def source_reader(storage: Storage | None) -> SourceReader: + if storage is None: + raise HTTPException( + status_code=501, + detail="Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL.", + ) + return SourceReader(storage) def user_scope(auth: UserAPIKeyAuth, write: bool = False) -> Scope: @@ -138,14 +143,12 @@ def validate_model(settings: LensSettings, auth: UserAPIKeyAuth) -> None: @router.get("", response_model=LensList) -async def list_lenses(auth: Auth) -> LensList: - from litellm.proxy import tracing_endpoints - +async def list_lenses(auth: Auth, storage: StorageDep) -> LensList: scope: Final = user_scope(auth) return LensList( lenses=tuple(e for e in await repository().lenses() if can_access(scope, e.scope)), workers=tuple(w for w in await repository().workers() if can_access(scope, w.scope)), - tracing_enabled=tracing_endpoints.receiver is not None, + tracing_enabled=storage is not None, ) @@ -265,10 +268,10 @@ class Preview(BaseModel): @router.post("/preview/sample", response_model=Sample) -async def preview_sample(body: Preview, auth: Auth) -> Sample: +async def preview_sample(body: Preview, auth: Auth, storage: StorageDep) -> Sample: validate_selection(body.settings) now: Final = min(body.as_of or datetime.now(timezone.utc), datetime.now(timezone.utc)) - return await source_reader().sample( + return await source_reader(storage).sample( user_scope(auth), body.settings, int((now - timedelta(hours=body.lookback_hours)).timestamp() * 1000), @@ -367,14 +370,14 @@ async def progress(lens_id: str, job_id: str, body: Progress, worker: WorkerAuth @router.get("/worker/{lens_id}/{job_id}/sample", response_model=Sample) -async def sample(lens_id: str, job_id: str, worker: WorkerAuth) -> Sample: +async def sample(lens_id: str, job_id: str, worker: WorkerAuth, storage: StorageDep) -> Sample: lens, job = await assigned(lens_id, job_id, worker) if job.sample is not None: return job.sample pages: list[Sample] = [] # mutable-ok: freeze selection after stable cursor traversal cursor = "" # rebind-ok: advance by immutable identity, never by shifting row positions while True: - page = await source_reader().sample( + page = await source_reader(storage).sample( lens.scope, job.settings, int(job.start.timestamp() * 1000), @@ -413,6 +416,7 @@ async def content( job_id: str, execution_id: str, worker: WorkerAuth, + storage: StorageDep, cursor: str = "", offset: int = Query(default=0, ge=0), ) -> ExecutionContent: @@ -421,7 +425,7 @@ async def content( execution: Final = next((e for e in selected.executions if e.id == execution_id), None) if execution is None: raise HTTPException(404, "Execution is outside this job's sample") - return await source_reader().content(lens.scope, execution, cursor, offset) + return await source_reader(storage).content(lens.scope, execution, cursor, offset) @router.post("/worker/{lens_id}/{job_id}/model", response_model=ModelResult) @@ -433,7 +437,7 @@ async def model(lens_id: str, job_id: str, body: ModelRequest, worker: WorkerAut @router.post("/worker/{lens_id}/{job_id}/result", response_model=Lens) -async def result(lens_id: str, job_id: str, body: Result, worker: WorkerAuth) -> Lens: +async def result(lens_id: str, job_id: str, body: Result, worker: WorkerAuth, storage: StorageDep) -> Lens: lens: Final = await get_lens(lens_id, worker.scope) old: Final = next((j for j in lens.jobs if j.id == job_id), None) if old and old.status in ("completed", "failed") and old.worker_id == worker.id: @@ -455,7 +459,7 @@ async def result(lens_id: str, job_id: str, body: Result, worker: WorkerAuth) -> raise HTTPException(422, "Finding references evidence outside the job") for finding in body.findings: - await validate_finding(lens, selected, finding) + await validate_finding(lens, selected, finding, storage) def finish(e: Lens) -> Lens: active: Final = current_job(e) @@ -523,12 +527,12 @@ async def claim_candidate(candidate: Lens, worker: Worker, now: datetime) -> Cla return None -async def validate_finding(lens: Lens, selected: Sample, finding: FindingDraft) -> None: +async def validate_finding(lens: Lens, selected: Sample, finding: FindingDraft, storage: Storage | None) -> None: previous: Final = next((f for f in lens.findings if f.id == finding.existing_finding_id), None) if finding.existing_finding_id and (previous is None or previous.check_id != finding.check_id): raise HTTPException(422, "Existing finding must belong to the same check") for evidence in finding.evidence: - if not await source_reader().verify_evidence( + if not await source_reader(storage).verify_evidence( lens.scope, next(e for e in selected.executions if e.id == evidence.execution_id), evidence ): raise HTTPException(422, "Evidence quote does not match stored content") @@ -536,7 +540,12 @@ async def validate_finding(lens: Lens, selected: Sample, finding: FindingDraft) @router.get("/{lens_id}/executions/{execution_id}", response_model=ExecutionContent) async def evidence_content( - lens_id: str, execution_id: str, auth: Auth, cursor: str = "", offset: int = Query(default=0, ge=0) + lens_id: str, + execution_id: str, + auth: Auth, + storage: StorageDep, + cursor: str = "", + offset: int = Query(default=0, ge=0), ) -> ExecutionContent: lens: Final = await get_lens(lens_id, user_scope(auth)) try: @@ -556,4 +565,4 @@ async def evidence_content( span_count=1, root_seen=source == "requests", ) - return await source_reader().content(lens.scope, execution, cursor, offset) + return await source_reader(storage).content(lens.scope, execution, cursor, offset) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ff0dd9df16f..cf09bdbef9b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -119,6 +119,7 @@ from litellm.proxy._types import ( PassThroughGenericEndpoint, ProxyErrorTypes, ProxyException, + ProxyLifespanState, SpecialModelNames, SupportedDBObjectType, TeamDefaultSettings, @@ -792,6 +793,7 @@ from litellm.proxy.spend_tracking.spend_management_endpoints import ( router as spend_management_router, ) from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload +from litellm.proxy.tracing_runtime import manage_tracing from litellm.proxy.types_utils.utils import get_instance_fn from litellm.proxy.ui_crud_endpoints.latest_release_endpoints import ( router as latest_release_endpoints_router, @@ -852,7 +854,6 @@ from litellm.secret_managers.main import ( secret_manager_would_be_consulted, str_to_bool, ) -from litellm.tracing import TraceReceiver from litellm.types.integrations.slack_alerting import AlertType, SlackAlertingArgs from litellm.types.llms.anthropic import ( AnthropicMessagesRequest, @@ -1222,7 +1223,7 @@ async def _connect_to_count_stored_values() -> SupportsRawQueries: @asynccontextmanager -async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: +async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]: global \ prisma_client, \ master_key, \ @@ -1529,9 +1530,6 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: _tagged.strategy._state_loaded = True asyncio.create_task(_adaptive_router_flusher_loop()) - ## [Optional] Initialize agent tracing - asyncio.create_task(ProxyStartupEvent.init_tracing(general_settings)) - ## [Optional] Initialize dd tracer ProxyStartupEvent._init_dd_tracer() @@ -1565,76 +1563,81 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: register_scheduled_sync(scheduler) - # End of startup event - yield + tracing_settings: Final = general_settings.get("tracing") + tracing_enabled: Final = TypeAdapter(bool).validate_python( + isinstance(tracing_settings, dict) and tracing_settings.get("store") == "clickhouse" + ) + async with manage_tracing(enabled=tracing_enabled) as receiver: + state: Final[ProxyLifespanState] = {"tracing_receiver": receiver} + yield state - if model_info_scheduler is not None and model_info_scheduler.running: - model_info_scheduler.remove_job("refresh_model_info") - if model_info_scheduler is not scheduler: - model_info_scheduler.shutdown(wait=False) + if model_info_scheduler is not None and model_info_scheduler.running: + model_info_scheduler.remove_job("refresh_model_info") + if model_info_scheduler is not scheduler: + model_info_scheduler.shutdown(wait=False) - # Shutdown event - stop starting scheduled jobs; the ones already running keep the drain window - if scheduler is not None: - pause_scheduled_jobs(scheduler) + # Shutdown event - stop starting scheduled jobs; the ones already running keep the drain window + if scheduler is not None: + pause_scheduled_jobs(scheduler) - # Shutdown event - drain in-flight requests before tearing down dependencies - # so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them. - GracefulShutdownManager.start_shutdown() - await GracefulShutdownManager.wait_for_drain() + # Shutdown event - drain in-flight requests before tearing down dependencies + # so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them. + GracefulShutdownManager.start_shutdown() + await GracefulShutdownManager.wait_for_drain() - # Shutdown event - close shared aiohttp session - if shared_aiohttp_session is not None: - try: - await shared_aiohttp_session.close() - verbose_proxy_logger.info("SESSION REUSE: Closed shared aiohttp session") - except Exception as e: - verbose_proxy_logger.error("Error closing shared aiohttp session: %s", e) + # Shutdown event - close shared aiohttp session + if shared_aiohttp_session is not None: + try: + await shared_aiohttp_session.close() + verbose_proxy_logger.info("SESSION REUSE: Closed shared aiohttp session") + except Exception as e: + verbose_proxy_logger.error("Error closing shared aiohttp session: %s", e) - # Shutdown event - stop RDS IAM token refresh background task - if ( - prisma_client is not None - and hasattr(prisma_client, "db") - and hasattr(prisma_client.db, "stop_token_refresh_task") - ): - try: - await prisma_client.db.stop_token_refresh_task() - except Exception as e: - verbose_proxy_logger.error("Error stopping token refresh task: %s", e) + # Shutdown event - stop RDS IAM token refresh background task + if ( + prisma_client is not None + and hasattr(prisma_client, "db") + and hasattr(prisma_client.db, "stop_token_refresh_task") + ): + try: + await prisma_client.db.stop_token_refresh_task() + except Exception as e: + verbose_proxy_logger.error("Error stopping token refresh task: %s", e) - # Shutdown event - stop Prisma DB health watchdog task - if prisma_client is not None and hasattr(prisma_client, "stop_db_health_watchdog_task"): - try: - await prisma_client.stop_db_health_watchdog_task() - except Exception as e: - verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e) + # Shutdown event - stop Prisma DB health watchdog task + if prisma_client is not None and hasattr(prisma_client, "stop_db_health_watchdog_task"): + try: + await prisma_client.stop_db_health_watchdog_task() + except Exception as e: + verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e) - if prisma_client is not None and hasattr(prisma_client, "stop_view_setup_task"): - try: - await prisma_client.stop_view_setup_task() - except Exception as e: - verbose_proxy_logger.error("Error stopping the spend view setup task: %s", e) + if prisma_client is not None and hasattr(prisma_client, "stop_view_setup_task"): + try: + await prisma_client.stop_view_setup_task() + except Exception as e: + verbose_proxy_logger.error("Error stopping the spend view setup task: %s", e) - await _drain_spend_event_producer_on_shutdown() + await _drain_spend_event_producer_on_shutdown() - # Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect - if scheduler is not None and scheduler_executor is not None: - try: - await stop_in_flight_scheduler_jobs(scheduler, scheduler_executor) - except Exception as e: - verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e) + # Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect + if scheduler is not None and scheduler_executor is not None: + try: + await stop_in_flight_scheduler_jobs(scheduler, scheduler_executor) + except Exception as e: + verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e) - await flush_spend_counters_on_shutdown() + await flush_spend_counters_on_shutdown() - await _flush_spend_logs_queue_on_shutdown() + await _flush_spend_logs_queue_on_shutdown() - await proxy_config.stop_config_sync_subscriber() + await proxy_config.stop_config_sync_subscriber() - await proxy_config.stop_auth_cache_invalidation_subscriber() + await proxy_config.stop_auth_cache_invalidation_subscriber() - await proxy_shutdown_event(worker_heartbeat=worker_heartbeat) + await proxy_shutdown_event(worker_heartbeat=worker_heartbeat) - if prometheus_multiproc_dir: - mark_worker_exit(os.getpid()) + if prometheus_multiproc_dir: + mark_worker_exit(os.getpid()) def _generate_stable_operation_id(route: "APIRoute") -> str: @@ -11357,39 +11360,6 @@ class ProxyStartupEvent: ) return connected_client - @classmethod - async def init_tracing(cls, general_settings: dict, receiver: TraceReceiver | None = None) -> None: - """ - Enable agent tracing (`POST/GET /v1/traces`) when configured: - - general_settings: - tracing: - store: clickhouse - """ - from litellm.integrations.clickhouse.clickhouse_spend_logger import ClickHouseSpendLogger - - manager: Final = litellm.logging_callback_manager - for callback in manager.get_custom_loggers_for_type(ClickHouseSpendLogger): - manager.remove_callback_from_all_lists(callback) - tracing_endpoints.receiver = None - settings: Final = general_settings.get("tracing") - if not isinstance(settings, dict) or settings.get("store") != "clickhouse": - return - try: - tracing: Final = receiver if receiver is not None else TraceReceiver.from_env() - await tracing.start() - except (KeyError, OSError, RuntimeError, ValueError) as error: - verbose_proxy_logger.warning("Agent tracing unavailable: %s", error) - return - tracing_endpoints.receiver = tracing - spend_logger: Final = ClickHouseSpendLogger(storage=tracing.store.storage) - manager.add_litellm_callback(spend_logger) - manager.add_litellm_success_callback(spend_logger) - manager.add_litellm_failure_callback(spend_logger) - manager.add_litellm_async_success_callback(spend_logger) - manager.add_litellm_async_failure_callback(spend_logger) - verbose_proxy_logger.info("Agent tracing enabled (store=clickhouse)") - @classmethod def _init_dd_tracer(cls): """ diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index b2b10ebfbc6..bd885282859 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -8,6 +8,7 @@ GET /v1/traces/{trace_id}/spans/{span_id} SpanDetail """ import time +from dataclasses import dataclass from typing import Annotated, Final from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response @@ -15,6 +16,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response from litellm.constants import OTLP_MAX_BODY_BYTES, OTLP_RETRY_AFTER_SECONDS from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.tracing_runtime import provide_receiver, require_receiver from litellm.tracing import ( Tenant, TraceReceiver, @@ -26,37 +28,42 @@ from litellm.tracing.types import SpanDetail, Trace, TracePage, TraceScope router = APIRouter(tags=["agent tracing"]) MS_PER_DAY: Final = 24 * 60 * 60 * 1000 -_ADMIN_ROLES: Final = (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) - -receiver: TraceReceiver | None = None -def get_receiver() -> TraceReceiver: - if receiver is None: - raise HTTPException( - status_code=501, - detail="Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL.", - ) - return receiver +@dataclass(frozen=True, slots=True) +class TraceAccessContext: + receiver: TraceReceiver | None + read_scope: TraceScope | None + write_tenant: Tenant | None + + def reader(self) -> tuple[TraceReceiver, TraceScope]: + tracing: Final = require_receiver(self.receiver) + if self.read_scope is None: + raise HTTPException(status_code=403, detail="Not allowed to view agent traces") + return tracing, self.read_scope + + def writer(self) -> tuple[TraceReceiver, Tenant]: + if self.write_tenant is None: + raise HTTPException(status_code=403, detail="Not allowed to ingest agent traces") + return require_receiver(self.receiver), self.write_tenant -def tenant_for(user_api_key_dict: UserAPIKeyAuth) -> Tenant: - return Tenant( - team_id=user_api_key_dict.team_id or "", - api_key_hash=user_api_key_dict.token or "", - org_id=user_api_key_dict.org_id or "", - ) - - -def scope_for(user_api_key_dict: UserAPIKeyAuth) -> TraceScope: - """Admins see everything; team members see their team; team-less keys see their own traces.""" - if user_api_key_dict.user_role in _ADMIN_ROLES: - return TraceScope(team_ids=(), api_key_hash="") - if user_api_key_dict.team_id: - return TraceScope(team_ids=(user_api_key_dict.team_id,), api_key_hash="") - if not user_api_key_dict.token: - raise HTTPException(status_code=403, detail="Not allowed to view agent traces") - return TraceScope(team_ids=("",), api_key_hash=user_api_key_dict.token) +async def provide_trace_access( + auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + tracing: Annotated[TraceReceiver | None, Depends(provide_receiver)], +) -> TraceAccessContext: + tenant: Final = Tenant(team_id=auth.team_id or "", api_key_hash=auth.token or "", org_id=auth.org_id or "") + match auth.user_role: + case LitellmUserRoles.PROXY_ADMIN: + return TraceAccessContext(tracing, TraceScope(team_ids=(), api_key_hash=""), tenant) + case LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY: + return TraceAccessContext(tracing, TraceScope(team_ids=(), api_key_hash=""), None) + case _ if auth.team_id: + return TraceAccessContext(tracing, TraceScope(team_ids=(auth.team_id,), api_key_hash=""), tenant) + case _ if auth.token: + return TraceAccessContext(tracing, TraceScope(team_ids=("",), api_key_hash=auth.token), tenant) + case _: + return TraceAccessContext(tracing, None, tenant) async def _read_otlp_body(request: Request) -> bytes: @@ -71,18 +78,16 @@ async def _read_otlp_body(request: Request) -> bytes: @router.post("/v1/traces", include_in_schema=False) async def ingest_otlp_traces( request: Request, - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + context: Annotated[TraceAccessContext, Depends(provide_trace_access)], ) -> Response: - if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY: - raise HTTPException(status_code=403, detail="Not allowed to ingest agent traces") - tracing: Final = get_receiver() + tracing, tenant = context.writer() content_type: Final = request.headers.get("content-type") try: await tracing.ingest( body=await _read_otlp_body(request), content_type=content_type, content_encoding=request.headers.get("content-encoding"), - tenant=tenant_for(user_api_key_dict), + tenant=tenant, ) except TracingPayloadTooLargeError as e: raise HTTPException(status_code=413, detail=str(e)) @@ -99,15 +104,16 @@ async def ingest_otlp_traces( @router.get("/v1/traces", response_model=None) async def list_agent_traces( - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + context: Annotated[TraceAccessContext, Depends(provide_trace_access)], start_ms: Annotated[int | None, Query(description="Window start, unix ms. Default: 24h ago")] = None, end_ms: Annotated[int | None, Query(description="Window end, unix ms. Default: now")] = None, cursor: Annotated[str | None, Query()] = None, ) -> TracePage: now_ms: Final = int(time.time() * 1000) try: - return await get_receiver().list_traces( - scope=scope_for(user_api_key_dict), + tracing, scope = context.reader() + return await tracing.list_traces( + scope=scope, start_ms=start_ms if start_ms is not None else now_ms - MS_PER_DAY, end_ms=end_ms if end_ms is not None else now_ms, cursor=cursor, @@ -119,10 +125,11 @@ async def list_agent_traces( @router.get("/v1/traces/{trace_id}", response_model=None) async def get_agent_trace( trace_id: str, - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + context: Annotated[TraceAccessContext, Depends(provide_trace_access)], trace_ref: Annotated[str, Query()] = "", ) -> Trace: - trace: Final = await get_receiver().get_trace(trace_id, scope_for(user_api_key_dict), trace_ref) + tracing, scope = context.reader() + trace: Final = await tracing.get_trace(trace_id, scope, trace_ref) if trace is None: raise HTTPException(status_code=404, detail=f"Trace {trace_id} not found") return trace @@ -132,10 +139,11 @@ async def get_agent_trace( async def get_agent_trace_span( trace_id: str, span_id: str, - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + context: Annotated[TraceAccessContext, Depends(provide_trace_access)], trace_ref: Annotated[str, Query()] = "", ) -> SpanDetail: - span: Final = await get_receiver().get_span(trace_id, span_id, scope_for(user_api_key_dict), trace_ref) + tracing, scope = context.reader() + span: Final = await tracing.get_span(trace_id, span_id, scope, trace_ref) if span is None: raise HTTPException(status_code=404, detail=f"Span {span_id} not found") return span diff --git a/litellm/proxy/tracing_runtime.py b/litellm/proxy/tracing_runtime.py new file mode 100644 index 00000000000..0b706d66a40 --- /dev/null +++ b/litellm/proxy/tracing_runtime.py @@ -0,0 +1,66 @@ +from collections.abc import AsyncGenerator, Callable +from contextlib import asynccontextmanager +from typing import Final + +from fastapi import HTTPException, Request +from pydantic import ConfigDict, TypeAdapter + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.integrations.clickhouse.clickhouse_spend_logger import ClickHouseSpendLogger +from litellm.rust_bridge.traces import ClickHouseStorage +from litellm.tracing import TraceReceiver + +_RECEIVER_ADAPTER: Final[TypeAdapter[TraceReceiver | None]] = TypeAdapter( + TraceReceiver | None, config=ConfigDict(arbitrary_types_allowed=True) +) +_UNAVAILABLE_DETAIL: Final = "Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL." + + +def require_receiver(tracing: TraceReceiver | None) -> TraceReceiver: + if tracing is None: + raise HTTPException(status_code=501, detail=_UNAVAILABLE_DETAIL) + return tracing + + +async def provide_receiver(request: Request) -> TraceReceiver | None: + return _RECEIVER_ADAPTER.validate_python(getattr(request.state, "tracing_receiver", None)) + + +async def provide_storage(request: Request) -> ClickHouseStorage | None: + tracing: Final = await provide_receiver(request) + return tracing.store.storage if tracing is not None else None + + +async def _start_receiver(factory: Callable[[], TraceReceiver]) -> TraceReceiver | None: + try: + tracing: Final = factory() + await tracing.start() + return tracing + except (KeyError, OSError, RuntimeError, ValueError) as error: + verbose_proxy_logger.warning("Agent tracing unavailable: %s", error) + return None + + +@asynccontextmanager +async def manage_tracing( + enabled: bool, receiver_factory: Callable[[], TraceReceiver] = TraceReceiver.from_env +) -> AsyncGenerator[TraceReceiver | None, None]: + tracing: Final = await _start_receiver(receiver_factory) if enabled else None + if tracing is None: + yield tracing + return + + spend_logger: Final = ClickHouseSpendLogger(storage=tracing.store.storage) + manager: Final = litellm.logging_callback_manager + manager.add_litellm_callback(spend_logger) + manager.add_litellm_success_callback(spend_logger) + manager.add_litellm_failure_callback(spend_logger) + manager.add_litellm_async_success_callback(spend_logger) + manager.add_litellm_async_failure_callback(spend_logger) + verbose_proxy_logger.info("Agent tracing enabled (store=clickhouse)") + try: + yield tracing + finally: + manager.remove_callback_from_all_lists(spend_logger) + await spend_logger.aclose() diff --git a/litellm/rust_bridge/traces.py b/litellm/rust_bridge/traces.py index 98607aa9206..1c20e408709 100644 --- a/litellm/rust_bridge/traces.py +++ b/litellm/rust_bridge/traces.py @@ -80,7 +80,7 @@ def decode_otlp( return _native().trace_decode_otlp(body, content_type, content_encoding, max_decompressed_bytes) -class TraceStorage: +class ClickHouseStorage: def __init__(self, database: str, url: str, reader_url: str | None = None) -> None: self._native: Final = _native().NativeTraceStorage(database, url, reader_url) diff --git a/litellm/tracing/AGENTS.md b/litellm/tracing/AGENTS.md index f69866c0419..ee1c8870edd 100644 --- a/litellm/tracing/AGENTS.md +++ b/litellm/tracing/AGENTS.md @@ -1,6 +1,6 @@ - Python owns tracing endpoints, authenticated tenant scope, framework normalization and API response shaping -- Trace ingestion awaits `TraceStorage.insert_rows` before returning success; propagate storage failures so OTLP exporters can retry +- Trace ingestion awaits `ClickHouseStorage.insert_rows` before returning success; propagate storage failures so OTLP exporters can retry - Spend logging keeps its separate batch queue in `litellm/integrations/clickhouse` -- Use `litellm.rust_bridge.traces.TraceStorage` for ClickHouse; keep schema, SQL, encoding and transport in `litellm-traces` +- Use `litellm.rust_bridge.traces.ClickHouseStorage` for ClickHouse; keep trace schema, SQL and encoding in `litellm-traces`, and generic transport in `litellm-storage-clickhouse` - Derive tenant fields from authentication and overwrite matching fields supplied by the exporter - Test confirmed writes, failures, tenant isolation and read behavior through public functions diff --git a/litellm/tracing/receiver.py b/litellm/tracing/receiver.py index 05d9dcb3307..700f33a8a0e 100644 --- a/litellm/tracing/receiver.py +++ b/litellm/tracing/receiver.py @@ -23,9 +23,9 @@ from litellm.constants import ( OTLP_OFFLOAD_DECODE_BYTES, ) from litellm.integrations.clickhouse.schema import ensure_schema -from litellm.rust_bridge.traces import TraceStorage +from litellm.rust_bridge.traces import ClickHouseStorage from litellm.tracing.decode import OTLPPayloadTooLargeError, decode_otlp -from litellm.tracing.store import ClickHouseTraceStore +from litellm.tracing.store import TraceStore from litellm.tracing.types import ( SpanDetail, SpanRow, @@ -60,14 +60,14 @@ class Tenant: class TraceReceiver: - def __init__(self, store: ClickHouseTraceStore) -> None: + def __init__(self, store: TraceStore) -> None: self.store = store @classmethod def from_env(cls) -> "TraceReceiver": return cls( - store=ClickHouseTraceStore( - TraceStorage( + store=TraceStore( + ClickHouseStorage( database=os.getenv("CLICKHOUSE_DATABASE", "litellm"), url=os.environ["CLICKHOUSE_URL"], reader_url=os.environ["CLICKHOUSE_READER_URL"], diff --git a/litellm/tracing/store.py b/litellm/tracing/store.py index fdb1c7820f2..9d1f64f77f0 100644 --- a/litellm/tracing/store.py +++ b/litellm/tracing/store.py @@ -16,7 +16,7 @@ from litellm.constants import AGENT_TRACING_LIST_PAGE_SIZE from litellm.integrations.clickhouse.schema import ( OTEL_TRACES_TABLE, ) -from litellm.rust_bridge.traces import TraceStorage +from litellm.rust_bridge.traces import ClickHouseStorage from litellm.tracing.types import ( AgentNode, Span, @@ -259,10 +259,10 @@ def trace_from_rows( ) -class ClickHouseTraceStore: +class TraceStore: """Stores spans and runs scoped trace reads.""" - def __init__(self, storage: TraceStorage) -> None: + def __init__(self, storage: ClickHouseStorage) -> None: self.storage = storage async def insert_spans(self, rows: Sequence[SpanRow]) -> None: diff --git a/tests/proxy_behavior/lens/test_lifecycle.py b/tests/proxy_behavior/lens/test_lifecycle.py index 3196b68bd83..849b2186a62 100644 --- a/tests/proxy_behavior/lens/test_lifecycle.py +++ b/tests/proxy_behavior/lens/test_lifecycle.py @@ -82,7 +82,7 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database: ) assert stored_worker is not None and stored_worker.id == worker.id assert worker.id == registration.worker.id - listing: Final = await endpoints.list_lenses(admin) + listing: Final = await endpoints.list_lenses(admin, storage=None) assert lens.id in tuple(e.id for e in listing.lenses) assert worker.id in tuple(w.id for w in listing.workers) claims: Final = await asyncio.gather( @@ -180,13 +180,13 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database: assert needs_billing.value.status_code == 409 assert await endpoints.heartbeat(lens.id, claimed.job.id, authenticated_legacy) finished: Final = await endpoints.result( - lens.id, claimed.job.id, Result(coverage=Coverage(screened=2)), authenticated_legacy + lens.id, claimed.job.id, Result(coverage=Coverage(screened=2)), authenticated_legacy, storage=None ) assert finished.jobs[0].status == "completed" assert finished.jobs[0].coverage.screened == 2 assert finished.last_scan_at == claimed.job.end assert finished.next_run_at > finished.jobs[0].finished_at - assert await endpoints.result(lens.id, claimed.job.id, Result(coverage=Coverage()), worker) == finished + assert await endpoints.result(lens.id, claimed.job.id, Result(coverage=Coverage()), worker, storage=None) == finished with pytest.raises(HTTPException) as stale: await endpoints.heartbeat(lens.id, claimed.job.id, worker) assert stale.value.status_code == 409 diff --git a/tests/proxy_migration_tests/test_prisma_toolchain.py b/tests/proxy_migration_tests/test_prisma_toolchain.py index 556c680a84a..ebe2390db16 100644 --- a/tests/proxy_migration_tests/test_prisma_toolchain.py +++ b/tests/proxy_migration_tests/test_prisma_toolchain.py @@ -312,7 +312,7 @@ def test_db_push_timeout_hint_names_the_per_command_budget( ) -> None: """``db push`` keeps the per-command budget, so its timeout hint has to name that variable.""" _, log_path = toolchain_env - monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x") + monkeypatch.delenv("DATABASE_URL", raising=False) monkeypatch.setenv(PRISMA_COMMAND_TIMEOUT_ENV_VAR, "1") monkeypatch.setenv("FAKE_PRISMA_FIRST_PUSH_SLEEP", "3") diff --git a/tests/test_litellm/integrations/clickhouse/test_clickhouse_batch_logger.py b/tests/test_litellm/integrations/clickhouse/test_clickhouse_batch_logger.py index bae94ba6100..5eb14e73855 100644 --- a/tests/test_litellm/integrations/clickhouse/test_clickhouse_batch_logger.py +++ b/tests/test_litellm/integrations/clickhouse/test_clickhouse_batch_logger.py @@ -3,6 +3,8 @@ Tests for the CustomBatchLogger-based ClickHouse base logger. """ import asyncio +from collections.abc import Mapping, Sequence +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -50,8 +52,7 @@ async def test_first_enqueued_row_flushes_after_synchronous_construction(): logger.enqueue([{"i": 1}]) await asyncio.wait_for(flushed.wait(), timeout=1) - if logger._flush_task is not None: - logger._flush_task.cancel() + await logger.aclose() @pytest.mark.asyncio @@ -79,3 +80,63 @@ async def test_failed_insert_is_requeued_then_dropped(): assert logger.rows_dropped == 2 assert logger.rows_written == 0 assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_close_waits_for_active_insert_and_stops_periodic_flush() -> None: + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def insert_rows(table: str, rows: Sequence[Mapping[str, object]]) -> None: + started.set() + await release.wait() + + insert: Final = AsyncMock(side_effect=insert_rows) + logger: Final = _logger(insert) + logger.flush_interval = 0.001 + logger.enqueue([{"i": 1}]) + await asyncio.wait_for(started.wait(), timeout=1) + closing: Final = asyncio.create_task(logger.aclose()) + await asyncio.sleep(0) + assert not closing.done() + release.set() + await asyncio.wait_for(closing, timeout=1) + assert logger.rows_written == 1 + insert.assert_awaited_once_with("test_table", [{"i": 1}]) + assert logger._flush_task is not None and logger._flush_task.done() + assert not logger._flush_task.cancelled() + + +@pytest.mark.asyncio +async def test_close_wakes_idle_worker_and_drains_queued_rows() -> None: + insert: Final = AsyncMock() + logger: Final = _logger(insert) + logger.flush_interval = 3600 + logger.enqueue([{"i": 1}]) + await asyncio.sleep(0) + + await asyncio.wait_for(logger.aclose(), timeout=1) + + insert.assert_awaited_once_with("test_table", [{"i": 1}]) + assert logger.rows_written == 1 + assert logger.log_queue == [] + assert logger._flush_task is not None and logger._flush_task.done() + assert not logger._flush_task.cancelled() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("recovers", [True, False]) +async def test_close_retries_every_batch_and_accounts_for_exhausted_rows(recovers: bool) -> None: + failure: Final = RuntimeError("ClickHouse unavailable") + insert: Final = AsyncMock(side_effect=[failure, None, None] if recovers else failure) + logger: Final = _logger(insert) + logger.batch_size = 1 + logger.log_queue.extend([{"request_id": "a"}, {"request_id": "b"}]) + + await logger.aclose() + + assert logger.log_queue == [] + assert logger.rows_written == (2 if recovers else 0) + assert logger.rows_dropped == (0 if recovers else 2) + assert insert.await_count == (3 if recovers else 2 * module.CLICKHOUSE_MAX_RETRIES) + assert {call.args[1][0]["request_id"] for call in insert.await_args_list} == {"a", "b"} diff --git a/tests/test_litellm/tracing/test_store.py b/tests/test_litellm/tracing/test_store.py index 2e80594ff9b..30d7a9b5b0b 100644 --- a/tests/test_litellm/tracing/test_store.py +++ b/tests/test_litellm/tracing/test_store.py @@ -8,7 +8,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest from litellm.tracing.store import ( - ClickHouseTraceStore, + TraceStore, agent_nodes, decode_cursor, encode_cursor, @@ -284,7 +284,7 @@ async def test_list_traces_sets_next_cursor_on_full_page(): "models": [], } client.query = AsyncMock(return_value=[row, {**row, "trace_id": "t1", "trace_ref": "ref1", "start_ms": 900}]) - store = ClickHouseTraceStore(client) + store = TraceStore(client) scope: TraceScope = {"team_ids": ("team-a",), "api_key_hash": ""} page = await store.list_traces(scope, 0, 2000, limit=2) @@ -303,7 +303,7 @@ async def test_list_traces_sets_next_cursor_on_full_page(): async def test_get_span_not_found_and_found(): client = MagicMock() client.query = AsyncMock(return_value=[]) - store = ClickHouseTraceStore(client) + store = TraceStore(client) scope: TraceScope = {"team_ids": (), "api_key_hash": ""} assert await store.get_span("t", "s", scope) is None stored_input = '[{"role": "user", "content": "hi"}]' @@ -355,7 +355,7 @@ async def test_trace_cost_is_scoped_and_counts_repeated_request_once(): }, ] client.query = AsyncMock(side_effect=[spans, spend]) - store = ClickHouseTraceStore(client) + store = TraceStore(client) scope: TraceScope = {"team_ids": ("team-a",), "api_key_hash": ""} trace = await store.get_trace("trace-1", scope) @@ -406,7 +406,7 @@ async def test_run_list_uses_matching_spend_and_leaves_missing_cost_unavailable( client.query = AsyncMock(side_effect=[rows, spend]) scope: TraceScope = {"team_ids": ("team-a",), "api_key_hash": ""} - page = await ClickHouseTraceStore(client).list_traces(scope, 0, 2000) + page = await TraceStore(client).list_traces(scope, 0, 2000) assert [run["spend"] for run in page["data"]] == [0.25, None] assert [call.args[0] for call in client.query.await_args_list] == ["list_traces", "spend_by_response_ids"] @@ -428,7 +428,7 @@ async def test_ambiguous_cache_response_id_keeps_cost_unavailable(): for request_id, cost in (("response-1", 0.25), ("response-1_cache_hit123", 0.0)) ] client.query = AsyncMock(side_effect=[[span], spend]) - store = ClickHouseTraceStore(client) + store = TraceStore(client) scope: TraceScope = {"team_ids": ("",), "api_key_hash": "key-a"} trace = await store.get_trace("trace-1", scope) diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index c3709ceae3f..9785bdd5e32 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -14,6 +14,7 @@ import logging import os import re from collections.abc import Mapping +from contextlib import nullcontext from dataclasses import dataclass from datetime import datetime from pathlib import Path @@ -44,48 +45,51 @@ from .conftest import normalize @pytest.mark.asyncio -async def test_tracing_config_automatically_logs_spend_without_callback_setting(): +@pytest.mark.parametrize("shutdown_error", [False, True]) +async def test_tracing_config_automatically_logs_spend_without_callback_setting(shutdown_error: bool) -> None: from litellm.integrations.clickhouse.clickhouse_spend_logger import ClickHouseSpendLogger - from litellm.proxy import tracing_endpoints - from litellm.proxy.proxy_server import ProxyStartupEvent + from litellm.proxy.tracing_runtime import manage_tracing from litellm.tracing import TraceReceiver - from litellm.tracing.store import ClickHouseTraceStore + from litellm.tracing.store import TraceStore - storage = MagicMock() + storage: Final = MagicMock() storage.ensure_schema = AsyncMock() storage.insert_rows = AsyncMock() - receiver = TraceReceiver(ClickHouseTraceStore(storage)) - prior_receiver = tracing_endpoints.receiver + receiver: Final = TraceReceiver(TraceStore(storage)) - try: - await ProxyStartupEvent.init_tracing({"tracing": {"store": "clickhouse"}}, receiver=receiver) - storage.ensure_schema.assert_awaited_once() - logger = next( - callback for callback in litellm._async_success_callback if isinstance(callback, ClickHouseSpendLogger) - ) - now = datetime.now() - await logger.async_log_success_event( - { - "standard_logging_object": { - "id": "response-1", - "startTime": now.timestamp(), - "endTime": now.timestamp(), - "response_cost": 0.25, - } - }, - None, - now, - now, - ) - await logger.flush_queue() - assert storage.insert_rows.await_args.args[0] == "spend_logs" - assert storage.insert_rows.await_args.args[1][0]["spend"] == 0.25 + outcome: Final = pytest.raises(RuntimeError, match="shutdown failure") if shutdown_error else nullcontext() + with outcome: + async with manage_tracing(enabled=True, receiver_factory=lambda: receiver): + storage.ensure_schema.assert_awaited_once() + logger: Final = next( + callback + for callback in litellm._async_success_callback + if isinstance(callback, ClickHouseSpendLogger) and callback.storage is storage + ) + now: Final = datetime.now() + await logger.async_log_success_event( + { + "standard_logging_object": { + "id": "response-1", + "startTime": now.timestamp(), + "endTime": now.timestamp(), + "response_cost": 0.25, + } + }, + None, + now, + now, + ) + storage.insert_rows.assert_not_awaited() - await ProxyStartupEvent.init_tracing({}) - assert all(not isinstance(callback, ClickHouseSpendLogger) for callback in litellm._async_success_callback) - finally: - await ProxyStartupEvent.init_tracing({}) - tracing_endpoints.receiver = prior_receiver + if shutdown_error: + raise RuntimeError("shutdown failure") + + assert storage.insert_rows.await_args.args[0] == "spend_logs" + assert storage.insert_rows.await_args.args[1][0]["spend"] == 0.25 + assert logger not in litellm._async_success_callback + assert logger._flush_task is not None and logger._flush_task.done() + assert not logger._flush_task.cancelled() # --------------------------------------------------------------------------- diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index 6391c1577f3..2e34172acfd 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -2,6 +2,9 @@ Tests for the agent tracing endpoints (litellm/proxy/tracing_endpoints.py). """ +from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager +from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest @@ -9,58 +12,79 @@ from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient from litellm.proxy import tracing_endpoints -from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy._types import LitellmUserRoles, ProxyLifespanState, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.tracing_runtime import manage_tracing, provide_storage +from litellm.rust_bridge.traces import ClickHouseStorage from litellm.tracing import TraceReceiver, TracingPayloadTooLargeError -from litellm.tracing.store import ClickHouseTraceStore +from litellm.tracing.store import TraceStore +from litellm.tracing.types import TraceScope TEAM_KEY = UserAPIKeyAuth( token="hashed-key", team_id="team-research", org_id="org-1", user_role=LitellmUserRoles.INTERNAL_USER ) -# ---------------------------------------------------------------- scope / tenant +@pytest.mark.parametrize( + ("auth", "scope", "can_write"), + ( + pytest.param( + UserAPIKeyAuth(token="admin-key", team_id="team-a", user_role=LitellmUserRoles.PROXY_ADMIN), + TraceScope(team_ids=(), api_key_hash=""), + True, + id="admin", + ), + pytest.param( + UserAPIKeyAuth(token="view-key", team_id="team-a", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY), + TraceScope(team_ids=(), api_key_hash=""), + False, + id="view-only-admin", + ), + pytest.param( + TEAM_KEY, + TraceScope(team_ids=("team-research",), api_key_hash=""), + True, + id="team-key", + ), + pytest.param( + UserAPIKeyAuth(token="hashed-key", user_role=LitellmUserRoles.INTERNAL_USER), + TraceScope(team_ids=("",), api_key_hash="hashed-key"), + True, + id="teamless-key", + ), + ), +) +def test_trace_read_and_write_permissions( + client: TestClient, receiver: MagicMock, auth: UserAPIKeyAuth, scope: TraceScope, can_write: bool +) -> None: + client.app.dependency_overrides[user_api_key_auth] = lambda: auth + read: Final = client.get("/v1/traces?start_ms=1&end_ms=2") + assert read.status_code == 200, read.text + receiver.list_traces.assert_awaited_once_with(scope=scope, start_ms=1, end_ms=2, cursor=None) -def test_scope_for_admin_sees_everything(): - for role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY): - auth = UserAPIKeyAuth(token="k", team_id="team-a", user_role=role) - assert tracing_endpoints.scope_for(auth) == {"team_ids": (), "api_key_hash": ""} - - -def test_scope_for_team_key_sees_its_team(): - assert tracing_endpoints.scope_for(TEAM_KEY) == {"team_ids": ("team-research",), "api_key_hash": ""} - - -def test_scope_for_teamless_key_sees_only_its_own_traces(): - auth = UserAPIKeyAuth(token="hashed-key", user_role=LitellmUserRoles.INTERNAL_USER) - assert tracing_endpoints.scope_for(auth) == {"team_ids": ("",), "api_key_hash": "hashed-key"} - - -def test_scope_for_no_team_no_token_is_forbidden(): - with pytest.raises(HTTPException) as e: - tracing_endpoints.scope_for(UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER)) - assert e.value.status_code == 403 - - -def test_tenant_for_comes_from_auth(): - tenant = tracing_endpoints.tenant_for(TEAM_KEY) - assert (tenant.team_id, tenant.api_key_hash, tenant.org_id) == ("team-research", "hashed-key", "org-1") - blank = tracing_endpoints.tenant_for(UserAPIKeyAuth()) - assert (blank.team_id, blank.api_key_hash, blank.org_id) == ("", "", "") - - -# ---------------------------------------------------------------- endpoints + write: Final = client.post("/v1/traces", json={}) + assert write.status_code == (200 if can_write else 403), write.text + if not can_write: + receiver.ingest.assert_not_awaited() + return + receiver.ingest.assert_awaited_once() + tenant: Final = receiver.ingest.await_args.kwargs["tenant"] + assert (tenant.team_id, tenant.api_key_hash, tenant.org_id) == ( + auth.team_id or "", + auth.token or "", + auth.org_id or "", + ) @pytest.fixture -def receiver(monkeypatch) -> MagicMock: +def receiver(client) -> MagicMock: fake = MagicMock() fake.ingest = AsyncMock(return_value=1) fake.list_traces = AsyncMock(return_value={"data": [], "next_cursor": None}) fake.get_trace = AsyncMock(return_value=None) fake.get_span = AsyncMock(return_value=None) - monkeypatch.setattr(tracing_endpoints, "receiver", fake) + client.app.dependency_overrides[tracing_endpoints.provide_receiver] = lambda: fake return fake @@ -72,8 +96,7 @@ def client() -> TestClient: return TestClient(app) -def test_501_when_tracing_not_enabled(client, monkeypatch): - monkeypatch.setattr(tracing_endpoints, "receiver", None) +def test_501_when_tracing_not_enabled(client): assert client.post("/v1/traces", content=b"").status_code == 501 assert client.get("/v1/traces").status_code == 501 @@ -149,13 +172,13 @@ def test_get_span_404_and_200(client, receiver): receiver.get_span.assert_awaited_with("t1", "s1", {"team_ids": ("team-research",), "api_key_hash": ""}, "") -def test_get_span_serves_ui_content_from_stored_payloads(client, monkeypatch): +def test_get_span_serves_ui_content_from_stored_payloads(client): storage = MagicMock() stored_output = '{"role": "ai", "content": "", "tool_calls": [{"name": "lookup", "args": {"id": 7}}]}' storage.query = AsyncMock( return_value=[{"span_id": "s1", "input": '{"city": "Paris"}', "output": stored_output, "attributes": {}}] ) - monkeypatch.setattr(tracing_endpoints, "receiver", TraceReceiver(ClickHouseTraceStore(storage))) + client.app.dependency_overrides[tracing_endpoints.provide_receiver] = lambda: TraceReceiver(TraceStore(storage)) body = client.get("/v1/traces/t1/spans/s1").json() assert body["output"] == stored_output assert body["input_ui"] == {"kind": "fields", "fields": [{"key": "city", "value": "Paris"}]} @@ -197,3 +220,259 @@ def test_view_only_admin_cannot_ingest_traces(client, receiver): response = client.post("/v1/traces", content=b"{}") assert response.status_code == 403 receiver.ingest.assert_not_called() + + +@pytest.mark.parametrize("status_code", [401, 403]) +def test_auth_failure_precedes_disabled_receiver(client: TestClient, status_code: int) -> None: + def unavailable() -> None: + return None + + def authenticate() -> UserAPIKeyAuth: + if status_code == 401: + raise HTTPException(status_code=401, detail="Invalid API key") + return UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + + client.app.dependency_overrides[user_api_key_auth] = authenticate + client.app.dependency_overrides[tracing_endpoints.provide_receiver] = unavailable + response: Final = client.post("/v1/traces", content=b"{}") + assert response.status_code == status_code + assert response.json() == { + "detail": "Invalid API key" if status_code == 401 else "Not allowed to ingest agent traces" + } + + +def test_disabled_receiver_precedes_read_scope_rejection(client: TestClient) -> None: + client.app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER + ) + response: Final = client.get("/v1/traces") + assert response.status_code == 501 + assert response.json() == { + "detail": "Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL." + } + + +@pytest.mark.requires_rust_extension +def test_injected_receiver_persists_authenticated_tenant(client: TestClient) -> None: + storage: Final = MagicMock(spec=ClickHouseStorage) + storage.insert_rows = AsyncMock() + tracing: Final = TraceReceiver(TraceStore(storage)) + client.app.dependency_overrides[tracing_endpoints.provide_receiver] = lambda: tracing + response: Final = client.post( + "/v1/traces", + json={ + "resourceSpans": [ + { + "resource": { + "attributes": [ + {"key": "litellm.team_id", "value": {"stringValue": "spoofed-team"}}, + {"key": "litellm.api_key_hash", "value": {"stringValue": "spoofed-key"}}, + {"key": "litellm.org_id", "value": {"stringValue": "spoofed-org"}}, + ] + }, + "scopeSpans": [ + { + "spans": [ + { + "traceId": "01" * 16, + "spanId": "02" * 8, + "name": "dependency-injection", + "startTimeUnixNano": "1000000000", + "endTimeUnixNano": "1000000001", + } + ] + } + ], + } + ], + }, + ) + assert response.status_code == 200, response.text + assert response.json() == {} + storage.insert_rows.assert_awaited_once() + table, rows = storage.insert_rows.await_args.args + assert table == "otel_traces" + assert len(rows) == 1 + assert rows[0]["TeamId"] == TEAM_KEY.team_id + assert rows[0]["ApiKeyHash"] == TEAM_KEY.token + assert rows[0]["ResourceAttributes"] == { + "litellm.team_id": TEAM_KEY.team_id, + "litellm.api_key_hash": TEAM_KEY.token, + "litellm.org_id": TEAM_KEY.org_id, + } + + +def test_lifespan_receivers_are_app_local() -> None: + first_storage: Final = MagicMock(spec=ClickHouseStorage) + first_storage.query = AsyncMock( + return_value=[ + { + "span_id": "first-span", + "input": "first-input", + "output": "", + "attributes": {}, + } + ] + ) + second_storage: Final = MagicMock(spec=ClickHouseStorage) + second_storage.query = AsyncMock( + return_value=[ + { + "span_id": "second-span", + "input": "second-input", + "output": "", + "attributes": {}, + } + ] + ) + first_receiver: Final = TraceReceiver(TraceStore(first_storage)) + second_receiver: Final = TraceReceiver(TraceStore(second_storage)) + first_storage.ensure_schema = AsyncMock() + second_storage.ensure_schema = AsyncMock() + + @asynccontextmanager + async def first_lifespan(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]: + async with manage_tracing(True, lambda: first_receiver) as receiver: + state: Final[ProxyLifespanState] = {"tracing_receiver": receiver} + yield state + + @asynccontextmanager + async def second_lifespan(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]: + async with manage_tracing(True, lambda: second_receiver) as receiver: + state: Final[ProxyLifespanState] = {"tracing_receiver": receiver} + yield state + + first_app: Final = FastAPI(lifespan=first_lifespan) + second_app: Final = FastAPI(lifespan=second_lifespan) + first_app.include_router(tracing_endpoints.router) + second_app.include_router(tracing_endpoints.router) + first_app.dependency_overrides[user_api_key_auth] = lambda: TEAM_KEY + second_app.dependency_overrides[user_api_key_auth] = lambda: TEAM_KEY + + with TestClient(first_app) as first_client: + with TestClient(second_app) as second_client: + second_response: Final = second_client.get("/v1/traces/t1/spans/second-span?trace_ref=second-run") + simultaneous: Final = first_client.get("/v1/traces/t1/spans/first-span?trace_ref=first-run") + first_response: Final = first_client.get("/v1/traces/t1/spans/first-span?trace_ref=first-run") + assert simultaneous.json() == first_response.json() + first_storage.ensure_schema.assert_awaited_once() + second_storage.ensure_schema.assert_awaited_once() + + assert first_response.status_code == second_response.status_code == 200 + assert first_response.json() == { + "span_id": "first-span", + "input": "first-input", + "output": "", + "attributes": {}, + "input_ui": {"kind": "text", "text": "first-input"}, + "output_ui": {"kind": "text", "text": ""}, + } + assert second_response.json() == { + "span_id": "second-span", + "input": "second-input", + "output": "", + "attributes": {}, + "input_ui": {"kind": "text", "text": "second-input"}, + "output_ui": {"kind": "text", "text": ""}, + } + assert first_storage.query.await_count == 2 + first_storage.query.assert_awaited_with( + "span_detail", + { + "team_ids": (TEAM_KEY.team_id,), + "api_key_hash": "", + "trace_id": "t1", + "span_id": "first-span", + "trace_ref": "first-run", + }, + ) + second_storage.query.assert_awaited_once_with( + "span_detail", + { + "team_ids": (TEAM_KEY.team_id,), + "api_key_hash": "", + "trace_id": "t1", + "span_id": "second-span", + "trace_ref": "second-run", + }, + ) + + +@pytest.mark.parametrize("auth", [TEAM_KEY, UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER)]) +def test_query_validation_precedes_trace_access_checks(client: TestClient, auth: UserAPIKeyAuth) -> None: + client.app.dependency_overrides[user_api_key_auth] = lambda: auth + response: Final = client.get("/v1/traces", params={"start_ms": "invalid"}) + assert response.status_code == 422 + assert response.json()["detail"][0]["loc"] == ["query", "start_ms"] + + +@pytest.mark.parametrize("enabled", [True, False]) +def test_unavailable_lifespan_receiver_returns_501(enabled: bool) -> None: + storage: Final = MagicMock(spec=ClickHouseStorage) + storage.ensure_schema = AsyncMock(side_effect=RuntimeError("storage unavailable")) + tracing: Final = TraceReceiver(TraceStore(storage)) + + @asynccontextmanager + async def lifespan(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]: + async with manage_tracing(enabled, lambda: tracing) as receiver: + state: Final[ProxyLifespanState] = {"tracing_receiver": receiver} + yield state + + app: Final = FastAPI(lifespan=lifespan) + app.include_router(tracing_endpoints.router) + app.dependency_overrides[user_api_key_auth] = lambda: TEAM_KEY + with TestClient(app) as client: + response: Final = client.get("/v1/traces") + assert response.status_code == 501 + assert storage.ensure_schema.await_count == int(enabled) + storage.query.assert_not_called() + + +def test_lens_reads_from_the_lifespan_storage() -> None: + from litellm.proxy.lens.endpoints import router as lens_router + + storage: Final = MagicMock(spec=ClickHouseStorage) + storage.ensure_schema = AsyncMock() + storage.lens_sample = AsyncMock(return_value=[]) + tracing: Final = TraceReceiver(TraceStore(storage)) + + @asynccontextmanager + async def lifespan(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]: + async with manage_tracing(True, lambda: tracing) as receiver: + state: Final[ProxyLifespanState] = {"tracing_receiver": receiver} + yield state + + app: Final = FastAPI(lifespan=lifespan) + app.include_router(lens_router) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + with TestClient(app) as client: + response: Final = client.post( + "/lens/preview/sample", + json={"settings": {"name": "Review", "model": "analysis", "context": "Find failed executions"}}, + ) + assert response.status_code == 200, response.text + assert response.json()["executions"] == [] + storage.lens_sample.assert_awaited_once() + assert storage.lens_sample.await_args.args[0]["all_teams"] == 1 + + +def test_lens_reads_from_injected_storage_without_receiver() -> None: + from litellm.proxy.lens.endpoints import router as lens_router + from litellm.proxy.lens.sources import Storage + + storage: Final = MagicMock(spec=Storage) + storage.lens_sample = AsyncMock(return_value=[]) + app: Final = FastAPI() + app.include_router(lens_router) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + app.dependency_overrides[provide_storage] = lambda: storage + + with TestClient(app) as client: + response: Final = client.post( + "/lens/preview/sample", + json={"settings": {"name": "Review", "model": "analysis", "context": "Find failed executions"}}, + ) + + assert response.status_code == 200, response.text + assert response.json()["executions"] == [] + storage.lens_sample.assert_awaited_once() From ec605826d4ebb69a0c3c79604eee6bc148cde4b9 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Thu, 1 Oct 2026 13:45:33 -0700 Subject: [PATCH 011/203] feat: improve trace ingestion and trace details (#43975) * refactor: separate OTLP HTTP decoding from trace codec * feat: complete trace ingestion and read paths * fix: encode OTLP protobuf errors in Rust * fix: raise OTLP body limit to 16 MiB * test: cover OTLP auth body parsing boundary * refactor: parse OTLP media type into enum * fix: enforce OTLP body size at HTTP boundary * perf: preserve shared OTLP metadata across ingestion * bench: compare owned and shared trace resource fanout * refactor: extract shared storage and Python conversion caches * refactor: keep shared storage owned by traces * test: keep trace loopback coverage in Rust * test(proxy): adapt trace coverage to injected access context * fix(tracing): satisfy stacked branch lint checks * refactor(tracing): use immutable ingestion payloads * fix(tracing): declare native error encoder export * test(proxy): resolve trace access through dependency * fix(tracing): align merged normalizer types and bridge tests * fix(tracing): address ingestion and diagnostic review findings * fix(proxy): preserve body parsing for partial request scopes * test(proxy): use valid HTTP scopes in request fixtures * test(proxy): complete auth request flow scopes --- litellm-rust/Cargo.lock | 7 + litellm-rust/Cargo.toml | 3 + .../crates/cache-azure-blob/Cargo.toml | 2 +- litellm-rust/crates/cache-gcs/Cargo.toml | 2 +- litellm-rust/crates/cache-response/Cargo.toml | 2 +- litellm-rust/crates/cache-s3/Cargo.toml | 2 +- litellm-rust/crates/core/Cargo.toml | 2 +- .../crates/gateway-inference/Cargo.toml | 2 +- .../host-python/src/conversion_cache.rs | 57 +++ litellm-rust/crates/host-python/src/lib.rs | 2 + .../host-python/tests/conversion_cache.rs | 121 ++++++ litellm-rust/crates/python-bridge/Cargo.toml | 3 +- litellm-rust/crates/python-bridge/src/lib.rs | 3 +- .../crates/python-bridge/src/routes/traces.rs | 162 +++++++- litellm-rust/crates/secrets-aws/Cargo.toml | 2 +- litellm-rust/crates/secrets-azure/Cargo.toml | 2 +- .../crates/secrets-cyberark/Cargo.toml | 2 +- litellm-rust/crates/secrets-google/Cargo.toml | 2 +- .../crates/secrets-hashicorp/Cargo.toml | 2 +- litellm-rust/crates/secrets/Cargo.toml | 2 +- .../crates/storage-clickhouse/src/insert.rs | 17 + .../crates/storage-clickhouse/src/lib.rs | 2 +- litellm-rust/crates/traces/Cargo.toml | 13 +- .../crates/traces/benches/resource-fanout.rs | 39 ++ .../crates/traces/query/span_error.sql | 13 + .../crates/traces/query/trace_spans.sql | 5 +- litellm-rust/crates/traces/src/error.rs | 2 +- litellm-rust/crates/traces/src/insert.rs | 279 +++++++++++--- litellm-rust/crates/traces/src/lib.rs | 4 +- litellm-rust/crates/traces/src/otlp.rs | 221 ----------- .../crates/traces/src/otlp/attributes.rs | 101 +++++ litellm-rust/crates/traces/src/otlp/limits.rs | 212 +++++++++++ litellm-rust/crates/traces/src/otlp/mod.rs | 42 +++ litellm-rust/crates/traces/src/otlp/span.rs | 166 +++++++++ litellm-rust/crates/traces/src/otlp/wire.rs | 43 +++ litellm-rust/crates/traces/src/shared.rs | 46 +++ litellm-rust/crates/traces/src/sql.rs | 3 + litellm-rust/crates/traces/tests/insert.rs | 105 +++++- .../crates/traces/tests/migrations.rs | 145 ++++++++ litellm-rust/crates/traces/tests/otlp.rs | 350 ++++++++++++++++-- litellm-rust/crates/traces/tests/shared.rs | 23 ++ litellm/constants.py | 4 +- .../proxy/common_utils/http_parsing_utils.py | 14 +- litellm/proxy/proxy_server.py | 21 ++ litellm/proxy/tracing_endpoints.py | 67 +++- litellm/rust_bridge/_native.pyi | 8 +- litellm/rust_bridge/traces.py | 23 +- litellm/tracing/decode.py | 319 +++++++++++++--- litellm/tracing/receiver.py | 111 +++++- litellm/tracing/store.py | 56 ++- litellm/tracing/types.py | 20 +- .../fixtures/langsmith_deep_agent_export.json | 60 +-- tests/test_litellm/tracing/test_decode.py | 100 ++++- tests/test_litellm/tracing/test_receiver.py | 69 +++- tests/test_litellm/tracing/test_store.py | 44 +++ tests/test_litellm_rust/test_traces.py | 115 +++++- .../test_otel_exception_handler.py | 15 +- .../unit/proxy/auth/test_user_api_key_auth.py | 20 +- .../test_user_api_key_auth_request_flow.py | 6 + .../common_utils/test_http_parsing_utils.py | 101 ++++- .../proxy_server/test_exception_handlers.py | 49 ++- tests/unit/proxy/test_proxy_reject_logging.py | 1 + tests/unit/proxy/test_proxy_server.py | 6 +- tests/unit/proxy/test_tracing_endpoints.py | 39 +- .../src/components/networking.tsx | 13 +- .../view_logs/TraceView/DetailContent.tsx | 67 +++- ...st.tsx => DetailPane.integration.test.tsx} | 33 +- .../view_logs/TraceView/traceTypes.ts | 8 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 63 ++++ 69 files changed, 3096 insertions(+), 569 deletions(-) create mode 100644 litellm-rust/crates/host-python/src/conversion_cache.rs create mode 100644 litellm-rust/crates/host-python/tests/conversion_cache.rs create mode 100644 litellm-rust/crates/traces/benches/resource-fanout.rs create mode 100644 litellm-rust/crates/traces/query/span_error.sql delete mode 100644 litellm-rust/crates/traces/src/otlp.rs create mode 100644 litellm-rust/crates/traces/src/otlp/attributes.rs create mode 100644 litellm-rust/crates/traces/src/otlp/limits.rs create mode 100644 litellm-rust/crates/traces/src/otlp/mod.rs create mode 100644 litellm-rust/crates/traces/src/otlp/span.rs create mode 100644 litellm-rust/crates/traces/src/otlp/wire.rs create mode 100644 litellm-rust/crates/traces/src/shared.rs create mode 100644 litellm-rust/crates/traces/tests/shared.rs rename ui/litellm-dashboard/src/components/view_logs/TraceView/{DetailPane.test.tsx => DetailPane.integration.test.tsx} (89%) diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 57e9e803e4a..ff0eafee47e 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -4090,6 +4090,7 @@ dependencies = [ "litellm-token-counter", "litellm-traces", "litellm-tracing", + "prost", "pyo3", "pyo3-async-runtimes", "qdrant-client", @@ -4384,6 +4385,7 @@ name = "litellm-traces" version = "0.1.0" dependencies = [ "base64 0.22.1", + "criterion", "flate2", "litellm-http", "litellm-storage-clickhouse", @@ -4393,10 +4395,12 @@ dependencies = [ "serde", "serde_json", "sha2 0.10.9", + "strum", "testcontainers-modules", "thiserror 2.0.19", "time", "tokio", + "wiremock", ] [[package]] @@ -4819,6 +4823,7 @@ dependencies = [ "js-sys", "pin-project-lite", "thiserror 2.0.19", + "tracing", ] [[package]] @@ -4833,6 +4838,8 @@ dependencies = [ "opentelemetry_sdk 0.33.0", "prost", "serde", + "tonic", + "tonic-prost", ] [[package]] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 450253ea768..8d837c2d31b 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -82,6 +82,7 @@ reqwest = { version = "0.12", default-features = false, features = ["json", "mul qdrant-client = { version = "1.19.0", default-features = false } uuid = { version = "1", features = ["v4"] } rstest = "0.26.1" +wiremock = "0.6.5" rstest_reuse = "0.7.0" rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"] } rustify = "=0.7.0" @@ -116,6 +117,8 @@ time = { version = "0.3.53", features = ["parsing"] } criterion = "0.8.2" fancy-regex = "0.19.2" veil = "0.3.0" +prost = "0.14.4" +opentelemetry-proto = "0.33" [profile.release] opt-level = 3 diff --git a/litellm-rust/crates/cache-azure-blob/Cargo.toml b/litellm-rust/crates/cache-azure-blob/Cargo.toml index 5bdfa16ef53..c28cb90d84a 100644 --- a/litellm-rust/crates/cache-azure-blob/Cargo.toml +++ b/litellm-rust/crates/cache-azure-blob/Cargo.toml @@ -26,4 +26,4 @@ litellm-cache-testing.workspace = true rstest.workspace = true serde_json.workspace = true tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/cache-gcs/Cargo.toml b/litellm-rust/crates/cache-gcs/Cargo.toml index 1a06683e615..91630879cbe 100644 --- a/litellm-rust/crates/cache-gcs/Cargo.toml +++ b/litellm-rust/crates/cache-gcs/Cargo.toml @@ -21,4 +21,4 @@ litellm-cache-testing.workspace = true rstest.workspace = true serde_json.workspace = true tokio.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/cache-response/Cargo.toml b/litellm-rust/crates/cache-response/Cargo.toml index 1379573e505..869c40a12ab 100644 --- a/litellm-rust/crates/cache-response/Cargo.toml +++ b/litellm-rust/crates/cache-response/Cargo.toml @@ -21,4 +21,4 @@ redis = "1.7.0" redis-test = "1.0.4" rstest.workspace = true tokio.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/cache-s3/Cargo.toml b/litellm-rust/crates/cache-s3/Cargo.toml index 680f2da8215..eb3a2fff1ac 100644 --- a/litellm-rust/crates/cache-s3/Cargo.toml +++ b/litellm-rust/crates/cache-s3/Cargo.toml @@ -23,6 +23,6 @@ tokio.workspace = true litellm-http = { workspace = true, features = ["test-support"] } litellm-cache-testing.workspace = true rstest.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true serde_json.workspace = true tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 8410aff1d6a..85362fd90d2 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -47,4 +47,4 @@ litellm-host-native.workspace = true litellm-llms = { workspace = true, features = ["test-support"] } rstest.workspace = true rstest_reuse.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/gateway-inference/Cargo.toml b/litellm-rust/crates/gateway-inference/Cargo.toml index c854f0ea1ad..4544152d059 100644 --- a/litellm-rust/crates/gateway-inference/Cargo.toml +++ b/litellm-rust/crates/gateway-inference/Cargo.toml @@ -30,4 +30,4 @@ futures-util.workspace = true tokio = { workspace = true, features = ["io-util"] } rstest.workspace = true tower = { version = "0.5.3", features = ["util"] } -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/host-python/src/conversion_cache.rs b/litellm-rust/crates/host-python/src/conversion_cache.rs new file mode 100644 index 00000000000..c78ed42bcee --- /dev/null +++ b/litellm-rust/crates/host-python/src/conversion_cache.rs @@ -0,0 +1,57 @@ +use std::collections::{HashMap, hash_map::Entry}; + +use pyo3::prelude::*; + +pub struct ToPythonCache<'a, 'py, T> { + entries: HashMap)>, +} + +impl Default for ToPythonCache<'_, '_, T> { + fn default() -> Self { + Self { + entries: HashMap::new(), + } + } +} + +impl<'a, 'py, T> ToPythonCache<'a, 'py, T> { + pub fn get_or_try_insert_with( + &mut self, + value: &'a T, + convert: impl FnOnce(&'a T) -> PyResult>, + ) -> PyResult<&Bound<'py, PyAny>> { + let identity = std::ptr::from_ref(value) as usize; + let entry = match self.entries.entry(identity) { + Entry::Occupied(entry) => entry.into_mut(), + Entry::Vacant(entry) => entry.insert((value, convert(value)?)), + }; + Ok(&entry.1) + } +} + +pub struct FromPythonCache<'py, T> { + entries: HashMap, T)>, +} + +impl Default for FromPythonCache<'_, T> { + fn default() -> Self { + Self { + entries: HashMap::new(), + } + } +} + +impl<'py, T> FromPythonCache<'py, T> { + pub fn get_or_try_insert_with( + &mut self, + value: &Bound<'py, PyAny>, + convert: impl FnOnce(&Bound<'py, PyAny>) -> PyResult, + ) -> PyResult<&T> { + let identity = value.as_ptr() as usize; + let entry = match self.entries.entry(identity) { + Entry::Occupied(entry) => entry.into_mut(), + Entry::Vacant(entry) => entry.insert((value.clone(), convert(value)?)), + }; + Ok(&entry.1) + } +} diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 4de404e3624..00543f64085 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -5,6 +5,7 @@ mod argument; mod binding; +mod conversion_cache; mod driver; mod error; mod file_reader; @@ -20,6 +21,7 @@ mod services; pub use argument::lookup; pub use binding::PythonBinding; +pub use conversion_cache::{FromPythonCache, ToPythonCache}; pub use driver::{CallOptions, run_call}; pub use error::{InvokeError, missing_state}; pub use file_reader::{FileContent, PythonFileReader, py_bytes}; diff --git a/litellm-rust/crates/host-python/tests/conversion_cache.rs b/litellm-rust/crates/host-python/tests/conversion_cache.rs new file mode 100644 index 00000000000..70ad838e001 --- /dev/null +++ b/litellm-rust/crates/host-python/tests/conversion_cache.rs @@ -0,0 +1,121 @@ +use std::{cell::Cell, rc::Rc}; + +use litellm_host_python::{FromPythonCache, Pythonized, ToPythonCache}; +use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; +use rstest::{fixture, rstest}; + +#[fixture] +fn python() { + Python::initialize(); +} + +#[rstest] +fn rust_identity_reuses_python_objects_without_merging_equal_values(#[from(python)] _python: ()) { + Python::attach(|py| { + let original = Rc::new(vec![1, 2]); + let cloned = original.clone(); + let equal = Rc::new(vec![1, 2]); + let mut cache = ToPythonCache::default(); + let first = cache + .get_or_try_insert_with(original.as_ref(), |value| { + Pythonized(value).into_pyobject(py) + }) + .unwrap() + .clone(); + let second = cache + .get_or_try_insert_with(cloned.as_ref(), |_| panic!("must reuse conversion")) + .unwrap() + .clone(); + let third = cache + .get_or_try_insert_with(equal.as_ref(), |value| Pythonized(value).into_pyobject(py)) + .unwrap(); + assert!(first.is(&second)); + assert!(!first.is(third)); + assert!(first.eq(third).unwrap()); + }); +} + +#[rstest] +fn python_identity_reuses_rust_values_without_merging_equal_objects(#[from(python)] _python: ()) { + Python::attach(|py| { + let original = PyDict::new(py); + original.set_item("value", 1).unwrap(); + let equal = original.copy().unwrap(); + let calls = Cell::new(0); + let mut cache = FromPythonCache::default(); + let convert = |value: &Bound<'_, PyAny>| { + calls.set(calls.get() + 1); + value.get_item("value")?.extract::().map(Rc::new) + }; + let first = cache + .get_or_try_insert_with(original.as_any(), convert) + .unwrap() + .clone(); + let second = cache + .get_or_try_insert_with(original.as_any(), convert) + .unwrap() + .clone(); + let third = cache + .get_or_try_insert_with(equal.as_any(), convert) + .unwrap(); + assert!(Rc::ptr_eq(&first, &second)); + assert!(!Rc::ptr_eq(&first, third)); + assert_eq!(&first, third); + assert_eq!(calls.get(), 2); + }); +} + +#[rstest] +fn python_sources_stay_alive_until_the_cache_is_dropped(#[from(python)] _python: ()) { + Python::attach(|py| { + let value = py + .eval(pyo3::ffi::c_str!("type('Tracked', (), {})()"), None, None) + .unwrap(); + let weak = py + .import("weakref") + .unwrap() + .call_method1("ref", (&value,)) + .unwrap(); + let mut cache = FromPythonCache::default(); + cache.get_or_try_insert_with(&value, |_| Ok(42)).unwrap(); + drop(value); + assert!(!weak.call0().unwrap().is_none()); + drop(cache); + assert!(weak.call0().unwrap().is_none()); + }); +} + +#[rstest] +#[case::to_python(true)] +#[case::from_python(false)] +fn failed_conversions_preserve_exceptions_and_can_be_retried( + #[from(python)] _python: (), + #[case] to_python: bool, +) { + Python::attach(|py| { + let failure = PyValueError::new_err("conversion failed"); + if to_python { + let source = vec![1, 2]; + let mut cache = ToPythonCache::default(); + let error = cache + .get_or_try_insert_with(&source, |_| Err(failure.clone_ref(py))) + .unwrap_err(); + assert!(error.value(py).is(failure.value(py))); + let result = cache + .get_or_try_insert_with(&source, |value| Pythonized(value).into_pyobject(py)) + .unwrap(); + assert_eq!(result.extract::>().unwrap(), source); + } else { + let source = PyDict::new(py).into_any(); + let mut cache = FromPythonCache::default(); + let error = cache + .get_or_try_insert_with(&source, |_| Err(failure.clone_ref(py))) + .unwrap_err(); + assert!(error.value(py).is(failure.value(py))); + assert_eq!( + *cache.get_or_try_insert_with(&source, |_| Ok(42)).unwrap(), + 42 + ); + } + }); +} diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index a1d1f63d6f3..2b505f08eca 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -52,6 +52,7 @@ litellm-llms-types.workspace = true litellm-host-python.workspace = true litellm-token-counter = { path = "../token-counter", default-features = false } pyo3.workspace = true +prost.workspace = true pyo3-async-runtimes.workspace = true reqwest.workspace = true redis = { version = "1.7.0", features = ["tls-rustls"] } @@ -73,7 +74,7 @@ futures-util.workspace = true rstest.workspace = true sha2.workspace = true tokio-tungstenite.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true aws-sdk-secretsmanager = "1.117.0" [[bench]] diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 0d4df996552..d269fa4015f 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -44,7 +44,7 @@ mod _native { #[pymodule_export] use crate::routes::token_counter::TokenCounter; #[pymodule_export] - use crate::routes::traces::{NativeTraceStorage, trace_decode_otlp}; + use crate::routes::traces::{NativeTraceStorage, trace_decode_otlp, trace_encode_error}; #[cfg(feature = "huggingface")] #[pymodule_export] use crate::tokenizer::HuggingFaceEncoding; @@ -111,6 +111,7 @@ mod tests { "NativeDiagnosticProcessor", "NativeTraceStorage", "trace_decode_otlp", + "trace_encode_error", "TokenCounter", "Tokenizer", "gil_stats", diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 47c924f4842..ca66e2e46be 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -1,13 +1,33 @@ use std::collections::BTreeMap; +use litellm_host_python::{FromPythonCache, ToPythonCache}; use litellm_http::ClientVariant; use litellm_storage_clickhouse::Storage; -use litellm_traces::{Error, InsertTable, Parameter, ReadQuery}; +use litellm_traces::{Error, InsertTable, Parameter, ReadQuery, Shared}; +use prost::Message; use pyo3::{ exceptions::{PyOverflowError, PyRuntimeError, PyValueError}, prelude::*, + types::{PyBytes, PyDict, PyList, PyMapping, PyString}, }; +#[derive(Message)] +struct OtlpErrorStatus { + #[prost(int32, tag = "1")] + code: i32, + #[prost(string, tag = "2")] + message: String, +} + +#[pyfunction] +pub fn trace_encode_error<'py>(py: Python<'py>, message: &str) -> Bound<'py, PyBytes> { + let status = OtlpErrorStatus { + code: 0, + message: message.to_owned(), + }; + PyBytes::new(py, &status.encode_to_vec()) +} + fn map_error(error: Error) -> PyErr { match error { Error::InvalidRow @@ -71,9 +91,7 @@ impl NativeTraceStorage { &self, py: Python<'py>, table: &str, - #[pyo3(from_py_with = litellm_host_python::from_py_argument)] rows: Vec< - BTreeMap, - >, + #[pyo3(from_py_with = insert_rows_from_py)] rows: Vec, ) -> PyResult> { let table = InsertTable::parse(table).map_err(map_error)?; let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; @@ -82,7 +100,8 @@ impl NativeTraceStorage { crate::execution::run_async( py, async move { - litellm_traces::insert_rows(&client, &connection, &database, table, rows).await + litellm_traces::insert_shared_rows(&client, &connection, &database, table, rows) + .await }, map_error, ) @@ -140,21 +159,132 @@ pub fn trace_decode_otlp<'py>( py: Python<'py>, body: &[u8], content_type: Option<&str>, - content_encoding: Option<&str>, - max_decompressed_bytes: usize, ) -> PyResult> { let spans = py - .detach(|| { - litellm_traces::decode_otlp( - body, - content_type, - content_encoding, - max_decompressed_bytes, - ) - }) + .detach(|| litellm_traces::decode_otlp(body, content_type)) .map_err(|error| match error { litellm_traces::DecodeError::TooLarge => PyOverflowError::new_err(error.to_string()), _ => PyValueError::new_err(error.to_string()), })?; - litellm_host_python::Pythonized(spans).into_pyobject(py) + spans_to_py(py, &spans).map(Bound::into_any) +} + +fn insert_rows_from_py(value: &Bound<'_, PyAny>) -> PyResult> { + let mut resources = FromPythonCache::default(); + value + .try_iter()? + .map(|row| { + let row = row?; + let mut fields = BTreeMap::new(); + for item in row.cast::()?.items()?.iter() { + let (key, value): (String, Bound<'_, PyAny>) = item.extract()?; + let converted = if matches!( + key.as_str(), + "ResourceAttributes" | "ScopeName" | "ScopeVersion" + ) { + resources + .get_or_try_insert_with(&value, |value| { + litellm_host_python::from_py_argument::(value) + .map(Shared::new) + })? + .clone() + } else { + Shared::new(litellm_host_python::from_py_argument(&value)?) + }; + fields.insert(key, converted); + } + Ok(fields) + }) + .collect() +} + +fn spans_to_py<'py>( + py: Python<'py>, + spans: &[litellm_traces::DecodedSpan], +) -> PyResult> { + let mut resources = ToPythonCache::default(); + let mut scopes = ToPythonCache::default(); + let result = PyList::empty(py); + for span in spans { + let resource = resources + .get_or_try_insert_with(span.resource_attributes.as_ref(), |value| { + litellm_host_python::Pythonized(value).into_pyobject(py) + })?; + let row = PyDict::new(py); + row.set_item("trace_id", &span.trace_id)?; + row.set_item("span_id", &span.span_id)?; + row.set_item("parent_span_id", &span.parent_span_id)?; + row.set_item("trace_state", &span.trace_state)?; + row.set_item("name", &span.name)?; + row.set_item("kind", &span.kind)?; + row.set_item("resource_attributes", resource)?; + for (key, value) in [ + ("scope_name", &span.scope_name), + ("scope_version", &span.scope_version), + ] { + let value = scopes.get_or_try_insert_with(value.as_ref(), |value| { + Ok(PyString::new(py, value).into_any()) + })?; + row.set_item(key, value)?; + } + row.set_item("attributes", &span.attributes)?; + row.set_item("start_ns", span.start_ns)?; + row.set_item("end_ns", span.end_ns)?; + row.set_item("status_code", &span.status_code)?; + row.set_item("status_message", &span.status_message)?; + row.set_item( + "events", + litellm_host_python::Pythonized(&span.events).into_pyobject(py)?, + )?; + result.append(row)?; + } + Ok(result) +} + +#[cfg(test)] +mod tests { + use super::*; + use rstest::rstest; + + #[rstest] + fn insert_projection_preserves_identity_without_merging_equal_resources() { + Python::initialize(); + Python::attach(|py| { + let resource = PyDict::new(py); + resource.set_item("service.name", "shared").unwrap(); + let equal_resource = resource.copy().unwrap(); + let rows = PyList::empty(py); + for value in [&resource, &resource, &equal_resource] { + let row = PyDict::new(py); + row.set_item("ResourceAttributes", value).unwrap(); + rows.append(row).unwrap(); + } + let projected = insert_rows_from_py(rows.as_any()).unwrap(); + assert!(Shared::shares_storage_with( + &projected[0]["ResourceAttributes"], + &projected[1]["ResourceAttributes"] + )); + assert!(!Shared::shares_storage_with( + &projected[0]["ResourceAttributes"], + &projected[2]["ResourceAttributes"] + )); + assert_eq!(projected[0], projected[2]); + }); + } + + #[rstest] + fn shared_conversion_preserves_every_decoded_field() { + Python::initialize(); + Python::attach(|py| { + let spans = litellm_traces::decode_otlp( + include_bytes!("../../../../../tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json"), + Some("application/json"), + ).unwrap(); + let expected = litellm_host_python::Pythonized(&spans) + .into_pyobject(py) + .unwrap(); + let actual = spans_to_py(py, &spans).unwrap(); + assert!(actual.eq(expected).unwrap()); + }); + } } diff --git a/litellm-rust/crates/secrets-aws/Cargo.toml b/litellm-rust/crates/secrets-aws/Cargo.toml index e7a394bd247..5d3bd413484 100644 --- a/litellm-rust/crates/secrets-aws/Cargo.toml +++ b/litellm-rust/crates/secrets-aws/Cargo.toml @@ -21,5 +21,5 @@ aws-credential-types = "1.3.0" base64.workspace = true rstest.workspace = true tokio.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true tempfile = "3" diff --git a/litellm-rust/crates/secrets-azure/Cargo.toml b/litellm-rust/crates/secrets-azure/Cargo.toml index efdf681e2bc..7ec03fb98da 100644 --- a/litellm-rust/crates/secrets-azure/Cargo.toml +++ b/litellm-rust/crates/secrets-azure/Cargo.toml @@ -20,7 +20,7 @@ percent-encoding = "2.3" [dev-dependencies] litellm-http = { workspace = true, features = ["test-support"] } -wiremock = "0.6.5" +wiremock.workspace = true rstest.workspace = true serde_json.workspace = true sha2.workspace = true diff --git a/litellm-rust/crates/secrets-cyberark/Cargo.toml b/litellm-rust/crates/secrets-cyberark/Cargo.toml index 0a91c61ade9..f630d5857d8 100644 --- a/litellm-rust/crates/secrets-cyberark/Cargo.toml +++ b/litellm-rust/crates/secrets-cyberark/Cargo.toml @@ -25,6 +25,6 @@ rcgen = "0.14.10" rstest.workspace = true tempfile = "3.27.0" tokio.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true serde.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/secrets-google/Cargo.toml b/litellm-rust/crates/secrets-google/Cargo.toml index 208b5ddd03f..3ce14fe7a12 100644 --- a/litellm-rust/crates/secrets-google/Cargo.toml +++ b/litellm-rust/crates/secrets-google/Cargo.toml @@ -28,4 +28,4 @@ reqwest.workspace = true litellm-http = { workspace = true, features = ["test-support"] } google-cloud-auth.workspace = true rstest.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/secrets-hashicorp/Cargo.toml b/litellm-rust/crates/secrets-hashicorp/Cargo.toml index c049ba127e5..7dd3d3c674f 100644 --- a/litellm-rust/crates/secrets-hashicorp/Cargo.toml +++ b/litellm-rust/crates/secrets-hashicorp/Cargo.toml @@ -21,4 +21,4 @@ veil.workspace = true rstest.workspace = true tempfile = "3" tokio.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/secrets/Cargo.toml b/litellm-rust/crates/secrets/Cargo.toml index f855a8a64a6..3655ce8bbc2 100644 --- a/litellm-rust/crates/secrets/Cargo.toml +++ b/litellm-rust/crates/secrets/Cargo.toml @@ -36,7 +36,7 @@ tokio = { workspace = true, features = ["fs"] } [dev-dependencies] litellm-http = { workspace = true, features = ["test-support"] } rstest.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true tempfile = "3" aws-sdk-kms = "1.120.0" google-cloud-kms-v1 = "1.14.0" diff --git a/litellm-rust/crates/storage-clickhouse/src/insert.rs b/litellm-rust/crates/storage-clickhouse/src/insert.rs index 3be73b086a0..81528ded907 100644 --- a/litellm-rust/crates/storage-clickhouse/src/insert.rs +++ b/litellm-rust/crates/storage-clickhouse/src/insert.rs @@ -26,6 +26,23 @@ pub async fn insert_encoded_rows( .write_all(encoded.as_bytes()) .map_err(|_| Error::InvalidRow)?; let body = encoder.finish().map_err(|_| Error::InvalidRow)?; + insert_compressed_rows(client, connection, database, table, token, body).await +} + +pub async fn insert_compressed_rows( + client: &Client, + connection: &Connection, + database: &str, + table: &str, + token: &str, + body: Vec, +) -> Result<(), Error> { + if !valid_identifier(database) { + return Err(Error::InvalidSchema); + } + if !valid_identifier(table) { + return Err(Error::InvalidTable); + } let mut url = connection.url().clone(); let existing_pairs: Vec<(String, String)> = url .query_pairs() diff --git a/litellm-rust/crates/storage-clickhouse/src/lib.rs b/litellm-rust/crates/storage-clickhouse/src/lib.rs index f2b34eddbf8..d11ee9d5cde 100644 --- a/litellm-rust/crates/storage-clickhouse/src/lib.rs +++ b/litellm-rust/crates/storage-clickhouse/src/lib.rs @@ -3,7 +3,7 @@ mod insert; mod read; pub use error::Error; -pub use insert::insert_encoded_rows; +pub use insert::{insert_compressed_rows, insert_encoded_rows}; pub use read::{Parameter, execute_read}; use url::Url; diff --git a/litellm-rust/crates/traces/Cargo.toml b/litellm-rust/crates/traces/Cargo.toml index b7f6e6ae52e..74de400764c 100644 --- a/litellm-rust/crates/traces/Cargo.toml +++ b/litellm-rust/crates/traces/Cargo.toml @@ -8,18 +8,25 @@ repository.workspace = true [dependencies] base64.workspace = true flate2.workspace = true -opentelemetry-proto = { version = "0.33.0", default-features = false, features = ["gen-tonic-messages", "trace", "with-serde"] } -prost = "0.14.4" +opentelemetry-proto = { workspace = true, features = ["gen-tonic-messages", "trace", "with-serde"] } +prost.workspace = true time = { workspace = true, features = ["formatting"] } litellm-http.workspace = true litellm-storage-clickhouse.workspace = true sha2.workspace = true -serde.workspace = true +serde = { workspace = true, features = ["rc"] } serde_json.workspace = true +strum.workspace = true thiserror.workspace = true [dev-dependencies] +criterion.workspace = true litellm-http = { workspace = true, features = ["test-support"] } rstest.workspace = true testcontainers-modules = { version = "0.15.0", features = ["clickhouse"] } tokio.workspace = true +wiremock.workspace = true + +[[bench]] +name = "resource-fanout" +harness = false diff --git a/litellm-rust/crates/traces/benches/resource-fanout.rs b/litellm-rust/crates/traces/benches/resource-fanout.rs new file mode 100644 index 00000000000..edf5d2eb055 --- /dev/null +++ b/litellm-rust/crates/traces/benches/resource-fanout.rs @@ -0,0 +1,39 @@ +use std::{collections::BTreeMap, hint::black_box, time::Duration}; + +use criterion::{BenchmarkId, Criterion, Throughput, criterion_group, criterion_main}; +use litellm_traces::Shared; + +fn fanout(resource: &T, spans: usize) -> Vec { + (0..spans).map(|_| resource.clone()).collect() +} + +fn resource_fanout(c: &mut Criterion) { + let mut group = c.benchmark_group("resource_fanout"); + for (attribute_bytes, spans) in [(256, 1), (256, 64), (8192, 1024), (16384, 1024)] { + let attributes = BTreeMap::from([ + ("service.name".to_owned(), "benchmark".to_owned()), + ("payload".to_owned(), "x".repeat(attribute_bytes)), + ]); + let owned = Box::new(attributes.clone()); + let shared = Shared::new(attributes); + let case = format!("{attribute_bytes}B_{spans}_spans"); + group.throughput(Throughput::Elements(spans as u64)); + group.bench_with_input(BenchmarkId::new("owned", &case), &owned, |b, resource| { + b.iter(|| black_box(fanout(black_box(resource), spans))); + }); + group.bench_with_input(BenchmarkId::new("shared", &case), &shared, |b, resource| { + b.iter(|| black_box(fanout(black_box(resource), spans))); + }); + } + group.finish(); +} + +criterion_group! { + name = benches; + config = Criterion::default() + .sample_size(20) + .warm_up_time(Duration::from_secs(1)) + .measurement_time(Duration::from_secs(2)); + targets = resource_fanout +} +criterion_main!(benches); diff --git a/litellm-rust/crates/traces/query/span_error.sql b/litellm-rust/crates/traces/query/span_error.sql new file mode 100644 index 00000000000..b4710006389 --- /dev/null +++ b/litellm-rust/crates/traces/query/span_error.sql @@ -0,0 +1,13 @@ +SELECT SpanId AS span_id, + substringUTF8(StatusMessage, {error_offset:UInt64} + 1, 16384) AS message, + lengthUTF8(StatusMessage) AS total_chars, + hex(SHA256(StatusMessage)) AS version +FROM otel_traces +WHERE TraceId = {trace_id:String} AND SpanId = {span_id:String} + AND (empty({team_ids:Array(String)}) OR TeamId IN {team_ids:Array(String)}) + AND ({api_key_hash:String} = '' OR ApiKeyHash = {api_key_hash:String}) + AND ({trace_ref:String} = '' OR + hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) = {trace_ref:String}) + AND ({error_version:String} = '' OR hex(SHA256(StatusMessage)) = {error_version:String}) +ORDER BY Timestamp, EngineReceivedMs, StatusMessage +LIMIT 1 diff --git a/litellm-rust/crates/traces/query/trace_spans.sql b/litellm-rust/crates/traces/query/trace_spans.sql index 409e6328198..dab3ac2e877 100644 --- a/litellm-rust/crates/traces/query/trace_spans.sql +++ b/litellm-rust/crates/traces/query/trace_spans.sql @@ -1,6 +1,7 @@ SELECT o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name, o.ObservationType AS type, o.AgentName AS agent, o.StatusCode AS status, - o.StatusMessage AS status_message, + substringUTF8(o.StatusMessage, 1, 128) AS status_message, + lengthUTF8(o.StatusMessage) > 128 AS error_truncated, toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns, o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model, o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens, @@ -12,5 +13,5 @@ WHERE o.TraceId = {trace_id:String} AND ({api_key_hash:String} = '' OR o.ApiKeyHash = {api_key_hash:String}) AND ({trace_ref:String} = '' OR hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) = {trace_ref:String}) -ORDER BY o.Timestamp +ORDER BY o.Timestamp, o.EngineReceivedMs, o.StatusMessage LIMIT 1 BY o.SpanId diff --git a/litellm-rust/crates/traces/src/error.rs b/litellm-rust/crates/traces/src/error.rs index 2ccfe0ea8d9..18fa4af9b53 100644 --- a/litellm-rust/crates/traces/src/error.rs +++ b/litellm-rust/crates/traces/src/error.rs @@ -2,6 +2,6 @@ pub enum DecodeError { #[error("invalid OTLP trace payload")] InvalidPayload, - #[error("OTLP trace payload exceeds the decompressed size limit")] + #[error("OTLP trace payload exceeds the decoding budget")] TooLarge, } diff --git a/litellm-rust/crates/traces/src/insert.rs b/litellm-rust/crates/traces/src/insert.rs index 6d9a2cab813..01a1eecfd7b 100644 --- a/litellm-rust/crates/traces/src/insert.rs +++ b/litellm-rust/crates/traces/src/insert.rs @@ -1,14 +1,23 @@ -use std::collections::BTreeMap; +use std::{ + borrow::Cow, + collections::BTreeMap, + io::{BufWriter, Write}, +}; +use serde::{Serialize, Serializer, ser::SerializeMap}; + +use flate2::{Compression, write::GzEncoder}; use litellm_http::Client; use serde_json::Value; use sha2::{Digest, Sha256}; use time::{OffsetDateTime, format_description::well_known::Rfc3339}; -use crate::{Connection, Error}; +use crate::{Connection, Error, Shared}; const MAX_INSERT_BYTES: usize = 64 * 1024 * 1024; +pub type InsertRow = BTreeMap>; + pub enum InsertTable { OtelTraces, SpendLogs, @@ -37,85 +46,176 @@ pub async fn insert_rows( database: &str, table: InsertTable, rows: Vec>, +) -> Result<(), Error> { + insert_shared_rows(client, connection, database, table, shared_rows(rows)).await +} + +pub async fn insert_shared_rows( + client: &Client, + connection: &Connection, + database: &str, + table: InsertTable, + rows: Vec, ) -> Result<(), Error> { if rows.is_empty() { return Ok(()); } - let token = format!( - "{:x}", - Sha256::digest(encode_rows_with_limit(rows.clone(), MAX_INSERT_BYTES)?.as_bytes()) - ); - let received_ms = OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000; - let rows = rows - .into_iter() - .map(|row| { - row.into_iter() - .filter(|(key, _)| key != "EngineReceivedMs") - .chain(std::iter::once(( - "EngineReceivedMs".to_owned(), - Value::from(received_ms as u64), - ))) - .collect() - }) - .collect(); - let encoded = encode_rows_with_limit(rows, MAX_INSERT_BYTES)?; - litellm_storage_clickhouse::insert_encoded_rows( + let received_ms = (OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) as u64; + let (token, body) = prepare_insert(&rows, received_ms, MAX_INSERT_BYTES)?; + litellm_storage_clickhouse::insert_compressed_rows( client, connection, database, table.name(), &token, - &encoded, + body, ) .await } -pub fn encode_rows(rows: Vec>) -> Result { - encode_rows_with_limit(rows, usize::MAX) +fn shared_rows(rows: Vec>) -> Vec { + rows.into_iter() + .map(|row| { + row.into_iter() + .map(|(key, value)| (key, Shared::new(value))) + .collect() + }) + .collect() } -fn encode_rows_with_limit( - rows: Vec>, - limit: usize, -) -> Result { - let mut body = Vec::new(); - for row in rows { - let encoded = row - .into_iter() - .map(|(name, value)| insert_value(&name, value).map(|value| (name, value))) - .collect::, _>>()?; - let record = serde_json::to_vec(&encoded).map_err(|_| Error::InvalidRow)?; - let size = body - .len() - .checked_add(record.len()) - .and_then(|size| size.checked_add(usize::from(!body.is_empty()))) - .ok_or(Error::InsertTooLarge)?; - if size > limit { - return Err(Error::InsertTooLarge); - } - if !body.is_empty() { - body.push(b'\n'); - } - body.extend_from_slice(&record); - } +pub fn encode_rows(rows: Vec>) -> Result { + let body = write_rows(&shared_rows(rows), None, Vec::new(), usize::MAX)?; String::from_utf8(body).map_err(|_| Error::InvalidRow) } -fn insert_value(name: &str, value: Value) -> Result { +fn prepare_insert( + rows: &[InsertRow], + received_ms: u64, + limit: usize, +) -> Result<(String, Vec), Error> { + let hash = write_rows(rows, None, HashWriter(Sha256::new()), limit)?; + let token = format!("{:x}", hash.0.finalize()); + let encoder = write_rows( + rows, + Some(received_ms), + BufWriter::new(GzEncoder::new(Vec::new(), Compression::default())), + limit, + )?; + let body = encoder + .into_inner() + .map_err(|_| Error::InvalidRow)? + .finish() + .map_err(|_| Error::InvalidRow)?; + Ok((token, body)) +} + +struct HashWriter(Sha256); + +impl Write for HashWriter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.0.update(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +struct LimitedWriter { + inner: W, + remaining: usize, + exceeded: bool, +} + +impl Write for LimitedWriter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + if bytes.len() > self.remaining { + self.exceeded = true; + return Err(std::io::Error::other(Error::InsertTooLarge)); + } + let written = self.inner.write(bytes)?; + self.remaining -= written; + Ok(written) + } + + fn flush(&mut self) -> std::io::Result<()> { + self.inner.flush() + } +} + +fn write_rows( + rows: &[InsertRow], + received_ms: Option, + writer: W, + limit: usize, +) -> Result { + let mut writer = LimitedWriter { + inner: writer, + remaining: limit, + exceeded: false, + }; + for (index, row) in rows.iter().enumerate() { + let result = (|| { + if index != 0 { + writer.write_all(b"\n").map_err(serde_json::Error::io)?; + } + serde_json::to_writer(&mut writer, &EncodedRow { row, received_ms }) + })(); + if result.is_err() { + return Err(if writer.exceeded { + Error::InsertTooLarge + } else { + Error::InvalidRow + }); + } + } + Ok(writer.inner) +} + +struct EncodedRow<'a> { + row: &'a InsertRow, + received_ms: Option, +} + +impl Serialize for EncodedRow<'_> { + fn serialize(&self, serializer: S) -> Result { + let mut map = serializer.serialize_map(None)?; + let mut received_ms = self.received_ms; + for (name, value) in self.row { + if name.as_str() >= "EngineReceivedMs" + && let Some(timestamp) = received_ms.take() + { + map.serialize_entry("EngineReceivedMs", ×tamp)?; + } + if name == "EngineReceivedMs" && self.received_ms.is_some() { + continue; + } + let value = insert_value(name, value).map_err(serde::ser::Error::custom)?; + map.serialize_entry(name, &value)?; + } + if let Some(timestamp) = received_ms { + map.serialize_entry("EngineReceivedMs", ×tamp)?; + } + map.end() + } +} + +fn insert_value<'a>(name: &str, value: &'a Value) -> Result, Error> { let multiplier = match name { "Timestamp" => 1, "start_time" | "end_time" | "completion_start_time" => 1_000_000, - _ => return Ok(value), + _ => return Ok(Cow::Borrowed(value)), }; if name == "completion_start_time" && value.is_null() { - return Ok(value); + return Ok(Cow::Borrowed(value)); } let timestamp = value.as_i64().ok_or(Error::InvalidRow)?; let datetime = OffsetDateTime::from_unix_timestamp_nanos(i128::from(timestamp) * multiplier) .map_err(|_| Error::InvalidRow)?; datetime .format(&Rfc3339) - .map(Value::String) + .map(|value| Cow::Owned(Value::String(value))) .map_err(|_| Error::InvalidRow) } @@ -126,20 +226,83 @@ mod tests { use rstest::rstest; use serde_json::json; - use super::encode_rows_with_limit; + use super::{shared_rows, write_rows}; use crate::Error; #[rstest] fn encoded_limit_counts_utf8_bytes_across_rows() { - let rows = vec![ + let rows = shared_rows(vec![ BTreeMap::from([("Input".to_owned(), json!("雪"))]), BTreeMap::from([("Input".to_owned(), json!("雪"))]), - ]; - let encoded = encode_rows_with_limit(rows.clone(), usize::MAX).expect("valid rows"); + ]); + let encoded = write_rows(&rows, None, Vec::new(), usize::MAX).expect("valid rows"); - assert!(encode_rows_with_limit(rows.clone(), encoded.len()).is_ok()); + assert!(write_rows(&rows, None, Vec::new(), encoded.len()).is_ok()); assert!(matches!( - encode_rows_with_limit(rows, encoded.len() - 1), + write_rows(&rows, None, Vec::new(), encoded.len() - 1), + Err(Error::InsertTooLarge) + )); + } + + #[rstest] + #[case::absent(None)] + #[case::submitted(Some(123))] + fn streamed_insert_preserves_token_and_stamps_receive_time(#[case] submitted: Option) { + use flate2::read::GzDecoder; + use sha2::{Digest, Sha256}; + use std::io::Read; + let mut row = BTreeMap::from([ + ("ApiKeyHash".into(), json!("key")), + ("ResourceAttributes".into(), json!({"message": "雪\n\""})), + ("Timestamp".into(), json!(1_234_567_890)), + ]); + if let Some(value) = submitted { + row.insert("EngineReceivedMs".into(), json!(value)); + } + let legacy = match submitted { + Some(_) => { + "{\"ApiKeyHash\":\"key\",\"EngineReceivedMs\":123,\"ResourceAttributes\":{\"message\":\"雪\\n\\\"\"},\"Timestamp\":\"1970-01-01T00:00:01.23456789Z\"}" + } + None => { + "{\"ApiKeyHash\":\"key\",\"ResourceAttributes\":{\"message\":\"雪\\n\\\"\"},\"Timestamp\":\"1970-01-01T00:00:01.23456789Z\"}" + } + }; + let rows = shared_rows(vec![row.clone(), row]); + let (token, body) = super::prepare_insert(&rows, 456, 4096).unwrap(); + assert_eq!( + token, + format!("{:x}", Sha256::digest(format!("{legacy}\n{legacy}"))) + ); + let mut decoded = String::new(); + GzDecoder::new(body.as_slice()) + .read_to_string(&mut decoded) + .unwrap(); + let expected = json!({ + "ApiKeyHash": "key", "EngineReceivedMs": 456, + "ResourceAttributes": {"message": "雪\n\""}, + "Timestamp": "1970-01-01T00:00:01.23456789Z", + }); + assert_eq!( + decoded + .lines() + .map(|line| serde_json::from_str::(line).unwrap()) + .collect::>(), + vec![expected.clone(), expected] + ); + assert_eq!( + rows[0] + .get("EngineReceivedMs") + .map(|value| value.as_u64().unwrap()), + submitted + ); + } + + #[rstest] + fn stamped_insert_enforces_the_encoded_limit() { + let rows = shared_rows(vec![BTreeMap::new()]); + assert!(super::prepare_insert(&rows, 1, 22).is_ok()); + assert!(matches!( + super::prepare_insert(&rows, 1, 21), Err(Error::InsertTooLarge) )); } diff --git a/litellm-rust/crates/traces/src/lib.rs b/litellm-rust/crates/traces/src/lib.rs index f5defb36cc2..1489b44c118 100644 --- a/litellm-rust/crates/traces/src/lib.rs +++ b/litellm-rust/crates/traces/src/lib.rs @@ -2,11 +2,13 @@ mod error; mod insert; mod otlp; mod schema; +mod shared; mod sql; pub use error::DecodeError; -pub use insert::{InsertTable, encode_rows, insert_rows}; +pub use insert::{InsertRow, InsertTable, encode_rows, insert_rows, insert_shared_rows}; pub use litellm_storage_clickhouse::{Connection, Error, Parameter, execute_read}; pub use otlp::{DecodedSpan, decode_otlp}; pub use schema::{ensure_schema, schema_statements}; +pub use shared::{Shared, SharedIdentity}; pub use sql::{LensQuery, ReadQuery, execute_named_read}; diff --git a/litellm-rust/crates/traces/src/otlp.rs b/litellm-rust/crates/traces/src/otlp.rs deleted file mode 100644 index f162256ef1f..00000000000 --- a/litellm-rust/crates/traces/src/otlp.rs +++ /dev/null @@ -1,221 +0,0 @@ -use std::{collections::BTreeMap, io::Read}; - -use base64::Engine; -use flate2::read::GzDecoder; -use opentelemetry_proto::tonic::{ - collector::trace::v1::ExportTraceServiceRequest, - common::v1::{AnyValue, KeyValue, any_value::Value as AttributeValue}, - trace::v1::{Span, span::SpanKind, status::StatusCode}, -}; -use prost::Message; -use serde::Serialize; -use serde_json::Value; - -use crate::DecodeError; - -#[derive(Serialize)] -pub struct DecodedEvent { - pub name: String, - pub attributes: BTreeMap, -} - -#[derive(Serialize)] -pub struct DecodedSpan { - pub trace_id: String, - pub span_id: String, - pub parent_span_id: String, - pub trace_state: String, - pub name: String, - pub kind: String, - pub resource_attributes: BTreeMap, - pub scope_name: String, - pub scope_version: String, - pub attributes: BTreeMap, - pub start_ns: u64, - pub end_ns: u64, - pub status_code: String, - pub status_message: String, - pub events: Vec, -} - -pub fn decode_otlp( - body: &[u8], - content_type: Option<&str>, - content_encoding: Option<&str>, - max_decompressed_bytes: usize, -) -> Result, DecodeError> { - let payload = if content_encoding == Some("gzip") || body.starts_with(&[0x1f, 0x8b]) { - let limit = u64::try_from(max_decompressed_bytes).map_err(|_| DecodeError::TooLarge)?; - let mut decoded = Vec::new(); - GzDecoder::new(body) - .take(limit + 1) - .read_to_end(&mut decoded) - .map_err(|_| DecodeError::InvalidPayload)?; - decoded - } else { - body.to_vec() - }; - if payload.len() > max_decompressed_bytes { - return Err(DecodeError::TooLarge); - } - let request = if content_type.is_some_and(|value| value.contains("json")) { - let value: Value = - serde_json::from_slice(&payload).map_err(|_| DecodeError::InvalidPayload)?; - serde_json::from_value(normalize_json_ids(value)?) - .map_err(|_| DecodeError::InvalidPayload)? - } else { - ExportTraceServiceRequest::decode(payload.as_slice()) - .map_err(|_| DecodeError::InvalidPayload)? - }; - Ok(request - .resource_spans - .into_iter() - .flat_map(|resource_spans| { - let resource_attributes = attributes( - resource_spans - .resource - .map(|resource| resource.attributes) - .unwrap_or_default(), - ); - resource_spans - .scope_spans - .into_iter() - .flat_map(move |scope_spans| { - let scope = scope_spans.scope.unwrap_or_default(); - let resource_attributes = resource_attributes.clone(); - scope_spans.spans.into_iter().map(move |span| { - decoded_span(span, &resource_attributes, &scope.name, &scope.version) - }) - }) - }) - .collect()) -} - -fn normalize_json_ids(value: Value) -> Result { - match value { - Value::Object(fields) => fields - .into_iter() - .map(|(name, value)| { - let normalized = if matches!(name.as_str(), "traceId" | "spanId" | "parentSpanId") { - let encoded = value.as_str().ok_or(DecodeError::InvalidPayload)?; - let bytes = base64::engine::general_purpose::STANDARD - .decode(encoded) - .map_err(|_| DecodeError::InvalidPayload)?; - Value::String(hex_bytes(&bytes)) - } else if name == "kind" && value.is_string() { - let kind = SpanKind::from_str_name(value.as_str().unwrap_or_default()) - .ok_or(DecodeError::InvalidPayload)?; - Value::from(kind as i32) - } else if name == "code" && value.is_string() { - let code = StatusCode::from_str_name(value.as_str().unwrap_or_default()) - .ok_or(DecodeError::InvalidPayload)?; - Value::from(code as i32) - } else { - normalize_json_ids(value)? - }; - Ok((name, normalized)) - }) - .collect::, _>>() - .map(Value::Object), - Value::Array(values) => values - .into_iter() - .map(normalize_json_ids) - .collect::, _>>() - .map(Value::Array), - value => Ok(value), - } -} - -fn hex_bytes(bytes: &[u8]) -> String { - bytes.iter().map(|byte| format!("{byte:02x}")).collect() -} - -fn decoded_span( - span: Span, - resource_attributes: &BTreeMap, - scope_name: &str, - scope_version: &str, -) -> DecodedSpan { - let status = span.status.unwrap_or_default(); - DecodedSpan { - trace_id: hex_bytes(&span.trace_id), - span_id: hex_bytes(&span.span_id), - parent_span_id: hex_bytes(&span.parent_span_id), - trace_state: span.trace_state, - name: span.name, - kind: SpanKind::try_from(span.kind) - .unwrap_or(SpanKind::Unspecified) - .as_str_name() - .to_owned(), - resource_attributes: resource_attributes.clone(), - scope_name: scope_name.to_owned(), - scope_version: scope_version.to_owned(), - attributes: attributes(span.attributes), - start_ns: span.start_time_unix_nano, - end_ns: span.end_time_unix_nano, - status_code: StatusCode::try_from(status.code) - .unwrap_or(StatusCode::Unset) - .as_str_name() - .to_owned(), - status_message: status.message, - events: span - .events - .into_iter() - .map(|event| DecodedEvent { - name: event.name, - attributes: attributes(event.attributes), - }) - .collect(), - } -} - -fn attributes(values: Vec) -> BTreeMap { - values - .into_iter() - .map(|entry| { - ( - entry.key, - entry.value.as_ref().map(attribute_text).unwrap_or_default(), - ) - }) - .collect() -} - -fn attribute_text(value: &AnyValue) -> String { - match value.value.as_ref() { - Some(AttributeValue::StringValue(value)) => value.clone(), - Some(AttributeValue::BoolValue(value)) => value.to_string(), - Some(AttributeValue::IntValue(value)) => value.to_string(), - Some(AttributeValue::DoubleValue(value)) => { - serde_json::to_string(value).unwrap_or_default() - } - Some(AttributeValue::BytesValue(value)) => String::from_utf8_lossy(value).into_owned(), - Some(AttributeValue::ArrayValue(value)) => format!( - "[{}]", - value - .values - .iter() - .map(|value| serde_json::to_string(&attribute_text(value)).unwrap_or_default()) - .collect::>() - .join(", ") - ), - Some(AttributeValue::KvlistValue(value)) => format!( - "{{{}}}", - value - .values - .iter() - .map(|entry| format!( - "{}: {}", - serde_json::to_string(&entry.key).unwrap_or_default(), - serde_json::to_string( - &entry.value.as_ref().map(attribute_text).unwrap_or_default() - ) - .unwrap_or_default() - )) - .collect::>() - .join(", ") - ), - Some(AttributeValue::StringValueStrindex(value)) => value.to_string(), - None => String::new(), - } -} diff --git a/litellm-rust/crates/traces/src/otlp/attributes.rs b/litellm-rust/crates/traces/src/otlp/attributes.rs new file mode 100644 index 00000000000..50e063e4582 --- /dev/null +++ b/litellm-rust/crates/traces/src/otlp/attributes.rs @@ -0,0 +1,101 @@ +use std::{collections::BTreeMap, io::Write}; + +use opentelemetry_proto::tonic::common::v1::{ + AnyValue, KeyValue, any_value::Value as AttributeValue, +}; +use serde::{ + Serialize, Serializer, + ser::{SerializeMap, SerializeSeq}, +}; + +use super::limits::{Budget, MAX_ATTRIBUTES}; +use crate::DecodeError; + +struct AttributeWriter<'a> { + body: Vec, + budget: &'a mut Budget, +} + +impl Write for AttributeWriter<'_> { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.budget + .consume(bytes.len()) + .map_err(std::io::Error::other)?; + self.body.extend_from_slice(bytes); + Ok(bytes.len()) + } + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +pub(super) fn attributes( + values: Vec, + budget: &mut Budget, +) -> Result, DecodeError> { + if values.len() > MAX_ATTRIBUTES { + return Err(DecodeError::TooLarge); + } + values + .into_iter() + .map(|entry| { + budget.consume(entry.key.len() + 96)?; + let text = match entry.value { + Some(AnyValue { + value: Some(AttributeValue::StringValue(value)), + }) => { + budget.consume(value.len())?; + value + } + Some(AnyValue { + value: Some(AttributeValue::BytesValue(value)), + }) => { + budget.consume(value.len().saturating_mul(3))?; + String::from_utf8_lossy(&value).into_owned() + } + value => { + let mut writer = AttributeWriter { + body: Vec::new(), + budget, + }; + serde_json::to_writer(&mut writer, &AttributeJson(value.as_ref())) + .map_err(|_| DecodeError::TooLarge)?; + String::from_utf8(writer.body).map_err(|_| DecodeError::InvalidPayload)? + } + }; + Ok((entry.key, text)) + }) + .collect() +} + +struct AttributeJson<'a>(Option<&'a AnyValue>); + +impl Serialize for AttributeJson<'_> { + fn serialize(&self, serializer: S) -> Result { + match self.0.and_then(|value| value.value.as_ref()) { + Some(AttributeValue::StringValue(value)) => serializer.serialize_str(value), + Some(AttributeValue::BoolValue(value)) => serializer.serialize_bool(*value), + Some(AttributeValue::IntValue(value)) => serializer.serialize_i64(*value), + Some(AttributeValue::DoubleValue(value)) => serializer.serialize_f64(*value), + Some(AttributeValue::BytesValue(value)) => { + serializer.serialize_str(&String::from_utf8_lossy(value)) + } + Some(AttributeValue::ArrayValue(value)) => { + let mut sequence = serializer.serialize_seq(Some(value.values.len()))?; + for entry in &value.values { + sequence.serialize_element(&AttributeJson(Some(entry)))?; + } + sequence.end() + } + Some(AttributeValue::KvlistValue(value)) => { + let mut map = serializer.serialize_map(Some(value.values.len()))?; + for entry in &value.values { + map.serialize_entry(&entry.key, &AttributeJson(entry.value.as_ref()))?; + } + map.end() + } + Some(AttributeValue::StringValueStrindex(value)) => serializer.serialize_i32(*value), + None => serializer.serialize_unit(), + } + } +} diff --git a/litellm-rust/crates/traces/src/otlp/limits.rs b/litellm-rust/crates/traces/src/otlp/limits.rs new file mode 100644 index 00000000000..f6b56ccf12d --- /dev/null +++ b/litellm-rust/crates/traces/src/otlp/limits.rs @@ -0,0 +1,212 @@ +use std::fmt; + +use prost::encoding::{DecodeContext, WireType, decode_key, decode_varint, skip_field}; +use serde::de::{DeserializeSeed, MapAccess, SeqAccess, Visitor}; + +use crate::{DecodeError, Shared}; + +pub(super) const MAX_DEPTH: usize = 32; +pub(super) const MAX_NODES: usize = 65_536; +pub(super) const MAX_SPANS: usize = 4_096; +pub(super) const MAX_ATTRIBUTES: usize = 256; +pub(super) const MAX_EVENTS: usize = 256; +pub(super) const MAX_DECODED_SPAN_BYTES: usize = 16 * 1024 * 1024; + +pub(super) fn json_preflight(payload: &[u8]) -> Result<(), DecodeError> { + let mut nodes = 0; + let mut exceeded = false; + let mut decoder = serde_json::Deserializer::from_slice(payload); + let result = JsonBudget { + nodes: &mut nodes, + exceeded: &mut exceeded, + depth: 0, + } + .deserialize(&mut decoder) + .and_then(|()| decoder.end()); + if exceeded { + return Err(DecodeError::TooLarge); + } + result.map_err(|_| DecodeError::InvalidPayload) +} + +struct JsonBudget<'a> { + nodes: &'a mut usize, + exceeded: &'a mut bool, + depth: usize, +} + +impl<'de> DeserializeSeed<'de> for JsonBudget<'_> { + type Value = (); + + fn deserialize>(self, decoder: D) -> Result<(), D::Error> { + *self.nodes += 1; + if *self.nodes > MAX_NODES || self.depth > MAX_DEPTH { + *self.exceeded = true; + return Err(serde::de::Error::custom("OTLP structure exceeds budget")); + } + decoder.deserialize_any(self) + } +} + +impl<'de> Visitor<'de> for JsonBudget<'_> { + type Value = (); + + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("OTLP JSON") + } + fn visit_bool(self, _: bool) -> Result<(), E> { + Ok(()) + } + fn visit_i64(self, _: i64) -> Result<(), E> { + Ok(()) + } + fn visit_u64(self, _: u64) -> Result<(), E> { + Ok(()) + } + fn visit_f64(self, _: f64) -> Result<(), E> { + Ok(()) + } + fn visit_str(self, _: &str) -> Result<(), E> { + Ok(()) + } + fn visit_unit(self) -> Result<(), E> { + Ok(()) + } + + fn visit_seq>(self, mut sequence: A) -> Result<(), A::Error> { + while sequence + .next_element_seed(JsonBudget { + nodes: self.nodes, + exceeded: self.exceeded, + depth: self.depth + 1, + })? + .is_some() + {} + Ok(()) + } + + fn visit_map>(self, mut map: A) -> Result<(), A::Error> { + while map + .next_key_seed(JsonBudget { + nodes: self.nodes, + exceeded: self.exceeded, + depth: self.depth + 1, + })? + .is_some() + { + map.next_value_seed(JsonBudget { + nodes: self.nodes, + exceeded: self.exceeded, + depth: self.depth + 1, + })?; + } + Ok(()) + } +} + +#[derive(Clone, Copy)] +enum MessageKind { + Export, + ResourceSpans, + Resource, + ScopeSpans, + Scope, + Span, + Event, + Link, + Status, + KeyValue, + AnyValue, + Array, + KvList, +} + +impl MessageKind { + fn child(self, tag: u32) -> Option { + match (self, tag) { + (Self::Export, 1) => Some(Self::ResourceSpans), + (Self::ResourceSpans, 1) => Some(Self::Resource), + (Self::ResourceSpans, 2) => Some(Self::ScopeSpans), + (Self::Resource, 1) + | (Self::Scope, 3) + | (Self::Span, 9) + | (Self::Event, 3) + | (Self::Link, 4) + | (Self::KvList, 1) => Some(Self::KeyValue), + (Self::ScopeSpans, 1) => Some(Self::Scope), + (Self::ScopeSpans, 2) => Some(Self::Span), + (Self::Span, 11) => Some(Self::Event), + (Self::Span, 13) => Some(Self::Link), + (Self::Span, 15) => Some(Self::Status), + (Self::KeyValue, 2) | (Self::Array, 1) => Some(Self::AnyValue), + (Self::AnyValue, 5) => Some(Self::Array), + (Self::AnyValue, 6) => Some(Self::KvList), + _ => None, + } + } +} + +pub(super) fn protobuf_preflight(payload: &[u8]) -> Result<(), DecodeError> { + scan_message(payload, MessageKind::Export, 0, &mut 0) +} + +fn scan_message( + mut payload: &[u8], + kind: MessageKind, + depth: usize, + nodes: &mut usize, +) -> Result<(), DecodeError> { + if depth > MAX_DEPTH { + return Err(DecodeError::TooLarge); + } + while !payload.is_empty() { + *nodes += 1; + if *nodes > MAX_NODES { + return Err(DecodeError::TooLarge); + } + let (tag, wire) = decode_key(&mut payload).map_err(|_| DecodeError::InvalidPayload)?; + if let (WireType::LengthDelimited, Some(child)) = (wire, kind.child(tag)) { + let length = decode_varint(&mut payload).map_err(|_| DecodeError::InvalidPayload)?; + let length = usize::try_from(length).map_err(|_| DecodeError::InvalidPayload)?; + let (message, rest) = payload + .split_at_checked(length) + .ok_or(DecodeError::InvalidPayload)?; + scan_message(message, child, depth + 1, nodes)?; + payload = rest; + } else { + skip_field(wire, tag, &mut payload, DecodeContext::default()) + .map_err(|_| DecodeError::InvalidPayload)?; + } + } + Ok(()) +} + +pub(super) struct Budget { + remaining: usize, +} + +impl Budget { + pub(super) fn new(remaining: usize) -> Self { + Self { remaining } + } + + pub(super) fn clone_shared( + &mut self, + value: &Shared, + allocated_bytes: impl FnOnce(&T) -> usize, + ) -> Result, DecodeError> { + let cloned = value.clone(); + if !value.shares_storage_with(&cloned) { + self.consume(allocated_bytes(value))?; + } + Ok(cloned) + } + + pub(super) fn consume(&mut self, bytes: usize) -> Result<(), DecodeError> { + self.remaining = self + .remaining + .checked_sub(bytes) + .ok_or(DecodeError::TooLarge)?; + Ok(()) + } +} diff --git a/litellm-rust/crates/traces/src/otlp/mod.rs b/litellm-rust/crates/traces/src/otlp/mod.rs new file mode 100644 index 00000000000..fcc42082151 --- /dev/null +++ b/litellm-rust/crates/traces/src/otlp/mod.rs @@ -0,0 +1,42 @@ +mod attributes; +mod limits; +mod span; +mod wire; + +use serde::Serialize; +use std::collections::BTreeMap; + +use crate::{DecodeError, Shared}; + +#[derive(Serialize)] +pub struct DecodedEvent { + pub name: String, + pub attributes: BTreeMap, +} + +#[derive(Serialize)] +pub struct DecodedSpan { + pub trace_id: String, + pub span_id: String, + pub parent_span_id: String, + pub trace_state: String, + pub name: String, + pub kind: String, + pub resource_attributes: Shared>, + pub scope_name: Shared, + pub scope_version: Shared, + pub attributes: BTreeMap, + pub start_ns: u64, + pub end_ns: u64, + pub status_code: String, + pub status_message: String, + pub events: Vec, +} + +pub fn decode_otlp( + body: &[u8], + content_type: Option<&str>, +) -> Result, DecodeError> { + let request = wire::decode(body, content_type)?; + span::flatten(request) +} diff --git a/litellm-rust/crates/traces/src/otlp/span.rs b/litellm-rust/crates/traces/src/otlp/span.rs new file mode 100644 index 00000000000..fa993f71e3c --- /dev/null +++ b/litellm-rust/crates/traces/src/otlp/span.rs @@ -0,0 +1,166 @@ +use std::collections::BTreeMap; + +use opentelemetry_proto::tonic::{ + collector::trace::v1::ExportTraceServiceRequest, + trace::v1::{ResourceSpans, ScopeSpans, Span, span::SpanKind, status::StatusCode}, +}; + +use super::{ + DecodedEvent, DecodedSpan, + attributes::attributes, + limits::{Budget, MAX_ATTRIBUTES, MAX_DECODED_SPAN_BYTES, MAX_EVENTS, MAX_SPANS}, +}; +use crate::{DecodeError, Shared}; + +pub(super) fn flatten(request: ExportTraceServiceRequest) -> Result, DecodeError> { + let mut budget = Budget::new(MAX_DECODED_SPAN_BYTES); + let mut spans = Vec::new(); + for resource in request.resource_spans { + append_resource(resource, &mut budget, &mut spans)?; + } + Ok(spans) +} + +fn append_resource( + resource: ResourceSpans, + budget: &mut Budget, + spans: &mut Vec, +) -> Result<(), DecodeError> { + let attributes = Shared::new(attributes( + resource + .resource + .map(|resource| resource.attributes) + .unwrap_or_default(), + budget, + )?); + for scope in resource.scope_spans { + append_scope(scope, &attributes, budget, spans)?; + } + Ok(()) +} + +fn append_scope( + scope_spans: ScopeSpans, + resource: &Shared>, + budget: &mut Budget, + spans: &mut Vec, +) -> Result<(), DecodeError> { + let scope = scope_spans.scope.unwrap_or_default(); + if scope.attributes.len() > MAX_ATTRIBUTES { + return Err(DecodeError::TooLarge); + } + budget.consume(scope.name.len() + scope.version.len())?; + let scope_name: Shared = scope.name.into(); + let scope_version: Shared = scope.version.into(); + for span in scope_spans.spans { + if spans.len() >= MAX_SPANS { + return Err(DecodeError::TooLarge); + } + validate_span(&span)?; + budget.consume( + span.name.len() + + span.trace_state.len() + + span + .status + .as_ref() + .map_or(0, |status| status.message.len()) + + size_of::() + + 128, + )?; + spans.push(decoded_span( + span, + resource, + &scope_name, + &scope_version, + budget, + )?); + } + Ok(()) +} + +fn valid_id(value: &[u8], length: usize) -> bool { + value.len() == length && value.iter().any(|byte| *byte != 0) +} + +fn validate_span(span: &Span) -> Result<(), DecodeError> { + if !valid_id(&span.trace_id, 16) + || !valid_id(&span.span_id, 8) + || (!span.parent_span_id.is_empty() && !valid_id(&span.parent_span_id, 8)) + || span.start_time_unix_nano > i64::MAX as u64 + || span.end_time_unix_nano > i64::MAX as u64 + || span.end_time_unix_nano < span.start_time_unix_nano + || span + .links + .iter() + .any(|link| !valid_id(&link.trace_id, 16) || !valid_id(&link.span_id, 8)) + { + return Err(DecodeError::InvalidPayload); + } + if span.events.len() > MAX_EVENTS + || span.links.len() > MAX_EVENTS + || span.attributes.len() > MAX_ATTRIBUTES + || span + .links + .iter() + .any(|link| link.attributes.len() > MAX_ATTRIBUTES) + || span + .events + .iter() + .any(|event| event.attributes.len() > MAX_ATTRIBUTES) + { + return Err(DecodeError::TooLarge); + } + Ok(()) +} + +fn hex_bytes(bytes: &[u8]) -> String { + bytes.iter().map(|byte| format!("{byte:02x}")).collect() +} + +fn decoded_span( + span: Span, + resource_attributes: &Shared>, + scope_name: &Shared, + scope_version: &Shared, + budget: &mut Budget, +) -> Result { + let status = span.status.unwrap_or_default(); + Ok(DecodedSpan { + trace_id: hex_bytes(&span.trace_id), + span_id: hex_bytes(&span.span_id), + parent_span_id: hex_bytes(&span.parent_span_id), + trace_state: span.trace_state, + name: span.name, + kind: SpanKind::try_from(span.kind) + .unwrap_or(SpanKind::Unspecified) + .as_str_name() + .to_owned(), + resource_attributes: budget.clone_shared(resource_attributes, |attributes| { + attributes + .iter() + .map(|(key, value)| key.len() + value.len() + 96) + .sum() + })?, + scope_name: budget.clone_shared(scope_name, String::len)?, + scope_version: budget.clone_shared(scope_version, String::len)?, + attributes: attributes(span.attributes, budget)?, + start_ns: span.start_time_unix_nano, + end_ns: span.end_time_unix_nano, + status_code: StatusCode::try_from(status.code) + .unwrap_or(StatusCode::Unset) + .as_str_name() + .to_owned(), + status_message: status.message, + events: span + .events + .into_iter() + .map(|event| { + budget.consume(event.name.len() + 96)?; + Ok(DecodedEvent { + name: event.name, + attributes: attributes(event.attributes, budget)?, + }) + }) + .collect::, DecodeError>>()?, + }) +} diff --git a/litellm-rust/crates/traces/src/otlp/wire.rs b/litellm-rust/crates/traces/src/otlp/wire.rs new file mode 100644 index 00000000000..bac29ba49e4 --- /dev/null +++ b/litellm-rust/crates/traces/src/otlp/wire.rs @@ -0,0 +1,43 @@ +use opentelemetry_proto::tonic::collector::trace::v1::ExportTraceServiceRequest; +use prost::Message; + +use super::limits::{json_preflight, protobuf_preflight}; +use crate::DecodeError; + +#[derive(strum::EnumString)] +#[strum(ascii_case_insensitive)] +enum OtlpMediaType { + #[strum(serialize = "application/json")] + Json, + #[strum( + serialize = "application/x-protobuf", + serialize = "application/protobuf" + )] + Protobuf, +} + +pub(super) fn decode( + body: &[u8], + content_type: Option<&str>, +) -> Result { + let media_type = content_type + .unwrap_or("application/x-protobuf") + .split(';') + .next() + .unwrap_or_default() + .trim() + .parse::() + .map_err(|_| DecodeError::InvalidPayload)?; + + let request = match media_type { + OtlpMediaType::Json => { + json_preflight(body)?; + serde_json::from_slice(body).map_err(|_| DecodeError::InvalidPayload)? + } + OtlpMediaType::Protobuf => { + protobuf_preflight(body)?; + ExportTraceServiceRequest::decode(body).map_err(|_| DecodeError::InvalidPayload)? + } + }; + Ok(request) +} diff --git a/litellm-rust/crates/traces/src/shared.rs b/litellm-rust/crates/traces/src/shared.rs new file mode 100644 index 00000000000..dafd08b72dc --- /dev/null +++ b/litellm-rust/crates/traces/src/shared.rs @@ -0,0 +1,46 @@ +use std::ops::Deref; + +use serde::Serialize; + +type Storage = std::sync::Arc; + +#[derive(Clone, Debug, PartialEq, Serialize)] +#[serde(transparent)] +pub struct Shared(Storage); + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct SharedIdentity(usize); + +impl Shared { + pub fn new(value: T) -> Self { + Self(Storage::new(value)) + } + + pub fn identity(&self) -> SharedIdentity { + SharedIdentity(std::ptr::from_ref(self.as_ref()) as usize) + } + + pub fn shares_storage_with(&self, other: &Self) -> bool { + self.identity() == other.identity() + } +} + +impl From for Shared { + fn from(value: T) -> Self { + Self::new(value) + } +} + +impl AsRef for Shared { + fn as_ref(&self) -> &T { + self.0.as_ref() + } +} + +impl Deref for Shared { + type Target = T; + + fn deref(&self) -> &T { + self.as_ref() + } +} diff --git a/litellm-rust/crates/traces/src/sql.rs b/litellm-rust/crates/traces/src/sql.rs index 9acb8de0a7a..36d6e3b4521 100644 --- a/litellm-rust/crates/traces/src/sql.rs +++ b/litellm-rust/crates/traces/src/sql.rs @@ -8,6 +8,7 @@ pub enum ReadQuery { ListTraces, TraceSpans, SpanDetail, + SpanError, SpendByResponseIds, } @@ -17,6 +18,7 @@ impl ReadQuery { "list_traces" => Ok(Self::ListTraces), "trace_spans" => Ok(Self::TraceSpans), "span_detail" => Ok(Self::SpanDetail), + "span_error" => Ok(Self::SpanError), "spend_by_response_ids" => Ok(Self::SpendByResponseIds), _ => Err(Error::InvalidQuery), } @@ -27,6 +29,7 @@ impl ReadQuery { Self::ListTraces => include_str!("../query/list_traces.sql"), Self::TraceSpans => include_str!("../query/trace_spans.sql"), Self::SpanDetail => include_str!("../query/span_detail.sql"), + Self::SpanError => include_str!("../query/span_error.sql"), Self::SpendByResponseIds => include_str!("../query/spend_by_response_ids.sql"), } } diff --git a/litellm-rust/crates/traces/tests/insert.rs b/litellm-rust/crates/traces/tests/insert.rs index cba678152b9..9dcb9cddf1f 100644 --- a/litellm-rust/crates/traces/tests/insert.rs +++ b/litellm-rust/crates/traces/tests/insert.rs @@ -1,8 +1,107 @@ -use std::collections::BTreeMap; +use std::{ + collections::BTreeMap, + io::{BufRead, BufReader}, +}; -use litellm_traces::encode_rows; -use rstest::rstest; +use flate2::read::GzDecoder; +use litellm_http::Client; +use litellm_traces::{ + Connection, Error, InsertRow, InsertTable, Shared, encode_rows, insert_shared_rows, +}; +use rstest::{fixture, rstest}; use serde_json::{Value, json}; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{header, method}, +}; + +#[fixture] +fn shared_rows(#[default(16 * 1024)] attribute_bytes: usize) -> Vec { + let resource = Shared::new(json!({"shared": "x".repeat(attribute_bytes)})); + (0..1024) + .map(|index| { + BTreeMap::from([ + ("ResourceAttributes".into(), resource.clone()), + ("SpanId".into(), Shared::new(json!(format!("{index:016x}")))), + ("Timestamp".into(), Shared::new(json!(1))), + ]) + }) + .collect() +} + +#[rstest] +#[case::one_request(1)] +#[case::concurrent_requests(2)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn shared_fanout_survives_gzip_insert_over_http( + shared_rows: Vec, + #[case] concurrency: usize, +) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(header("Content-Encoding", "gzip")) + .respond_with(ResponseTemplate::new(200)) + .expect(concurrency as u64) + .mount(&server) + .await; + let client = Client::no_redirect_for_test(); + let connection = Connection::parse(&server.uri()).unwrap(); + let expected_resource = shared_rows[0]["ResourceAttributes"].clone(); + let expected_count = shared_rows.len(); + let mut requests = tokio::task::JoinSet::new(); + for _ in 0..concurrency { + let client = client.clone(); + let connection = connection.clone(); + let rows = shared_rows.clone(); + requests.spawn(async move { + insert_shared_rows( + &client, + &connection, + "traces", + InsertTable::OtelTraces, + rows, + ) + .await + }); + } + while let Some(result) = requests.join_next().await { + result.unwrap().unwrap(); + } + let received = server.received_requests().await.unwrap(); + assert_eq!(received.len(), concurrency); + for request in received { + let decoder = GzDecoder::new(request.body.as_slice()); + let mut count = 0; + for (index, line) in BufReader::new(decoder).lines().enumerate() { + let row: Value = serde_json::from_str(&line.unwrap()).unwrap(); + assert_eq!(&row["ResourceAttributes"], expected_resource.as_ref()); + assert_eq!(row["SpanId"], format!("{index:016x}")); + assert_eq!(row["Timestamp"], "1970-01-01T00:00:00.000000001Z"); + assert!(row["EngineReceivedMs"].as_u64().unwrap() > 0); + count += 1; + } + assert_eq!(count, expected_count); + } +} + +#[rstest] +#[tokio::test] +async fn shared_fanout_over_insert_limit_never_reaches_http( + #[with(64 * 1024)] shared_rows: Vec, +) { + let server = MockServer::start().await; + let connection = Connection::parse(&server.uri()).unwrap(); + let result = insert_shared_rows( + &Client::no_redirect_for_test(), + &connection, + "traces", + InsertTable::OtelTraces, + shared_rows, + ) + .await; + assert!(matches!(result, Err(Error::InsertTooLarge))); + assert!(server.received_requests().await.unwrap().is_empty()); +} #[rstest] #[case::span("Timestamp", json!(1_234_567_890), json!("1970-01-01T00:00:01.23456789Z"))] diff --git a/litellm-rust/crates/traces/tests/migrations.rs b/litellm-rust/crates/traces/tests/migrations.rs index cc8fe51a469..01e8982423f 100644 --- a/litellm-rust/crates/traces/tests/migrations.rs +++ b/litellm-rust/crates/traces/tests/migrations.rs @@ -835,3 +835,148 @@ async fn lens_content_keeps_output_visible_after_long_input( assert_eq!(recovered, original); Ok(()) } + +#[rstest] +#[case::ascii(10, format!("ParentCommand: {}", "x".repeat(460_000)))] +#[case::multibyte(1_000, "\u{1f9ea}".repeat(1_024))] +#[case::escaped(1_000, "\0\n\"\\".repeat(1_024))] +#[tokio::test] +async fn trace_error_previews_preserve_paginated_diagnostics( + #[future(awt)] database: TestResult, + #[case] span_count: usize, + #[case] message: String, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + let rows = (0..span_count) + .map(|index| { + serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp + index as i64, "TraceId": "diagnostic-trace", + "SpanId": format!("span-{index}"), "SpanName": "tool", + "StatusCode": "STATUS_CODE_ERROR", "StatusMessage": message, + })) + }) + .collect::>, _>>()?; + insert_rows(&database, "otel_traces", rows).await?; + let reader = Connection::reader(&database.url, "trace_test")?; + let mut parameters = BTreeMap::from([ + ( + "trace_id".into(), + Parameter::Text("diagnostic-trace".into()), + ), + ("team_ids".into(), Parameter::Strings(vec![])), + ("api_key_hash".into(), Parameter::Text(String::new())), + ("trace_ref".into(), Parameter::Text(String::new())), + ]); + let body = execute_named_read( + &database.client, + &reader, + ReadQuery::TraceSpans, + ¶meters, + ) + .await?; + let response: serde_json::Value = serde_json::from_str(&body)?; + let spans = response["data"].as_array().expect("trace spans"); + assert_eq!(spans.len(), span_count); + let prefix: String = message.chars().take(128).collect(); + assert!(!prefix.is_empty()); + assert!( + spans + .iter() + .all(|span| span["status_message"] == prefix && span["error_truncated"] == 1) + ); + parameters.insert("span_id".into(), Parameter::Text("span-0".into())); + parameters.insert("error_version".into(), Parameter::Text(String::new())); + let mut recovered = String::new(); + loop { + parameters.insert( + "error_offset".into(), + Parameter::Integer(recovered.chars().count() as i64), + ); + let body = execute_named_read(&database.client, &reader, ReadQuery::SpanError, ¶meters) + .await?; + assert!(body.len() < 128 * 1024); + let response: serde_json::Value = serde_json::from_str(&body)?; + let chunk = response["data"][0]["message"] + .as_str() + .expect("diagnostic chunk"); + assert!(!chunk.is_empty()); + recovered.push_str(chunk); + let version = response["data"][0]["version"] + .as_str() + .expect("diagnostic version"); + parameters.insert("error_version".into(), Parameter::Text(version.into())); + if recovered.chars().count() >= message.chars().count() { + break; + } + } + assert_eq!(recovered, message); + parameters.insert( + "api_key_hash".into(), + Parameter::Text("unrelated-key".into()), + ); + let denied = + execute_named_read(&database.client, &reader, ReadQuery::SpanError, ¶meters).await?; + assert_eq!( + serde_json::from_str::(&denied)?["data"], + serde_json::json!([]) + ); + Ok(()) +} + +#[rstest] +#[case::different_start(1, 0)] +#[case::different_receive(0, 1)] +#[case::tied_timestamps(0, 0)] +#[tokio::test] +async fn duplicate_span_preview_matches_diagnostic( + #[future(awt)] database: TestResult, + #[case] start_delta: i64, + #[case] receive_delta: i64, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + let message = "a".repeat(200); + let rows = [ + (start_delta, receive_delta, "z".repeat(200)), + (0, 0, message.clone()), + ] + .into_iter() + .map(|(start_delta, receive_delta, message)| { + serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp + start_delta, "EngineReceivedMs": 100 + receive_delta, + "TraceId": "duplicate-trace", "SpanId": "duplicate-span", "StatusMessage": message, + })) + }) + .collect::>, _>>()?; + insert_rows(&database, "otel_traces", rows).await?; + let reader = Connection::reader(&database.url, "trace_test")?; + let parameters = BTreeMap::from([ + ("trace_id".into(), Parameter::Text("duplicate-trace".into())), + ("span_id".into(), Parameter::Text("duplicate-span".into())), + ("team_ids".into(), Parameter::Strings(vec![])), + ("api_key_hash".into(), Parameter::Text(String::new())), + ("trace_ref".into(), Parameter::Text(String::new())), + ("error_version".into(), Parameter::Text(String::new())), + ("error_offset".into(), Parameter::Integer(0)), + ]); + let preview = execute_named_read( + &database.client, + &reader, + ReadQuery::TraceSpans, + ¶meters, + ) + .await?; + let diagnostic = + execute_named_read(&database.client, &reader, ReadQuery::SpanError, ¶meters).await?; + let preview: serde_json::Value = serde_json::from_str(&preview)?; + let diagnostic: serde_json::Value = serde_json::from_str(&diagnostic)?; + assert_eq!(preview["data"].as_array().unwrap().len(), 1); + assert_eq!(preview["data"][0]["status_message"], message[..128]); + assert_eq!(diagnostic["data"][0]["message"], message); + Ok(()) +} diff --git a/litellm-rust/crates/traces/tests/otlp.rs b/litellm-rust/crates/traces/tests/otlp.rs index 002ba159ef9..8aa2cbedeb3 100644 --- a/litellm-rust/crates/traces/tests/otlp.rs +++ b/litellm-rust/crates/traces/tests/otlp.rs @@ -1,33 +1,19 @@ -use flate2::{Compression, write::GzEncoder}; +use litellm_traces::Shared; use litellm_traces::decode_otlp; use rstest::rstest; -use std::io::Write; const FIXTURE: &[u8] = include_bytes!( "../../../../tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json" ); #[rstest] -#[case::json(FIXTURE, Some("application/json"), None)] -#[case::gzip_json(FIXTURE, Some("application/json"), Some("gzip"))] -fn decodes_neutral_spans( - #[case] body: &[u8], - #[case] content_type: Option<&str>, - #[case] content_encoding: Option<&str>, -) { - let payload = if content_encoding == Some("gzip") { - let mut encoder = GzEncoder::new(Vec::new(), Compression::default()); - encoder.write_all(body).expect("gzip input"); - encoder.finish().expect("gzip payload") - } else { - body.to_vec() - }; - let spans = decode_otlp(&payload, content_type, content_encoding, 8 * 1024 * 1024) - .expect("valid OTLP export"); +#[case::json(FIXTURE, Some("application/json"))] +fn decodes_neutral_spans(#[case] body: &[u8], #[case] content_type: Option<&str>) { + let spans = decode_otlp(body, content_type).expect("valid OTLP export"); assert_eq!(spans.len(), 6); assert_eq!(spans[0].trace_id, "4bad42b84e9de3ba46fc870185f8f023"); assert_eq!(spans[0].resource_attributes["service.name"], "agent-demo"); - assert_eq!(spans[0].scope_name, "langsmith"); + assert_eq!(spans[0].scope_name.as_ref(), "langsmith"); assert!( spans .iter() @@ -36,12 +22,322 @@ fn decodes_neutral_spans( } #[rstest] -#[case::invalid(b"not protobuf", None, 8 * 1024 * 1024)] -#[case::too_large(FIXTURE, Some("application/json"), 1)] -fn rejects_invalid_or_oversized_payload( - #[case] body: &[u8], - #[case] content_type: Option<&str>, - #[case] limit: usize, -) { - assert!(decode_otlp(body, content_type, None, limit).is_err()); +fn accepts_trace_larger_than_eight_mib(mut span: opentelemetry_proto::tonic::trace::v1::Span) { + use prost::Message; + + span.name = "x".repeat(9 * 1024 * 1024); + let body = request_with(span).encode_to_vec(); + let decoded = decode_otlp(&body, None).expect("16 MiB default accepts a 9 MiB trace"); + assert_eq!(decoded[0].name.len(), 9 * 1024 * 1024); +} + +#[rstest] +fn rejects_invalid_payload() { + assert!(decode_otlp(b"not protobuf", None).is_err()); +} + +#[rstest] +fn decoder_does_not_enforce_the_http_body_limit() { + let body = format!("{{\"ignored\":\"{}\"}}", "x".repeat(16 * 1024 * 1024 + 1)); + assert!( + decode_otlp(body.as_bytes(), Some("application/json")) + .unwrap() + .is_empty() + ); +} + +fn request_with( + span: opentelemetry_proto::tonic::trace::v1::Span, +) -> opentelemetry_proto::tonic::collector::trace::v1::ExportTraceServiceRequest { + use opentelemetry_proto::tonic::{ + collector::trace::v1::ExportTraceServiceRequest, + trace::v1::{ResourceSpans, ScopeSpans}, + }; + ExportTraceServiceRequest { + resource_spans: vec![ResourceSpans { + scope_spans: vec![ScopeSpans { + spans: vec![span], + ..Default::default() + }], + ..Default::default() + }], + } +} + +#[rstest::fixture] +fn span() -> opentelemetry_proto::tonic::trace::v1::Span { + opentelemetry_proto::tonic::trace::v1::Span { + trace_id: vec![1; 16], + span_id: vec![2; 8], + start_time_unix_nano: 1, + end_time_unix_nano: 2, + ..Default::default() + } +} + +#[rstest] +fn standard_json_and_protobuf_preserve_the_same_identifiers( + span: opentelemetry_proto::tonic::trace::v1::Span, +) { + use prost::Message; + let request = request_with(span); + let json = serde_json::to_vec(&request).unwrap(); + let binary = request.encode_to_vec(); + let json_spans = decode_otlp(&json, Some("application/json; charset=utf-8")).unwrap(); + let binary_spans = decode_otlp(&binary, Some("application/x-protobuf")).unwrap(); + assert_eq!( + serde_json::to_value(&json_spans).unwrap(), + serde_json::to_value(&binary_spans).unwrap() + ); + assert_eq!(json_spans[0].trace_id, "01".repeat(16)); + assert_eq!(json_spans[0].span_id, "02".repeat(8)); +} + +#[rstest] +#[case::json("APPLICATION/JSON; charset=utf-8", b"{}")] +#[case::protobuf("application/x-protobuf; charset=binary", b"")] +#[case::protobuf_alias("APPLICATION/PROTOBUF", b"")] +fn supported_content_types_select_the_decoder(#[case] content_type: &str, #[case] body: &[u8]) { + assert!(decode_otlp(body, Some(content_type)).is_ok()); +} + +#[rstest] +#[case::missing_content_type(None)] +#[case::unsupported_content_type(Some("text/plain"))] +fn content_type_defaults_to_protobuf_and_rejects_unknown_values( + #[case] content_type: Option<&str>, +) { + let result = decode_otlp(b"", content_type); + assert_eq!(result.is_ok(), content_type.is_none()); +} + +#[rstest] +#[case::short_trace(vec![1; 15], vec![2;8], 1, 2)] +#[case::zero_trace(vec![0; 16], vec![2;8], 1, 2)] +#[case::short_span(vec![1; 16], vec![2;7], 1, 2)] +#[case::timestamp_overflow(vec![1;16], vec![2;8], i64::MAX as u64 + 1, i64::MAX as u64 + 1)] +#[case::negative_duration(vec![1;16], vec![2;8], 3, 2)] +fn rejects_ids_and_timestamps_that_cannot_be_stored( + #[case] trace_id: Vec, + #[case] span_id: Vec, + #[case] start: u64, + #[case] end: u64, +) { + use prost::Message; + let span = opentelemetry_proto::tonic::trace::v1::Span { + trace_id, + span_id, + start_time_unix_nano: start, + end_time_unix_nano: end, + ..Default::default() + }; + assert!(matches!( + decode_otlp(&request_with(span).encode_to_vec(), None), + Err(litellm_traces::DecodeError::InvalidPayload) + )); +} + +#[rstest] +fn resource_fanout_shares_one_allocation(span: opentelemetry_proto::tonic::trace::v1::Span) { + use opentelemetry_proto::tonic::{ + common::v1::{AnyValue, KeyValue, any_value::Value}, + resource::v1::Resource, + }; + use prost::Message; + let mut request = request_with(span.clone()); + request.resource_spans[0].resource = Some(Resource { + attributes: vec![KeyValue { + key: "shared".into(), + value: Some(AnyValue { + value: Some(Value::StringValue("x".repeat(16 * 1024))), + }), + ..Default::default() + }], + ..Default::default() + }); + request.resource_spans[0].scope_spans[0].spans = vec![span; 1024]; + let second_scope = request.resource_spans[0].scope_spans[0].clone(); + request.resource_spans[0].scope_spans.push(second_scope); + request + .resource_spans + .push(request.resource_spans[0].clone()); + let body = request.encode_to_vec(); + let decoded = decode_otlp(&body, None).expect("shared resources do not expand with span count"); + assert_eq!(decoded.len(), 4096); + assert!(decoded[..2048].iter().all(|span| { + Shared::shares_storage_with(&span.resource_attributes, &decoded[0].resource_attributes) + })); + assert!(!Shared::shares_storage_with( + &decoded[0].resource_attributes, + &decoded[2048].resource_attributes + )); + assert_eq!( + *decoded[0].resource_attributes, + *decoded[2048].resource_attributes + ); +} + +#[rstest] +fn nested_values_are_serialized_once(span: opentelemetry_proto::tonic::trace::v1::Span) { + use opentelemetry_proto::tonic::common::v1::{ + AnyValue, ArrayValue, KeyValue, any_value::Value, + }; + use prost::Message; + let nested = (0..8).fold( + AnyValue { + value: Some(Value::StringValue("quoted \"value\"".into())), + }, + |child, _| AnyValue { + value: Some(Value::ArrayValue(ArrayValue { + values: vec![child], + })), + }, + ); + let mut request = request_with(span); + request.resource_spans[0].scope_spans[0].spans[0].attributes = vec![KeyValue { + key: "nested".into(), + value: Some(nested), + ..Default::default() + }]; + let spans = decode_otlp(&request.encode_to_vec(), None).unwrap(); + let expected = (0..8).fold(serde_json::json!("quoted \"value\""), |child, _| { + serde_json::json!([child]) + }); + assert_eq!( + serde_json::from_str::(&spans[0].attributes["nested"]).unwrap(), + expected + ); + assert!(spans[0].attributes["nested"].len() < 64); +} + +#[rstest] +#[case::nesting(format!("{}0{}", "[".repeat(40), "]".repeat(40)).into_bytes())] +#[case::nodes(format!("[{}]", vec!["0"; 65537].join(",")).into_bytes())] +fn rejects_json_structure_before_building_a_tree(#[case] body: Vec) { + assert!(matches!( + decode_otlp(&body, Some("application/json")), + Err(litellm_traces::DecodeError::TooLarge) + )); +} + +#[rstest] +#[case::depth(40, 1)] +#[case::nodes(0, 65537)] +fn protobuf_preflight_rejects_expansion_before_prost_allocates( + span: opentelemetry_proto::tonic::trace::v1::Span, + #[case] depth: usize, + #[case] count: usize, +) { + use opentelemetry_proto::tonic::common::v1::{ + AnyValue, ArrayValue, KeyValue, any_value::Value, + }; + use prost::Message; + let value = (0..depth).fold( + AnyValue { + value: Some(Value::BoolValue(true)), + }, + |child, _| AnyValue { + value: Some(Value::ArrayValue(ArrayValue { + values: vec![child], + })), + }, + ); + let mut request = request_with(span); + request.resource_spans[0].scope_spans[0].spans[0].attributes = vec![KeyValue { + key: "deep".into(), + value: Some(value), + ..Default::default() + }]; + request.resource_spans = vec![request.resource_spans[0].clone(); count]; + let body = request.encode_to_vec(); + assert!(matches!( + decode_otlp(&body, None), + Err(litellm_traces::DecodeError::TooLarge) + )); +} + +#[rstest] +fn scope_fanout_shares_name_and_version(span: opentelemetry_proto::tonic::trace::v1::Span) { + use opentelemetry_proto::tonic::common::v1::InstrumentationScope; + use prost::Message; + let mut request = request_with(span.clone()); + request.resource_spans[0].scope_spans[0].scope = Some(InstrumentationScope { + name: "n".repeat(16 * 1024), + version: "v".repeat(16 * 1024), + ..Default::default() + }); + request.resource_spans[0].scope_spans[0].spans = vec![span; 1024]; + let decoded = decode_otlp(&request.encode_to_vec(), None).unwrap(); + assert!( + decoded + .iter() + .all(|span| Shared::shares_storage_with(&span.scope_name, &decoded[0].scope_name)) + ); + assert!( + decoded.iter().all(|span| Shared::shares_storage_with( + &span.scope_version, + &decoded[0].scope_version + )) + ); + assert_eq!(decoded[0].scope_name.len(), 16 * 1024); + assert_eq!(decoded[0].scope_version.len(), 16 * 1024); +} + +#[rstest] +fn unique_attribute_expansion_still_respects_decoded_budget( + span: opentelemetry_proto::tonic::trace::v1::Span, +) { + use opentelemetry_proto::tonic::common::v1::{AnyValue, KeyValue, any_value::Value}; + use prost::Message; + let mut request = request_with(span.clone()); + request.resource_spans[0].scope_spans[0].spans = (0..1024) + .map(|index| { + let mut span = span.clone(); + span.attributes = vec![KeyValue { + key: "unique".into(), + value: Some(AnyValue { + value: Some(Value::StringValue(format!( + "{index:04}{}", + "x".repeat(16_300) + ))), + }), + ..Default::default() + }]; + span + }) + .collect(); + let body = request.encode_to_vec(); + assert!(body.len() < 16 * 1024 * 1024); + assert!(matches!( + decode_otlp(&body, None), + Err(litellm_traces::DecodeError::TooLarge) + )); +} + +#[rstest] +fn escaped_attribute_expansion_is_bounded_below_four_mib( + span: opentelemetry_proto::tonic::trace::v1::Span, +) { + use opentelemetry_proto::tonic::common::v1::{ + AnyValue, ArrayValue, KeyValue, any_value::Value, + }; + use prost::Message; + let mut request = request_with(span); + request.resource_spans[0].scope_spans[0].spans[0].attributes = vec![KeyValue { + key: "escaped".into(), + value: Some(AnyValue { + value: Some(Value::ArrayValue(ArrayValue { + values: vec![AnyValue { + value: Some(Value::StringValue("\0".repeat(3 * 1024 * 1024))), + }], + })), + }), + ..Default::default() + }]; + let body = request.encode_to_vec(); + assert!(body.len() < 4 * 1024 * 1024); + assert!(matches!( + decode_otlp(&body, None), + Err(litellm_traces::DecodeError::TooLarge) + )); } diff --git a/litellm-rust/crates/traces/tests/shared.rs b/litellm-rust/crates/traces/tests/shared.rs new file mode 100644 index 00000000000..2e76e6321db --- /dev/null +++ b/litellm-rust/crates/traces/tests/shared.rs @@ -0,0 +1,23 @@ +use litellm_traces::Shared; +use rstest::rstest; + +#[rstest] +fn clones_preserve_values_and_serialize_transparently() { + let original = Shared::new(vec!["value".to_owned()]); + let cloned = original.clone(); + assert_eq!(cloned.as_ref(), original.as_ref()); + assert_eq!( + serde_json::to_value(&cloned).unwrap(), + serde_json::json!(["value"]) + ); +} + +#[rstest] +fn clones_share_storage_without_merging_equal_values() { + let original = Shared::new("value".to_owned()); + let cloned = original.clone(); + let equal = Shared::new("value".to_owned()); + assert!(original.shares_storage_with(&cloned)); + assert!(!original.shares_storage_with(&equal)); + assert_eq!(*original, *equal); +} diff --git a/litellm/constants.py b/litellm/constants.py index 7e1e63a112b..76ab419272f 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -52,10 +52,10 @@ CLICKHOUSE_MAX_BUFFERED_ROWS: Final = get_env_int("CLICKHOUSE_MAX_BUFFERED_ROWS" CLICKHOUSE_MAX_RETRIES: Final = get_env_int("CLICKHOUSE_MAX_RETRIES", 3) AGENT_TRACING_RETENTION_DAYS: Final = get_env_int("AGENT_TRACING_RETENTION_DAYS", 30) AGENT_TRACING_SPEND_LOG_RETENTION_DAYS: Final = get_env_int("AGENT_TRACING_SPEND_LOG_RETENTION_DAYS", 90) -OTLP_MAX_BODY_BYTES: Final = get_env_int("OTLP_MAX_BODY_BYTES", 8 * 1024 * 1024) +OTLP_MAX_BODY_BYTES: Final = get_env_int("OTLP_MAX_BODY_BYTES", 16 * 1024 * 1024) OTLP_MAX_ATTRIBUTE_VALUE_BYTES: Final = get_env_int("OTLP_MAX_ATTRIBUTE_VALUE_BYTES", 64 * 1024) OTLP_RETRY_AFTER_SECONDS: Final = get_env_int("OTLP_RETRY_AFTER_SECONDS", 2) -OTLP_OFFLOAD_DECODE_BYTES: Final = get_env_int("OTLP_OFFLOAD_DECODE_BYTES", 256 * 1024) +OTLP_MAX_CONCURRENT_INGESTS: Final = get_env_int("OTLP_MAX_CONCURRENT_INGESTS", 2) AGENT_TRACING_INPUT_PREVIEW_CHARS: Final = get_env_int("AGENT_TRACING_INPUT_PREVIEW_CHARS", 240) AGENT_TRACING_LIST_PAGE_SIZE: Final = get_env_int("AGENT_TRACING_LIST_PAGE_SIZE", 50) DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10)) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 6e5114b2f87..aa4d6a39f25 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -6,6 +6,7 @@ from typing import Annotated, Any, Final, Literal, Union, get_args, get_origin import orjson from fastapi import Request, UploadFile, status +from starlette._utils import get_route_path from typing_extensions import NotRequired, ReadOnly, Required, assert_never from litellm._logging import verbose_proxy_logger @@ -167,6 +168,10 @@ def _parse_binary_body(body: bytes) -> dict: return {} +def is_otlp_trace_request(request: Request) -> bool: + return request.method == "POST" and get_route_path(request.scope) == "/v1/traces" + + async def _read_request_body(request: Request | None) -> dict: """ Safely read the request body and parse it as JSON. @@ -181,6 +186,9 @@ async def _read_request_body(request: Request | None) -> dict: if request is None: return {} + if is_otlp_trace_request(request): + return {} + # Check if we already read and parsed the body _cached_request_body: Final[dict | None] = _safe_get_request_parsed_body(request=request) if _cached_request_body is not None: @@ -189,11 +197,7 @@ async def _read_request_body(request: Request | None) -> dict: _request_headers: Final[dict] = _safe_get_request_headers(request=request) content_type: Final = _request_headers.get("content-type", "") - if _normalize_media_type(content_type) in _BINARY_CONTENT_TYPES or ( - request.scope.get("path") == "/v1/traces" - and request.scope.get("method") == "POST" - and _request_headers.get("content-encoding", "").lower() == "gzip" - ): + if _normalize_media_type(content_type) in _BINARY_CONTENT_TYPES: parsed_body = _parse_binary_body(await request.body()) elif _is_form_content_type(content_type): try: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index cf09bdbef9b..9abf949ec1e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -720,6 +720,9 @@ try: except ImportError: build_billing_metrics_recorder = None shutdown_billing_metrics_recorder = None +from fastapi.exception_handlers import http_exception_handler +from starlette.exceptions import HTTPException as StarletteHTTPException + from litellm.proxy import tracing_endpoints from litellm.proxy.middleware.admission_control_middleware import ( AdmissionControlMiddleware, @@ -1895,6 +1898,9 @@ async def openai_exception_handler(request: Request, exc: ProxyException): ) status_code: Final = int(exc.code) if exc.code else status.HTTP_500_INTERNAL_SERVER_ERROR _close_dangling_otel_server_span(request, status_code, exc=exc) + otlp_response: Final = tracing_endpoints.otlp_error_response(request, status_code, headers) + if otlp_response is not None: + return otlp_response return JSONResponse( status_code=status_code, content={"error": error_dict}, @@ -1902,6 +1908,15 @@ async def openai_exception_handler(request: Request, exc: ProxyException): ) +@app.exception_handler(StarletteHTTPException) +async def otlp_http_exception_handler(request: Request, exc: StarletteHTTPException) -> Response: + response: Final = tracing_endpoints.otlp_error_response(request, exc.status_code, exc.headers) + if response is not None: + _close_dangling_otel_server_span(request, exc.status_code, exc=exc) + return response + return await http_exception_handler(request, exc) + + def _log_model_access_denial(exc: ProxyException) -> None: if not isinstance(exc, ModelAccessDeniedProxyException): return @@ -2023,6 +2038,9 @@ async def otel_request_validation_exception_handler(request: Request, exc: Reque _close_dangling_otel_server_span(request, problem.status, exc=public_exc) return problem_response(problem) _close_dangling_otel_server_span(request, 422, exc=public_exc) + otlp_response: Final = tracing_endpoints.otlp_error_response(request, 422) + if otlp_response is not None: + return otlp_response return JSONResponse(status_code=422, content={"detail": public_errors}) @@ -2046,6 +2064,9 @@ async def otel_unhandled_exception_handler(request: Request, exc: Exception): ) ) _close_dangling_otel_server_span(request, 500, exc=exc) + otlp_response: Final = tracing_endpoints.otlp_error_response(request, 500) + if otlp_response is not None: + return otlp_response return JSONResponse( status_code=500, content={ diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index bd885282859..46d29c50c1b 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -8,14 +8,18 @@ GET /v1/traces/{trace_id}/spans/{span_id} SpanDetail """ import time +from collections.abc import Mapping from dataclasses import dataclass +from http.client import responses +from types import MappingProxyType from typing import Annotated, Final from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response -from litellm.constants import OTLP_MAX_BODY_BYTES, OTLP_RETRY_AFTER_SECONDS +from litellm.constants import OTLP_RETRY_AFTER_SECONDS from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.http_parsing_utils import is_otlp_trace_request from litellm.proxy.tracing_runtime import provide_receiver, require_receiver from litellm.tracing import ( Tenant, @@ -23,7 +27,7 @@ from litellm.tracing import ( TracingPayloadTooLargeError, ) from litellm.tracing.decode import InvalidOTLPPayloadError, encode_otlp_response -from litellm.tracing.types import SpanDetail, Trace, TracePage, TraceScope +from litellm.tracing.types import SpanDetail, SpanErrorPage, Trace, TracePage, TraceScope router = APIRouter(tags=["agent tracing"]) @@ -66,13 +70,25 @@ async def provide_trace_access( return TraceAccessContext(tracing, None, tenant) -async def _read_otlp_body(request: Request) -> bytes: - body: Final = bytearray() - async for chunk in request.stream(): - if len(body) + len(chunk) > OTLP_MAX_BODY_BYTES: - raise TracingPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") - body.extend(chunk) - return bytes(body) +def otlp_error_response( + request: Request, status_code: int, headers: Mapping[str, str] | None = None +) -> Response | None: + if not is_otlp_trace_request(request): + return None + body, media_type = encode_otlp_response( + request.headers.get("content-type"), responses.get(status_code, "Trace request failed") + ) + return Response(content=body, status_code=status_code, media_type=media_type, headers=headers) + + +def _otlp_error(content_type: str | None, status_code: int, message: str, retry: bool = False) -> Response: + body, media_type = encode_otlp_response(content_type, message) + return Response( + content=body, + status_code=status_code, + media_type=media_type, + headers=MappingProxyType({"Retry-After": str(OTLP_RETRY_AFTER_SECONDS)}) if retry else None, + ) @router.post("/v1/traces", include_in_schema=False) @@ -80,24 +96,23 @@ async def ingest_otlp_traces( request: Request, context: Annotated[TraceAccessContext, Depends(provide_trace_access)], ) -> Response: - tracing, tenant = context.writer() content_type: Final = request.headers.get("content-type") try: + tracing, tenant = context.writer() await tracing.ingest( - body=await _read_otlp_body(request), + body=request.stream(), content_type=content_type, content_encoding=request.headers.get("content-encoding"), tenant=tenant, ) except TracingPayloadTooLargeError as e: - raise HTTPException(status_code=413, detail=str(e)) + return _otlp_error(content_type, 413, str(e)) except InvalidOTLPPayloadError as error: - raise HTTPException(status_code=400, detail=str(error)) from error + return _otlp_error(content_type, 400, str(error)) except RuntimeError: - raise HTTPException( - status_code=503, - headers={"Retry-After": str(OTLP_RETRY_AFTER_SECONDS)}, - ) + return _otlp_error(content_type, 503, "Trace ingestion is temporarily unavailable", retry=True) + except HTTPException as error: + return _otlp_error(content_type, error.status_code, str(error.detail)) body, media_type = encode_otlp_response(content_type) return Response(content=body, media_type=media_type) @@ -147,3 +162,21 @@ async def get_agent_trace_span( if span is None: raise HTTPException(status_code=404, detail=f"Span {span_id} not found") return span + + +@router.get("/v1/traces/{trace_id}/spans/{span_id}/error", response_model=SpanErrorPage) +async def get_agent_trace_span_error( + trace_id: str, + span_id: str, + context: Annotated[TraceAccessContext, Depends(provide_trace_access)], + trace_ref: Annotated[str, Query()] = "", + cursor: Annotated[str | None, Query(max_length=512)] = None, +) -> SpanErrorPage: + try: + tracing, scope = context.reader() + page: Final = await tracing.get_span_error(trace_id, span_id, scope, trace_ref, cursor) + except ValueError as error: + raise HTTPException(status_code=400, detail=str(error)) from error + if page is None: + raise HTTPException(status_code=404, detail="Span diagnostic not found or no longer available") + return page diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 206c0f78ed8..a3d8ba0e582 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -21,15 +21,14 @@ class RustUpstreamError(Exception): ... class ForkedAfterNativeRuntimeStarted(RuntimeError): ... class ProcessReservedForForking(RuntimeError): ... -def trace_decode_otlp( - body: bytes, content_type: str | None, content_encoding: str | None, max_decompressed_bytes: int -) -> list[DecodedSpan]: ... +def trace_decode_otlp(body: bytes, content_type: str | None) -> list[DecodedSpan]: ... +def trace_encode_error(message: str) -> bytes: ... @final class NativeTraceStorage: def __new__(cls, database: str, url: str, reader_url: str | None = None) -> NativeTraceStorage: ... def ensure_schema(self, trace_retention_days: int, spend_log_retention_days: int) -> Future[None]: ... - def insert_rows(self, table: str, rows: Sequence[Mapping[str, JsonValue]]) -> Future[None]: ... + def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> Future[None]: ... def lens_query(self, name: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ... def query(self, query: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ... @@ -353,6 +352,7 @@ __all__ = [ "reserve_process_for_forking", "responses", "trace_decode_otlp", + "trace_encode_error", "transcription", ] diff --git a/litellm/rust_bridge/traces.py b/litellm/rust_bridge/traces.py index 1c20e408709..6724db41ad3 100644 --- a/litellm/rust_bridge/traces.py +++ b/litellm/rust_bridge/traces.py @@ -31,7 +31,7 @@ class DecodedSpan(TypedDict): events: ReadOnly[list[DecodedEvent]] -ReadQueryName = Literal["list_traces", "trace_spans", "span_detail", "spend_by_response_ids"] +ReadQueryName = Literal["list_traces", "trace_spans", "span_detail", "span_error", "spend_by_response_ids"] class NativeStore(Protocol): @@ -39,7 +39,7 @@ class NativeStore(Protocol): def ensure_schema(self, trace_retention_days: int, spend_log_retention_days: int) -> Awaitable[None]: ... - def insert_rows(self, table: str, rows: Sequence[Mapping[str, JsonValue]]) -> Awaitable[None]: ... + def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> Awaitable[None]: ... def lens_query(self, name: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Awaitable[str]: ... @@ -53,17 +53,16 @@ class NativeTraces(Protocol): self, body: bytes, content_type: str | None, - content_encoding: str | None, - max_decompressed_bytes: int, ) -> list[DecodedSpan]: ... + def trace_encode_error(self, message: str) -> bytes: ... + class QueryResponse(BaseModel): model_config = ConfigDict(frozen=True) data: list[dict[str, JsonValue]] -INSERT_ROWS: Final = TypeAdapter(list[dict[str, JsonValue]]) QUERY_PARAMETERS: Final = TypeAdapter(dict[str, str | int | list[str]]) @@ -74,10 +73,14 @@ def _native() -> NativeTraces: return cast(NativeTraces, native) # cast-ok: the native extension is validated against this protocol at call sites -def decode_otlp( - body: bytes, content_type: str | None, content_encoding: str | None, max_decompressed_bytes: int -) -> list[DecodedSpan]: - return _native().trace_decode_otlp(body, content_type, content_encoding, max_decompressed_bytes) +def decode_otlp(body: bytes, content_type: str | None) -> list[DecodedSpan]: + return _native().trace_decode_otlp(body, content_type) + + +def encode_error(message: str) -> bytes: + if get_native_bridge() is None: + return b"" + return _native().trace_encode_error(message) class ClickHouseStorage: @@ -88,7 +91,7 @@ class ClickHouseStorage: await self._native.ensure_schema(trace_retention_days, spend_log_retention_days) async def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> None: - await self._native.insert_rows(table, INSERT_ROWS.validate_python(rows)) + await self._native.insert_rows(table, rows) async def query( self, name: ReadQueryName, parameters: Mapping[str, object] | None = None diff --git a/litellm/tracing/decode.py b/litellm/tracing/decode.py index ce46bed5e22..d8b5f70de68 100644 --- a/litellm/tracing/decode.py +++ b/litellm/tracing/decode.py @@ -8,37 +8,43 @@ Pure functions, no I/O. Two steps: Deep Agents), OTEL GenAI semconv, OpenInference. """ +import gzip import json +import zlib from collections.abc import Mapping +from dataclasses import dataclass +from io import BytesIO from itertools import accumulate from types import MappingProxyType from typing import Final from pydantic import JsonValue, TypeAdapter, ValidationError -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm.constants import OTLP_MAX_ATTRIBUTE_VALUE_BYTES, OTLP_MAX_BODY_BYTES from litellm.rust_bridge.traces import DecodedSpan from litellm.rust_bridge.traces import decode_otlp as native_decode_otlp -from litellm.tracing.normalizers import select_normalizer -from litellm.tracing.normalizers.base import to_int -from litellm.tracing.types import SpanRow +from litellm.rust_bridge.traces import encode_error as native_encode_error +from litellm.tracing.normalizers.messages import content_text +from litellm.tracing.types import SpanRow, SpanType +_FRAMEWORK_SUFFIXES: Final = ( + ".wrap_model_call", + ".wrap_tool_call", + ".before_agent", + ".after_agent", + ".before_model", + ".after_model", +) +_LLM_OPERATIONS: Final = frozenset({"chat", "text_completion", "generate_content"}) +_LC_ROLES: Final = MappingProxyType({"human": "user", "ai": "assistant", "system": "system", "tool": "tool"}) +_OPENINFERENCE_TYPES: Final[Mapping[str, SpanType]] = MappingProxyType({"AGENT": "agent", "LLM": "llm", "TOOL": "tool"}) + + +_JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) _MESSAGE_LIST: Final = TypeAdapter(tuple[dict[str, JsonValue], ...]) _MAX_JSON_ESCAPE_BYTES: Final = 6 - -# attributes whose content we lift into Input/Output and drop from SpanAttributes -_HEAVY_ATTRIBUTES: Final = frozenset( - { - "gen_ai.prompt", - "gen_ai.completion", - "gen_ai.tool.definitions", - "gen_ai.input.messages", - "gen_ai.output.messages", - "input.value", - "output.value", - } -) +_MAX_TOKENS: Final = (1 << 32) - 1 class InvalidOTLPPayloadError(ValueError): @@ -49,12 +55,39 @@ class OTLPPayloadTooLargeError(OverflowError): pass +class MessageExtras(TypedDict): + tool_calls: ReadOnly[NotRequired[JsonValue]] + name: ReadOnly[NotRequired[str]] + + +class NormalizedMessage(MessageExtras): + role: ReadOnly[str] + content: ReadOnly[str] + + +class OTLPError(TypedDict): + message: ReadOnly[str] + + +@dataclass(frozen=True, slots=True) +class NormalizedSpan: + kind: SpanType + agent: str = "" + model: str = "" + request_id: str = "" + input: str = "" + output: str = "" + input_tokens: int = 0 + output_tokens: int = 0 + consumed: frozenset[str] = frozenset() + + def _truncate(value: str) -> str: - size = len(value.encode("utf-8")) - if size <= OTLP_MAX_ATTRIBUTE_VALUE_BYTES: + encoded: Final = value.encode("utf-8") + if len(encoded) <= OTLP_MAX_ATTRIBUTE_VALUE_BYTES: return value - kept = value.encode("utf-8")[:OTLP_MAX_ATTRIBUTE_VALUE_BYTES].decode("utf-8", "ignore") - return f"{kept}…[truncated {size - OTLP_MAX_ATTRIBUTE_VALUE_BYTES} bytes]" + kept: Final = encoded[:OTLP_MAX_ATTRIBUTE_VALUE_BYTES].decode("utf-8", "ignore") + return f"{kept}…[truncated {len(encoded) - len(kept.encode('utf-8'))} bytes]" def _size(value: str) -> int: @@ -137,9 +170,9 @@ def _truncate_payload(value: str) -> str: def decode_otlp( body: bytes, content_type: str | None = None, content_encoding: str | None = None ) -> tuple[SpanRow, ...]: - """Decode an OTLP trace export and normalize every span.""" + payload: Final = _decode_content_encoding(body, content_encoding) try: - spans: Final = native_decode_otlp(body, content_type, content_encoding, OTLP_MAX_BODY_BYTES) + spans: Final = native_decode_otlp(payload, content_type) except OverflowError as error: raise OTLPPayloadTooLargeError(str(error)) from error except ValueError as error: @@ -147,19 +180,34 @@ def decode_otlp( return tuple(_span_row(span) for span in spans) +def _decode_content_encoding(body: bytes, content_encoding: str | None) -> bytes: + if len(body) > OTLP_MAX_BODY_BYTES: + raise OTLPPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") + if content_encoding is None or content_encoding.lower() == "identity": + return body + if content_encoding.lower() != "gzip": + raise InvalidOTLPPayloadError("Unsupported OTLP content encoding") + try: + with gzip.GzipFile(fileobj=BytesIO(body)) as stream: + payload: Final = stream.read(OTLP_MAX_BODY_BYTES + 1) + except (EOFError, OSError, zlib.error) as error: + raise InvalidOTLPPayloadError("Invalid OTLP gzip body") from error + if len(payload) > OTLP_MAX_BODY_BYTES: + raise OTLPPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") + return payload + + def _exception_message(span: DecodedSpan) -> str: - """`span.record_exception()` writes an `exception` event; surface it when status.message is empty.""" for event in span["events"]: if event["name"] == "exception": - attributes = event["attributes"] - return attributes.get("exception.message") or attributes.get("exception.type", "") + return event["attributes"].get("exception.message") or event["attributes"].get("exception.type", "") return "" def _span_row(span: DecodedSpan) -> SpanRow: - attributes = span["attributes"] - resource = span["resource_attributes"] - row = SpanRow( + attributes: Final = span["attributes"] + normalized: Final = normalize(span) + return SpanRow( Timestamp=span["start_ns"], TraceId=span["trace_id"], SpanId=span["span_id"], @@ -167,44 +215,201 @@ def _span_row(span: DecodedSpan) -> SpanRow: TraceState=span["trace_state"], SpanName=span["name"], SpanKind=span["kind"], - ServiceName=resource.get("service.name", ""), - ResourceAttributes=resource, + ServiceName=span["resource_attributes"].get("service.name", ""), + ResourceAttributes=span["resource_attributes"], ScopeName=span["scope_name"], ScopeVersion=span["scope_version"], - SpanAttributes=attributes, - Duration=max(span["end_ns"] - span["start_ns"], 0), + SpanAttributes=MappingProxyType( + {key: _truncate(value) for key, value in attributes.items() if key not in normalized.consumed} + ), + Duration=span["end_ns"] - span["start_ns"], StatusCode=span["status_code"], StatusMessage=span["status_message"] or _exception_message(span), TeamId="", ApiKeyHash="", - ObservationType="chain", - AgentName="", - LiteLLMRequestId="", - Model="", - InputTokens=0, - OutputTokens=0, - Input="", - Output="", + ObservationType=normalized.kind, + AgentName=normalized.agent, + Model=normalized.model, + LiteLLMRequestId=attributes.get("gen_ai.response.id") or normalized.request_id, + InputTokens=normalized.input_tokens, + OutputTokens=normalized.output_tokens, + Input=_truncate_payload(normalized.input), + Output=_truncate(normalized.output), ) - normalize(row, attributes) - row["SpanAttributes"] = {k: _truncate(v) for k, v in attributes.items() if k not in _HEAVY_ATTRIBUTES} - row["Input"], row["Output"] = _truncate_payload(row["Input"]), _truncate(row["Output"]) - return row -def _set_tokens(row: SpanRow, attributes: Mapping[str, str]) -> None: - row["InputTokens"] = to_int(attributes.get("gen_ai.usage.input_tokens")) - row["OutputTokens"] = to_int(attributes.get("gen_ai.usage.output_tokens")) +def _loads(value: str) -> JsonValue: + if len(value.encode("utf-8")) > OTLP_MAX_BODY_BYTES: + return None + try: + return _JSON.validate_json(value) + except ValidationError: + return None -def normalize(row: SpanRow, attributes: Mapping[str, str]) -> None: - select_normalizer(row["ScopeName"], attributes).normalize(row, attributes) - if not row["InputTokens"] and not row["OutputTokens"]: - _set_tokens(row, attributes) +def _text(value: JsonValue) -> str: + return value if isinstance(value, str) else "" -def encode_otlp_response(content_type: str | None) -> tuple[bytes, str]: - """Empty ExportTraceServiceResponse in the caller's encoding.""" - if content_type and "json" in content_type: - return b"{}", "application/json" - return b"", "application/x-protobuf" +def _message(value: JsonValue) -> NormalizedMessage | None: + if not isinstance(value, dict): + return None + kwargs: Final = value.get("kwargs", value) + if not isinstance(kwargs, dict): + return None + kind: Final = _text(kwargs.get("type")) or _text(kwargs.get("role")) + if not kind: + return None + calls: Final = kwargs.get("tool_calls") + if calls is not None and (not isinstance(calls, list) or not all(isinstance(call, dict) for call in calls)): + return None + role: Final = _LC_ROLES.get(kind, kind) + content: Final = kwargs.get("content", "") + name: Final = kwargs.get("name") + tool_calls: Final = MessageExtras(tool_calls=calls) if calls else MessageExtras() + tool_name: Final = MessageExtras(name=name) if role == "tool" and isinstance(name, str) else MessageExtras() + message: Final[NormalizedMessage] = { + "role": role, + "content": content_text(content), + **tool_calls, + **tool_name, + } + return message + + +def _messages(value: JsonValue, raw: str) -> str: + if not isinstance(value, list): + return raw + messages: Final = tuple(_message(item) for item in value) + return json.dumps(messages) if all(message is not None for message in messages) else raw + + +def _langsmith_type(span: DecodedSpan) -> SpanType: + attributes: Final = span["attributes"] + kind: Final = attributes.get("langsmith.span.kind", "chain") + if kind in ("llm", "tool"): + return "llm" if kind == "llm" else "tool" + if not span["parent_span_id"] or span["name"] == attributes.get("langsmith.metadata.lc_agent_name"): + return "agent" + return "framework" if span["name"].endswith(_FRAMEWORK_SUFFIXES) else "chain" + + +def _langsmith_io(kind: SpanType, attributes: Mapping[str, str]) -> tuple[str, str, str]: + raw_prompt: Final = attributes.get("gen_ai.prompt", "") + raw_completion: Final = attributes.get("gen_ai.completion", "") + prompt: Final = _loads(raw_prompt) + completion: Final = _loads(raw_completion) + messages: Final = prompt.get("messages") if isinstance(prompt, dict) else None + if kind == "llm": + batch: Final = ( + messages[0] if isinstance(messages, list) and messages and isinstance(messages[0], list) else messages + ) + generations: Final = completion.get("generations") if isinstance(completion, dict) else None + first: Final = generations[0] if isinstance(generations, list) and generations else None + item: Final = first[0] if isinstance(first, list) and first else first + message: Final = item.get("message") if isinstance(item, dict) else None + parsed: Final = _message(message) + kwargs: Final = message.get("kwargs", message) if isinstance(message, dict) else None + metadata: Final = kwargs.get("response_metadata") if isinstance(kwargs, dict) else None + request_id: Final = _text(metadata.get("id")) if isinstance(metadata, dict) else "" + return _messages(batch, raw_prompt), json.dumps(parsed) if parsed is not None else raw_completion, request_id + if kind == "tool": + output: Final = completion.get("output", completion) if isinstance(completion, dict) else completion + update: Final = output.get("update") if isinstance(output, dict) else None + updates: Final = update.get("messages") if isinstance(update, dict) else None + final: Final = updates[-1] if isinstance(updates, list) and updates else output + content: Final = final.get("content", final) if isinstance(final, dict) else final + return ( + raw_prompt, + (content if isinstance(content, str) else json.dumps(content)) if content is not None else raw_completion, + "", + ) + if kind == "agent": + outputs: Final = completion.get("messages") if isinstance(completion, dict) else None + last: Final = _message(outputs[-1]) if isinstance(outputs, list) and outputs else None + return _messages(messages, raw_prompt), json.dumps(last) if last is not None else raw_completion, "" + return raw_prompt, raw_completion, "" + + +def _to_int(value: str | None) -> int: + try: + number: Final = int(value) if value else 0 + except ValueError: + return 0 + if not 0 <= number <= _MAX_TOKENS: + raise InvalidOTLPPayloadError("OTLP token count is outside the storage range") + return number + + +def normalize(span: DecodedSpan) -> NormalizedSpan: + attributes: Final = span["attributes"] + fallback: Final[SpanType] = "agent" if not span["parent_span_id"] else "chain" + input_tokens: Final = _to_int(attributes.get("gen_ai.usage.input_tokens")) + output_tokens: Final = _to_int(attributes.get("gen_ai.usage.output_tokens")) + if span["scope_name"] == "langsmith" or "langsmith.span.kind" in attributes: + kind: Final = _langsmith_type(span) + prompt, completion, request_id = _langsmith_io(kind, attributes) + return NormalizedSpan( + kind, + attributes.get("langsmith.metadata.lc_agent_name", ""), + attributes.get("gen_ai.request.model", ""), + request_id, + prompt, + completion, + input_tokens, + output_tokens, + frozenset({"gen_ai.prompt", "gen_ai.completion"}), + ) + if "openinference.span.kind" in attributes: + return NormalizedSpan( + _OPENINFERENCE_TYPES.get(attributes["openinference.span.kind"].upper(), fallback), + attributes.get("agent.name", ""), + attributes.get("llm.model_name", ""), + "", + attributes.get("input.value", ""), + attributes.get("output.value", ""), + _to_int(attributes.get("llm.token_count.prompt")) + if "llm.token_count.prompt" in attributes + else input_tokens, + _to_int(attributes.get("llm.token_count.completion")) + if "llm.token_count.completion" in attributes + else output_tokens, + frozenset({"input.value", "output.value"}), + ) + operation: Final = attributes.get("gen_ai.operation.name", "") + genai_kind: Final[SpanType] = ( + "llm" + if operation in _LLM_OPERATIONS + else "tool" + if operation == "execute_tool" + else "agent" + if operation == "invoke_agent" + else fallback + ) + input_key: Final = ( + "gen_ai.input.messages" if attributes.get("gen_ai.input.messages") else "gen_ai.tool.call.arguments" + ) + output_key: Final = ( + "gen_ai.output.messages" if attributes.get("gen_ai.output.messages") else "gen_ai.tool.call.result" + ) + return NormalizedSpan( + genai_kind, + attributes.get("gen_ai.agent.name", ""), + attributes.get("gen_ai.request.model") or attributes.get("gen_ai.response.model", ""), + "", + attributes.get(input_key, ""), + attributes.get(output_key, ""), + input_tokens, + output_tokens, + frozenset({input_key, output_key}), + ) + + +def encode_otlp_response(content_type: str | None, error: str | None = None) -> tuple[bytes, str]: + media_type: Final = (content_type or "application/x-protobuf").split(";", 1)[0].strip().lower() + if media_type == "application/json": + response: Final[OTLPError] = {"message": error or ""} + return (json.dumps(response).encode() if error else b"{}"), "application/json" + if error is None: + return b"", "application/x-protobuf" + return native_encode_error(error), "application/x-protobuf" diff --git a/litellm/tracing/receiver.py b/litellm/tracing/receiver.py index 700f33a8a0e..6cef84ec6d0 100644 --- a/litellm/tracing/receiver.py +++ b/litellm/tracing/receiver.py @@ -14,13 +14,17 @@ The proxy endpoints are thin wrappers: auth -> build tenant/scope -> call one me import asyncio import os +from collections.abc import AsyncIterable, Callable, Mapping +from io import BytesIO +from threading import BoundedSemaphore +from types import MappingProxyType from typing import Final from litellm.constants import ( AGENT_TRACING_RETENTION_DAYS, AGENT_TRACING_SPEND_LOG_RETENTION_DAYS, OTLP_MAX_BODY_BYTES, - OTLP_OFFLOAD_DECODE_BYTES, + OTLP_MAX_CONCURRENT_INGESTS, ) from litellm.integrations.clickhouse.schema import ensure_schema from litellm.rust_bridge.traces import ClickHouseStorage @@ -28,6 +32,7 @@ from litellm.tracing.decode import OTLPPayloadTooLargeError, decode_otlp from litellm.tracing.store import TraceStore from litellm.tracing.types import ( SpanDetail, + SpanErrorPage, SpanRow, Trace, TracePage, @@ -39,6 +44,10 @@ class TracingPayloadTooLargeError(Exception): pass +class TracingOverloadedError(RuntimeError): + pass + + class Tenant: """Who sent the spans. Always taken from auth, never from span attributes.""" @@ -48,20 +57,49 @@ class Tenant: self.org_id = org_id def stamp(self, row: SpanRow) -> SpanRow: - row["TeamId"] = self.team_id - row["ApiKeyHash"] = self.api_key_hash - row["ResourceAttributes"] = { - **row["ResourceAttributes"], - "litellm.team_id": self.team_id, - "litellm.api_key_hash": self.api_key_hash, - "litellm.org_id": self.org_id, + return self.stamp_rows((row,))[0] + + def stamp_rows(self, rows: tuple[SpanRow, ...]) -> tuple[SpanRow, ...]: + resources: Final = MappingProxyType({id(row["ResourceAttributes"]): row["ResourceAttributes"] for row in rows}) + stamped: Final = MappingProxyType( + { + identity: MappingProxyType( + { + **attributes, + "litellm.team_id": self.team_id, + "litellm.api_key_hash": self.api_key_hash, + "litellm.org_id": self.org_id, + } + ) + for identity, attributes in resources.items() + } + ) + return tuple(self._stamp_row(row, stamped[id(row["ResourceAttributes"])]) for row in rows) + + def _stamp_row(self, row: SpanRow, resource: Mapping[str, str]) -> SpanRow: + stamped: Final[SpanRow] = { + **row, + "TeamId": self.team_id, + "ApiKeyHash": self.api_key_hash, + "ResourceAttributes": resource, } - return row + return stamped class TraceReceiver: - def __init__(self, store: TraceStore) -> None: + def __init__( + self, + store: TraceStore, + max_concurrent_ingests: int = OTLP_MAX_CONCURRENT_INGESTS, + decoder: Callable[[bytes, str | None, str | None], tuple[SpanRow, ...]] = decode_otlp, + body_read_timeout: float = 30, + ) -> None: + if max_concurrent_ingests < 1: + raise ValueError("OTLP ingestion concurrency must be positive") self.store = store + self._decoder: Final = decoder + self._body_read_timeout: Final = body_read_timeout + self._ingest_slots: Final = BoundedSemaphore(max_concurrent_ingests) @classmethod def from_env(cls) -> "TraceReceiver": @@ -84,24 +122,45 @@ class TraceReceiver: async def ingest( self, - body: bytes, + body: bytes | AsyncIterable[bytes], content_type: str | None, content_encoding: str | None, tenant: Tenant, ) -> int: - """Decode an OTLP trace export and store its authenticated spans.""" - if len(body) > OTLP_MAX_BODY_BYTES: + if not self._ingest_slots.acquire(blocking=False): + raise TracingOverloadedError("OTLP ingestion is at capacity") + task: Final = asyncio.create_task(self._ingest(body, content_type, content_encoding, tenant)) + task.add_done_callback(self._release_ingest) + return await asyncio.shield(task) + + def _release_ingest(self, task: asyncio.Task[int]) -> None: + self._ingest_slots.release() + if not task.cancelled(): + task.exception() + + async def _ingest( + self, + body: bytes | AsyncIterable[bytes], + content_type: str | None, + content_encoding: str | None, + tenant: Tenant, + ) -> int: + try: + payload: Final = ( + body + if isinstance(body, bytes) + else await asyncio.wait_for(_read_body(body), timeout=self._body_read_timeout) + ) + except asyncio.TimeoutError as error: + raise TracingOverloadedError("OTLP body upload timed out") from error + if len(payload) > OTLP_MAX_BODY_BYTES: raise TracingPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") try: - rows: Final = ( - await asyncio.to_thread(decode_otlp, body, content_type, content_encoding) - if len(body) > OTLP_OFFLOAD_DECODE_BYTES - else decode_otlp(body, content_type, content_encoding) - ) + rows: Final = await asyncio.to_thread(self._decoder, payload, content_type, content_encoding) except OTLPPayloadTooLargeError as error: raise TracingPayloadTooLargeError(str(error)) from error try: - await self.store.insert_spans(tuple(tenant.stamp(r) for r in rows)) + await self.store.insert_spans(tenant.stamp_rows(rows)) except OverflowError as error: raise TracingPayloadTooLargeError(str(error)) from error return len(rows) @@ -114,3 +173,17 @@ class TraceReceiver: async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None: return await self.store.get_span(trace_id, span_id, scope, trace_ref) + + async def get_span_error( + self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "", cursor: str | None = None + ) -> SpanErrorPage | None: + return await self.store.get_span_error(trace_id, span_id, scope, trace_ref, cursor) + + +async def _read_body(chunks: AsyncIterable[bytes]) -> bytes: + with BytesIO() as body: + async for chunk in chunks: + if body.tell() + len(chunk) > OTLP_MAX_BODY_BYTES: + raise TracingPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") + body.write(chunk) + return body.getvalue() diff --git a/litellm/tracing/store.py b/litellm/tracing/store.py index 9d1f64f77f0..91420ffd025 100644 --- a/litellm/tracing/store.py +++ b/litellm/tracing/store.py @@ -9,7 +9,7 @@ from itertools import chain from types import MappingProxyType from typing import Any, Final -from pydantic import BaseModel, ConfigDict, TypeAdapter +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter from litellm._logging import verbose_logger from litellm.constants import AGENT_TRACING_LIST_PAGE_SIZE @@ -21,6 +21,7 @@ from litellm.tracing.types import ( AgentNode, Span, SpanDetail, + SpanErrorPage, SpanRow, SpanStatus, Trace, @@ -35,6 +36,20 @@ SPEND_WINDOW_MS: Final = 30 * 60 * 1000 _STATUS: Final = MappingProxyType({"STATUS_CODE_OK": "ok", "STATUS_CODE_ERROR": "error"}) +class _ErrorCursor(BaseModel): + model_config = ConfigDict(frozen=True) + offset: int = Field(ge=0, le=(1 << 63) - 1) + version: str = Field(pattern=r"^[A-F0-9]{64}$") + + +class _ErrorRow(BaseModel): + model_config = ConfigDict(frozen=True) + span_id: str + message: str + total_chars: int + version: str + + class _SpendRow(BaseModel): model_config = ConfigDict(frozen=True) @@ -134,6 +149,7 @@ def span_from_row(row: dict[str, Any], trace_start_ns: int, spend_rows: Sequence duration_ms=int(row["duration_ns"]) / NANOS_PER_MS, status=_status(row["status"]), error=row.get("status_message") or None, + error_truncated=bool(row.get("error_truncated", False)), input_preview=row["input_preview"], model=row["model"] or None, input_tokens=int(row["input_tokens"]), @@ -349,3 +365,41 @@ class TraceStore: output_ui=to_ui_content(rows[0]["output"]), attributes=rows[0]["attributes"], ) + + async def get_span_error( + self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "", cursor: str | None = None + ) -> SpanErrorPage | None: + try: + position: Final = ( + _ErrorCursor.model_validate_json(base64.b64decode(cursor, altchars=b"-_", validate=True)) + if cursor + else None + ) + except (ValueError, binascii.Error) as error: + raise ValueError("Invalid diagnostic cursor") from error + rows: Final = await self.storage.query( + "span_error", + MappingProxyType( + { + **scope, + "trace_id": trace_id, + "span_id": span_id, + "trace_ref": trace_ref, + "error_offset": position.offset if position else 0, + "error_version": position.version if position else "", + } + ), + ) + if not rows: + return None + row: Final = _ErrorRow.model_validate(rows[0]) + offset: Final = (position.offset if position else 0) + len(row.message) + continuation: Final = _ErrorCursor(offset=offset, version=row.version) if offset < row.total_chars else None + return SpanErrorPage( + span_id=row.span_id, + message=row.message, + total_chars=row.total_chars, + next_cursor=base64.urlsafe_b64encode(continuation.model_dump_json().encode()).decode() + if continuation + else None, + ) diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py index fcf0d83fb8f..ff965483013 100644 --- a/litellm/tracing/types.py +++ b/litellm/tracing/types.py @@ -9,7 +9,7 @@ A trace is one agent run. It's made of spans (agent / llm / tool / chain / frame """ -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from typing import Literal from typing_extensions import NotRequired, ReadOnly, TypedDict @@ -29,7 +29,8 @@ class Span(TypedDict): start_offset_ms: ReadOnly[float] # relative to trace start duration_ms: ReadOnly[float] status: ReadOnly[SpanStatus] - error: ReadOnly[str | None] # exception message when status == "error" + error: ReadOnly[str | None] + error_truncated: ReadOnly[bool] input_preview: ReadOnly[str] model: ReadOnly[str | None] input_tokens: ReadOnly[int] @@ -91,6 +92,13 @@ class SpanDetail(TypedDict): attributes: ReadOnly[dict[str, str]] +class SpanErrorPage(TypedDict): + span_id: ReadOnly[str] + message: ReadOnly[str] + total_chars: ReadOnly[int] + next_cursor: ReadOnly[str | None] + + class TraceScope(TypedDict): """Who is asking. Empty team_ids = all teams (admins only).""" @@ -109,15 +117,15 @@ class SpanRow(TypedDict): SpanName: ReadOnly[str] SpanKind: ReadOnly[str] ServiceName: ReadOnly[str] - ResourceAttributes: dict[str, str] + ResourceAttributes: ReadOnly[Mapping[str, str]] ScopeName: ReadOnly[str] ScopeVersion: ReadOnly[str] - SpanAttributes: dict[str, str] + SpanAttributes: ReadOnly[Mapping[str, str]] Duration: ReadOnly[int] # ns StatusCode: ReadOnly[str] StatusMessage: ReadOnly[str] - TeamId: str - ApiKeyHash: str + TeamId: ReadOnly[str] + ApiKeyHash: ReadOnly[str] ObservationType: SpanType AgentName: str LiteLLMRequestId: str diff --git a/tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json b/tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json index 9bd8e67633b..48d8ef0f1dc 100644 --- a/tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json +++ b/tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json @@ -42,10 +42,10 @@ }, "spans": [ { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "XnnztbUEmF4=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "5e79f3b5b504985e", "name": "deep_research_agent", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790742989377137920", "endTimeUnixNano": "1790743040762587136", "attributes": [ @@ -123,16 +123,16 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 }, { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "imocMZQNB68=", - "parentSpanId": "Hfr3D90RhPI=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "8a6a1c31940d07af", + "parentSpanId": "1dfaf70fdd1184f2", "name": "ChatOpenAI", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790742989383207936", "endTimeUnixNano": "1790742998893985024", "attributes": [ @@ -354,16 +354,16 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 }, { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "zwThqgPzRPo=", - "parentSpanId": "g0UfMjWEf2w=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "cf04e1aa03f344fa", + "parentSpanId": "83451f3235847f6c", "name": "FilesystemMiddleware.wrap_model_call", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790742989379030016", "endTimeUnixNano": "1790742998895730944", "attributes": [ @@ -477,16 +477,16 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 }, { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "svs6j1ovzgE=", - "parentSpanId": "Vt73x+GSQ0o=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "b2fb3a8f5a2fce01", + "parentSpanId": "56def7c7e192434a", "name": "task", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790742998900896000", "endTimeUnixNano": "1790743034076956160", "attributes": [ @@ -624,16 +624,16 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 }, { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "gUmbSS/ZP4U=", - "parentSpanId": "svs6j1ovzgE=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "81499b492fd93f85", + "parentSpanId": "b2fb3a8f5a2fce01", "name": "researcher", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790742998901422080", "endTimeUnixNano": "1790743034076699904", "attributes": [ @@ -759,16 +759,16 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 }, { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "/mLyrQOgEWw=", - "parentSpanId": "SUm+6tN4+TU=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "fe62f2ad03a0116c", + "parentSpanId": "4949beead378f935", "name": "search_docs", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790743004976721920", "endTimeUnixNano": "1790743004977214208", "attributes": [ @@ -912,7 +912,7 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 } @@ -921,4 +921,4 @@ ] } ] -} \ No newline at end of file +} diff --git a/tests/test_litellm/tracing/test_decode.py b/tests/test_litellm/tracing/test_decode.py index 7185341d038..21a79dd6b87 100644 --- a/tests/test_litellm/tracing/test_decode.py +++ b/tests/test_litellm/tracing/test_decode.py @@ -5,13 +5,14 @@ The fixture is a trimmed real export from a Deep Agents run (LangSmith OTEL mode deep_research_agent -> task (tool) -> researcher (subagent) -> search_docs (tool). """ +import base64 import gzip import json from pathlib import Path from unittest.mock import patch import pytest -from google.protobuf.json_format import Parse +from google.protobuf.json_format import ParseDict from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest from opentelemetry.proto.common.v1.common_pb2 import AnyValue, KeyValue from opentelemetry.proto.trace.v1.trace_pb2 import ResourceSpans, ScopeSpans, Span, Status @@ -31,7 +32,14 @@ def _fixture_json() -> bytes: def _fixture_protobuf() -> bytes: request = ExportTraceServiceRequest() - Parse(_fixture_json().decode(), request) + payload = json.loads(_fixture_json()) + for resource in payload["resourceSpans"]: + for scope in resource["scopeSpans"]: + for span in scope["spans"]: + for field in ("traceId", "spanId", "parentSpanId"): + if field in span: + span[field] = base64.b64encode(bytes.fromhex(span[field])).decode() + ParseDict(payload, request) return request.SerializeToString() @@ -191,7 +199,7 @@ def test_plain_tool_input_output(rows_by_name): def test_heavy_attributes_are_lifted_out_of_span_attributes(rows_by_name): for row in rows_by_name.values(): - assert not set(row["SpanAttributes"]) & decode._HEAVY_ATTRIBUTES + assert not set(row["SpanAttributes"]) & {"gen_ai.prompt", "gen_ai.completion"} assert rows_by_name["ChatOpenAI"]["SpanAttributes"]["langsmith.span.kind"] == "llm" @@ -217,12 +225,40 @@ def test_content_type_defaults_to_protobuf(): assert len(decode_otlp(_fixture_protobuf(), None)) == 6 -@pytest.mark.parametrize("content_encoding", ["gzip", None]) -def test_gzip_body_by_header_or_magic_bytes(content_encoding): - rows = decode_otlp(gzip.compress(_fixture_protobuf()), "application/x-protobuf", content_encoding) +def test_gzip_body_by_header(): + rows = decode_otlp(gzip.compress(_fixture_protobuf()), "application/x-protobuf", "gzip") assert len(rows) == 6 +def test_gzip_requires_content_encoding_header(): + with pytest.raises(decode.InvalidOTLPPayloadError): + decode_otlp(gzip.compress(_fixture_protobuf()), "application/x-protobuf") + + +def test_invalid_gzip_body_is_rejected(): + with pytest.raises(decode.InvalidOTLPPayloadError): + decode_otlp(b"not gzip", "application/x-protobuf", "gzip") + + +def test_gzip_expansion_respects_body_limit(): + with patch.object(decode, "OTLP_MAX_BODY_BYTES", 1024): + with pytest.raises(decode.OTLPPayloadTooLargeError): + decode_otlp(gzip.compress(b" " * 16384), "application/json", "gzip") + + +def test_concatenated_gzip_members_are_decoded(): + body = _fixture_json() + midpoint = len(body) // 2 + compressed = gzip.compress(body[:midpoint]) + gzip.compress(body[midpoint:]) + assert len(decode_otlp(compressed, "application/json", "gzip")) == 6 + + +@pytest.mark.parametrize("encoding", ["br", "gzip, identity"]) +def test_unsupported_content_encoding_is_rejected(encoding): + with pytest.raises(decode.InvalidOTLPPayloadError): + decode_otlp(_fixture_protobuf(), "application/x-protobuf", encoding) + + def test_long_values_are_truncated_with_marker(): with patch.object(decode, "OTLP_MAX_ATTRIBUTE_VALUE_BYTES", 100): rows = {r["SpanName"]: r for r in decode_otlp(_fixture_json(), "application/json")} @@ -402,7 +438,7 @@ def test_non_string_attribute_values_are_stringified(): assert row["SpanAttributes"]["flag"] == "true" assert row["SpanAttributes"]["ratio"] == "0.5" assert row["SpanAttributes"]["raw"] == "abc" - assert json.loads(row["SpanAttributes"]["list"]) == ["a", "1"] + assert json.loads(row["SpanAttributes"]["list"]) == ["a", 1] # ---------------------------------------------------------------- helpers @@ -412,3 +448,53 @@ def test_encode_otlp_response_matches_request_encoding(): assert encode_otlp_response("application/json") == (b"{}", "application/json") assert encode_otlp_response("application/x-protobuf") == (b"", "application/x-protobuf") assert encode_otlp_response(None) == (b"", "application/x-protobuf") + body, media_type = encode_otlp_response("application/x-protobuf", "invalid trace") + assert media_type == "application/x-protobuf" + from google.rpc.status_pb2 import Status + + assert Status.FromString(body).message == "invalid trace" + + +@pytest.mark.parametrize( + "attributes, expected", + [ + ({"langsmith__span__kind": "llm"}, "llm"), + ({"langsmith__span__kind": "tool"}, "tool"), + ({"gen_ai__operation__name": "chat"}, "llm"), + ({"gen_ai__operation__name": "execute_tool"}, "tool"), + ({"openinference__span__kind": "LLM"}, "llm"), + ], +) +def test_explicit_root_span_semantics_and_response_id_are_preserved(attributes, expected): + exported = _span("root", b"\x01" * 8, gen_ai__response__id="response-123", **attributes) + (row,) = decode_otlp(_export(exported)) + assert (row["ObservationType"], row["LiteLLMRequestId"]) == (expected, "response-123") + + +@pytest.mark.parametrize( + "payload", + [ + '{"messages": 7}', + '{"messages": {"0": "wrong"}}', + '{"messages": [{"kwargs": []}]}', + '{"messages": [{"role": "assistant", "tool_calls": [1]}]}', + ], +) +def test_malformed_framework_messages_preserve_raw_content_without_rejecting_the_batch(payload): + exported = _span("agent", b"\x01" * 8, langsmith__span__kind="chain", gen_ai__prompt=payload) + (row,) = decode_otlp(_export(exported)) + assert row["Input"] == payload + + +def test_unrecognized_heavy_attributes_are_retained(): + exported = _span("root", b"\x01" * 8, gen_ai__prompt="unknown convention", gen_ai__tool__definitions="tools") + (row,) = decode_otlp(_export(exported)) + assert row["SpanAttributes"]["gen_ai.prompt"] == "unknown convention" + assert row["SpanAttributes"]["gen_ai.tool.definitions"] == "tools" + + +@pytest.mark.parametrize("count", [-1, 1 << 32]) +def test_token_counts_outside_storage_range_are_rejected(count): + exported = _span("root", b"\x01" * 8, gen_ai__usage__input_tokens=count) + with pytest.raises(decode.InvalidOTLPPayloadError, match="storage range"): + decode_otlp(_export(exported)) diff --git a/tests/test_litellm/tracing/test_receiver.py b/tests/test_litellm/tracing/test_receiver.py index d492844db79..0d9aa8d034d 100644 --- a/tests/test_litellm/tracing/test_receiver.py +++ b/tests/test_litellm/tracing/test_receiver.py @@ -2,7 +2,10 @@ Tests for TraceReceiver.ingest (litellm/tracing/receiver.py) with a fake store. """ +import asyncio +from collections.abc import AsyncIterator from pathlib import Path +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -89,18 +92,6 @@ async def test_ingest_rejects_oversized_body(): store.insert_spans.assert_not_awaited() -@pytest.mark.asyncio -async def test_large_body_is_decoded_off_the_event_loop(): - store = _fake_store() - with ( - patch.object(receiver_module, "OTLP_OFFLOAD_DECODE_BYTES", 0), - patch.object(receiver_module.asyncio, "to_thread", wraps=receiver_module.asyncio.to_thread) as to_thread, - ): - count = await TraceReceiver(store).ingest(FIXTURE.read_bytes(), "application/json", None, TENANT) - assert count == 6 - to_thread.assert_called_once() - - @pytest.mark.asyncio async def test_empty_export_writes_nothing(): store = _fake_store() @@ -115,3 +106,57 @@ async def test_reads_delegate_to_store(): scope: TraceScope = {"team_ids": ("team-research",), "api_key_hash": ""} assert await tracing.get_trace("t1", scope) is None store.get_trace.assert_awaited_once_with("t1", scope, "") + + +@pytest.mark.asyncio +async def test_cancelled_request_keeps_its_worker_slot_until_decode_finishes(): + import asyncio + import threading + + from litellm.tracing.receiver import TracingOverloadedError + + loop = asyncio.get_running_loop() + owner = threading.get_ident() + started = asyncio.Event() + stored = asyncio.Event() + release = threading.Event() + + def decoder(body, content_type, content_encoding): + assert threading.get_ident() != owner + loop.call_soon_threadsafe(started.set) + assert release.wait(5) + return () + + store = _fake_store() + store.insert_spans.side_effect = lambda _: stored.set() + tracing = TraceReceiver(store, max_concurrent_ingests=1, decoder=decoder) + pending = asyncio.create_task(tracing.ingest(b"small gzip", None, "gzip", TENANT)) + try: + await asyncio.wait_for(started.wait(), 5) + pending.cancel() + with pytest.raises(asyncio.CancelledError): + await pending + with pytest.raises(TracingOverloadedError): + await tracing.ingest(b"", None, None, TENANT) + finally: + release.set() + await asyncio.wait_for(stored.wait(), 5) + await asyncio.sleep(0) + assert await tracing.ingest(b"", None, None, TENANT) == 0 + + +@pytest.mark.asyncio +async def test_expired_upload_releases_ingestion_slot_without_writing() -> None: + from litellm.tracing.receiver import TracingOverloadedError + + async def unfinished_body() -> AsyncIterator[bytes]: + await asyncio.Event().wait() + yield b"" + + store: Final = _fake_store() + receiver: Final = TraceReceiver(store, max_concurrent_ingests=1, body_read_timeout=0) + with pytest.raises(TracingOverloadedError, match="upload timed out"): + await receiver.ingest(unfinished_body(), "application/json", None, TENANT) + store.insert_spans.assert_not_awaited() + assert await receiver.ingest(b"{}", "application/json", None, TENANT) == 0 + store.insert_spans.assert_awaited_once_with(()) diff --git a/tests/test_litellm/tracing/test_store.py b/tests/test_litellm/tracing/test_store.py index 30d7a9b5b0b..3f43e42842c 100644 --- a/tests/test_litellm/tracing/test_store.py +++ b/tests/test_litellm/tracing/test_store.py @@ -436,3 +436,47 @@ async def test_ambiguous_cache_response_id_keeps_cost_unavailable(): assert trace is not None assert trace["summary"]["spend"] is None assert trace["spans"][0]["spend"] is None + + +@pytest.mark.asyncio +async def test_diagnostic_continuation_preserves_content_version_scope_and_unicode_offset(): + from hashlib import sha256 + + message = "first 🧪\nlast" + version = sha256(message.encode()).hexdigest().upper() + client = MagicMock() + client.query = AsyncMock( + side_effect=[ + [{"span_id": "span-1", "message": "first 🧪", "total_chars": len(message), "version": version}], + [{"span_id": "span-1", "message": "\nlast", "total_chars": len(message), "version": version}], + ] + ) + store = TraceStore(client) + scope = {"team_ids": ("team-a",), "api_key_hash": "key-a"} + first = await store.get_span_error("trace-1", "span-1", scope, "scoped-run") + assert first is not None and first["next_cursor"] is not None + last = await store.get_span_error("trace-1", "span-1", scope, "scoped-run", first["next_cursor"]) + assert last is not None + assert first["message"] + last["message"] == message + assert last["next_cursor"] is None + client.query.assert_awaited_with( + "span_error", + { + **scope, + "trace_id": "trace-1", + "span_id": "span-1", + "trace_ref": "scoped-run", + "error_offset": len(first["message"]), + "error_version": version, + }, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("cursor", ["garbage", "e30=", "WzEsMl0="]) +async def test_malformed_diagnostic_cursor_never_reaches_storage(cursor): + client = MagicMock() + client.query = AsyncMock() + with pytest.raises(ValueError, match="Invalid diagnostic cursor"): + await TraceStore(client).get_span_error("trace", "span", {"team_ids": (), "api_key_hash": ""}, cursor=cursor) + client.query.assert_not_awaited() diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index fc750d88e42..e6492c9bca6 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -2,12 +2,17 @@ import base64 import gzip import json import time +from types import MappingProxyType from typing import Final from urllib.parse import parse_qs, urlsplit import pytest -from litellm.rust_bridge._native import NativeTraceStorage +from litellm.rust_bridge._native import NativeTraceStorage, trace_decode_otlp +from litellm.rust_bridge.traces import ClickHouseStorage +from litellm.tracing import Tenant, TraceReceiver, TracingPayloadTooLargeError +from litellm.tracing.decode import decode_otlp +from litellm.tracing.store import TraceStore from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec pytestmark = pytest.mark.requires_rust_extension @@ -61,7 +66,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")) @@ -73,9 +80,10 @@ 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 @@ -93,5 +101,100 @@ async def test_insert_encodes_and_sends_rows(recording_server: RecordingServer) "Timestamp": "1970-01-01T00:00:01.23456789Z", "EngineReceivedMs": row["EngineReceivedMs"], } - 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" + + +def _resource_export(attribute_bytes: int, span_count: int, groups: int = 1) -> bytes: + span: Final = { + "traceId": "01" * 16, + "spanId": "02" * 8, + "name": "shared-resource", + "startTimeUnixNano": "1", + "endTimeUnixNano": "2", + } + resource: Final = { + "resource": { + "attributes": [ + {"key": "shared", "value": {"stringValue": "x" * attribute_bytes}}, + {"key": "litellm.team_id", "value": {"stringValue": "spoofed"}}, + ] + }, + "scopeSpans": [ + { + "scope": {"name": "scope-" * 32, "version": "v" * 128}, + "spans": [{**span, "spanId": f"{index + 1:016x}"} for index in range(span_count)], + } + ], + } + return json.dumps({"resourceSpans": [resource] * groups}).encode() + + +def test_decode_and_tenant_stamping_share_resources_without_crossing_groups() -> None: + body: Final = _resource_export(128, 2, 2) + native: Final = trace_decode_otlp(body, "application/json") + assert native[0]["scope_name"] is native[1]["scope_name"] + assert native[0]["scope_version"] is native[1]["scope_version"] + assert native[0]["resource_attributes"] is native[1]["resource_attributes"] + assert native[2]["resource_attributes"] is native[3]["resource_attributes"] + assert native[0]["resource_attributes"] is not native[2]["resource_attributes"] + rows: Final = decode_otlp(body, "application/json") + first: Final = Tenant("team-a", "key-a", "org-a").stamp_rows(rows) + second: Final = Tenant("team-b", "key-b", "org-b").stamp_rows(rows) + assert first[0]["ResourceAttributes"] is first[1]["ResourceAttributes"] + assert first[2]["ResourceAttributes"] is first[3]["ResourceAttributes"] + assert first[0]["ResourceAttributes"] is not first[2]["ResourceAttributes"] + assert first[0]["ResourceAttributes"] is not second[0]["ResourceAttributes"] + assert first[0]["ResourceAttributes"] == { + "shared": "x" * 128, + "litellm.team_id": "team-a", + "litellm.api_key_hash": "key-a", + "litellm.org_id": "org-a", + } + assert second[0]["ResourceAttributes"]["litellm.team_id"] == "team-b" + assert rows[0]["ResourceAttributes"] == {"shared": "x" * 128, "litellm.team_id": "spoofed"} + + +@pytest.mark.asyncio +async def test_resource_fanout_reaches_insert_with_identical_values(recording_server: RecordingServer) -> None: + body: Final = _resource_export(16 * 1024, 1024) + receiver: Final = TraceReceiver(TraceStore(ClickHouseStorage("trace_test", recording_server.base_url))) + tenant: Final = Tenant("team-a", "key-a", "org-a") + assert await receiver.ingest(body, "application/json", None, tenant) == 1024 + encoded: Final = gzip.decompress(recording_server.requests[0].raw_body) + actual: Final = tuple(json.loads(line) for line in encoded.splitlines()) + expected: Final = tenant.stamp_rows(decode_otlp(body, "application/json")) + assert len(encoded) < 64 * 1024 * 1024 + assert tuple({key: value for key, value in row.items() if key != "EngineReceivedMs"} for row in actual) == tuple( + {**row, "Timestamp": "1970-01-01T00:00:00.000000001Z"} for row in expected + ) + assert len({row["EngineReceivedMs"] for row in actual}) == 1 + + +@pytest.mark.asyncio +async def test_shared_resource_still_hits_insert_limit_before_transport(recording_server: RecordingServer) -> None: + recording_server.expected_requests = 0 + body: Final = _resource_export(64 * 1024, 1024) + receiver: Final = TraceReceiver(TraceStore(ClickHouseStorage("trace_test", recording_server.base_url))) + with pytest.raises(TracingPayloadTooLargeError, match="encoded size limit"): + await receiver.ingest(body, "application/json", None, Tenant("team-a", "key-a")) + assert recording_server.requests == [] + + +@pytest.mark.asyncio +async def test_insert_validates_values_without_pydantic_copy(recording_server: RecordingServer) -> None: + storage: Final = ClickHouseStorage("trace_test", recording_server.base_url) + invalid: Final = object() + with pytest.raises(ValueError, match=type(invalid).__name__): + await storage.insert_rows("otel_traces", [{"ResourceAttributes": invalid}]) + attributes: Final = MappingProxyType({"service.name": "trace-test"}) + await storage.insert_rows( + "otel_traces", + (MappingProxyType({"Timestamp": 1, "ResourceAttributes": attributes, "SpanAttributes": attributes}),), + ) + stored: Final = json.loads(gzip.decompress(recording_server.requests[0].raw_body)) + assert stored["Timestamp"] == "1970-01-01T00:00:00.000000001Z" + assert stored["ResourceAttributes"] == attributes + assert stored["SpanAttributes"] == attributes diff --git a/tests/unit/integrations/open_telemetry/test_otel_exception_handler.py b/tests/unit/integrations/open_telemetry/test_otel_exception_handler.py index dc99df24c50..56d067b16cf 100644 --- a/tests/unit/integrations/open_telemetry/test_otel_exception_handler.py +++ b/tests/unit/integrations/open_telemetry/test_otel_exception_handler.py @@ -3,10 +3,9 @@ that fail after auth but before the route handler runs (e.g. /model/new TypeError or RequestValidationError).""" import asyncio -import types import pytest -from fastapi import HTTPException +from fastapi import HTTPException, Request from fastapi.exceptions import RequestValidationError import litellm.proxy.proxy_server as proxy_server_module @@ -23,13 +22,11 @@ from litellm.integrations._types.open_inference import ErrorAttributes from ._helpers import assert_server_span_attrs, get_server_span -def _fake_request(parent_otel_span=None, path="/key/generate"): - """A real Request always carries a url; the validation handler reads its path to - decide whether the caller is on a surface with its own error contract.""" - state = types.SimpleNamespace() - if parent_otel_span is not None: - state.parent_otel_span = parent_otel_span - return types.SimpleNamespace(state=state, url=types.SimpleNamespace(path=path)) +def _fake_request(parent_otel_span: object | None = None, path: str = "/key/generate") -> Request: + return Request({ + "type": "http", "method": "POST", "path": path, "headers": [], + "state": {"parent_otel_span": parent_otel_span}, + }) @pytest.fixture diff --git a/tests/unit/proxy/auth/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py index 08c2f02a83c..1cfef5b3a6c 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -124,7 +124,7 @@ async def test_check_blocked_team(): setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "prisma_client", "hello-world") - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") await user_api_key_auth(request=request, api_key="Bearer " + user_key) @@ -162,7 +162,7 @@ async def test_team_object_has_object_permission_id(): setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "prisma_client", "test-client") - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") with patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock) as mock_common_checks: @@ -263,7 +263,7 @@ async def test_aaauser_personal_budgets(key_ownership): setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "prisma_client", _NoMembershipRowPrisma()) - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") test_user_cache = getattr(litellm.proxy.proxy_server, "user_api_key_cache") @@ -294,7 +294,7 @@ async def test_user_api_key_auth_fails_with_prohibited_params(prohibited_param): setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") # Create request with prohibited parameter in body - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") async def return_body(): @@ -334,7 +334,7 @@ async def test_auth_with_allowed_routes(route, should_raise_error): setattr(proxy_server, "master_key", "sk-1234") setattr(proxy_server, "general_settings", general_settings) - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": route, "headers": []}) request._url = URL(url=route) if should_raise_error: @@ -411,7 +411,7 @@ def test_ui_token_route_access(route, user_role, should_be_allowed): from starlette.datastructures import URL from fastapi import Request - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": route, "headers": []}) request._url = URL(url=route) if should_be_allowed: @@ -494,7 +494,7 @@ async def test_auth_not_connected_to_db(): {"allow_requests_on_db_unavailable": True}, ) - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") valid_token = await user_api_key_auth(request=request, api_key="Bearer " + user_key) @@ -676,7 +676,7 @@ async def test_soft_budget_alert(): setattr(litellm.proxy.proxy_server, "prisma_client", AsyncMock()) # Create request - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") # Track if budget_alerts was called @@ -1162,7 +1162,7 @@ async def test_x_litellm_api_key(): ignored_key = "aj12445" # Create request with headers as bytes - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") valid_token = await user_api_key_auth( @@ -1336,7 +1336,7 @@ async def test_user_model_budget_is_enforced_through_user_api_key_auth(over_budg ttl=600, ) - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") async def return_body(): diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index ef6832ef77b..781d0a13bfd 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -6267,6 +6267,7 @@ async def test_user_api_key_auth_sets_end_user_id_when_builder_skips_it(): "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") @@ -6321,6 +6322,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder() "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") @@ -6376,6 +6378,7 @@ async def test_user_api_key_auth_authenticates_before_raising_malformed_body_err "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") @@ -6435,6 +6438,7 @@ async def _run_auth_with_malformed_body(post_call_failure_hook): "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") @@ -6507,6 +6511,7 @@ async def test_user_api_key_auth_malformed_body_with_rejected_key_still_returns_ "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") @@ -6557,6 +6562,7 @@ async def test_user_api_key_auth_does_not_double_log_a_malformed_body_from_a_rej "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") diff --git a/tests/unit/proxy/common_utils/test_http_parsing_utils.py b/tests/unit/proxy/common_utils/test_http_parsing_utils.py index e3851f6c21a..fd747d5a6f2 100644 --- a/tests/unit/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/unit/proxy/common_utils/test_http_parsing_utils.py @@ -2,14 +2,14 @@ import gzip import io import json from collections.abc import Mapping -from typing import Literal, get_type_hints +from typing import Final, Literal, get_type_hints from unittest.mock import AsyncMock, MagicMock, patch import orjson import pytest -from fastapi import Request from fastapi.testclient import TestClient from starlette.datastructures import FormData +from starlette.requests import Request @@ -1109,7 +1109,7 @@ class TestGetRequestBody: mock_request.method = "POST" mock_request.body = AsyncMock(return_value=orjson.dumps(payload)) mock_request.headers = {"content-type": "application/json; charset=utf-8"} - mock_request.scope = {} + mock_request.scope = {"type": "http", "method": "POST", "path": "/v1/chat/completions"} result = await get_request_body(mock_request) assert result == payload @@ -1120,7 +1120,7 @@ class TestGetRequestBody: mock_request.method = "POST" mock_request.headers = {"content-type": "multipart/form-data; boundary=x"} mock_request.form = AsyncMock(return_value=FormData({"k": "v"})) - mock_request.scope = {} + mock_request.scope = {"type": "http", "method": "POST", "path": "/v1/chat/completions"} result = await get_request_body(mock_request) assert result == {"k": "v"} @@ -1273,3 +1273,96 @@ def test_shared_inference_model_selection_preserves_handler_precedence( from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model assert resolve_inference_model(body, settings, cli, path, kind=kind) == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "method,path,skip_parse", + [ + ("POST", "/v1/traces", True), + ("GET", "/v1/traces", False), + ("POST", "/v1/messages", False), + ("POST", "/v1/traces/other", False), + ], +) +@pytest.mark.parametrize("root_path", ["", "/tenant-a"]) +async def test_only_trace_ingest_skips_json_body(method: str, path: str, skip_parse: bool, root_path: str) -> None: + body: Final = b'{"key":"value"}' + receive: Final = AsyncMock(return_value={"type": "http.request", "body": body, "more_body": False}) + request: Final = Request( + { + "type": "http", "method": method, "path": root_path + path, "root_path": root_path, + "headers": [(b"content-type", b"application/json")], + }, + receive, + ) + + parsed: Final = await _read_request_body(request) + if skip_parse: + assert parsed == {} + receive.assert_not_awaited() + else: + assert parsed == {"key": "value"} + receive.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("content_type, encoding", [ + ("application/json", ""), ("application/x-protobuf", ""), ("application/json", "gzip"), +]) +async def test_otlp_auth_does_not_consume_chunked_bodies_before_the_receiver_limit(content_type, encoding): + from litellm.constants import OTLP_MAX_BODY_BYTES + from litellm.tracing import Tenant, TraceReceiver, TracingPayloadTooLargeError + + received = [] + chunk = b"x" * (OTLP_MAX_BODY_BYTES // 2 + 1) + + async def receive(): + received.append(1) + assert len(received) <= 2, "receiver must reject without consuming subsequent chunks" + return {"type": "http.request", "body": chunk, "more_body": True} + + request = Request({"type": "http", "method": "POST", "path": "/v1/traces", "headers": [ + (b"content-type", content_type.encode()), (b"content-encoding", encoding.encode()), + ]}, receive) + assert await _read_request_body(request) == {} + assert received == [] + store = MagicMock() + store.insert_spans = AsyncMock() + with pytest.raises(TracingPayloadTooLargeError): + await TraceReceiver(store).ingest(request.stream(), content_type, encoding, Tenant("team", "key")) + assert len(received) == 2 + store.insert_spans.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_auth_body_read_and_trace_handler_leave_stream_for_receiver_limit() -> None: + from litellm.constants import OTLP_MAX_BODY_BYTES + from litellm.proxy import tracing_endpoints + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import _read_request_body_deferring_parse_failure + from litellm.tracing import TraceReceiver + + chunk: Final = b"x" * (OTLP_MAX_BODY_BYTES // 2 + 1) + receive: Final = AsyncMock( + side_effect=[{"type": "http.request", "body": chunk, "more_body": True}] * 2 + ) + request: Final = Request( + {"type": "http", "method": "POST", "path": "/v1/traces", "headers": [(b"content-type", b"application/json")]}, + receive, + ) + store: Final = MagicMock() + store.insert_spans = AsyncMock() + context: Final = await tracing_endpoints.provide_trace_access( + auth=UserAPIKeyAuth(token="key", team_id="team"), tracing=TraceReceiver(store) + ) + + parsed, parse_error = await _read_request_body_deferring_parse_failure(request) + assert parsed == {} + assert parse_error is None + receive.assert_not_awaited() + + response: Final = await tracing_endpoints.ingest_otlp_traces(request, context) + assert response.status_code == 413 + assert receive.await_count == 2 + store.insert_spans.assert_not_awaited() diff --git a/tests/unit/proxy/proxy_server/test_exception_handlers.py b/tests/unit/proxy/proxy_server/test_exception_handlers.py index 16cb1146ff5..0aff43057f9 100644 --- a/tests/unit/proxy/proxy_server/test_exception_handlers.py +++ b/tests/unit/proxy/proxy_server/test_exception_handlers.py @@ -16,7 +16,7 @@ from unittest.mock import MagicMock import httpx import pytest -from fastapi import HTTPException +from fastapi import HTTPException, Request from fastapi.exceptions import RequestValidationError from litellm.proxy._types import ProxyException @@ -31,10 +31,10 @@ from .conftest import normalize def _make_request(parent_otel_span=None, path="/chat/completions"): - """A real Request always carries a url; the validation handler reads its path to - decide whether the caller is on a surface with its own error contract.""" - state = SimpleNamespace(parent_otel_span=parent_otel_span) - return SimpleNamespace(state=state, url=SimpleNamespace(path=path)) + return Request({ + "type": "http", "method": "POST", "path": path, "headers": [], + "state": {"parent_otel_span": parent_otel_span}, + }) # --------------------------------------------------------------------------- @@ -477,3 +477,42 @@ async def test_otel_unhandled_exception_handler_reraises_http_exception_invalid( request = _make_request() with pytest.raises(HTTPException): await otel_unhandled_exception_handler(request=request, exc=HTTPException(status_code=418, detail="teapot")) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("media_type", ["application/json", "application/x-protobuf"]) +@pytest.mark.parametrize("root_path", ["", "/tenant-a"]) +@pytest.mark.parametrize("native_available", [True, False]) +@pytest.mark.parametrize("error", [ + ProxyException("database credentials: secret", "auth_error", None, 401), + HTTPException(403, "database credentials: secret"), +]) +async def test_otlp_auth_errors_hide_internal_details_and_survive_missing_native( + media_type: str, root_path: str, native_available: bool, + error: ProxyException | HTTPException, monkeypatch: pytest.MonkeyPatch, +) -> None: + from google.rpc.status_pb2 import Status + + from litellm.proxy.proxy_server import otlp_http_exception_handler + from litellm.rust_bridge import loader + + if not native_available: + monkeypatch.setattr(loader, "_cached_bridge", None) + request: Final = Request({ + "type": "http", "method": "POST", "path": root_path + "/v1/traces", "root_path": root_path, + "headers": [(b"content-type", media_type.encode())], + }) + response: Final = ( + await openai_exception_handler(request, error) + if isinstance(error, ProxyException) + else await otlp_http_exception_handler(request, error) + ) + assert response.status_code == (401 if isinstance(error, ProxyException) else 403) + assert response.headers["content-type"].startswith(media_type) + message: Final = ( + json.loads(response.body)["message"] + if media_type == "application/json" + else Status.FromString(response.body).message + ) + expected: Final = "Unauthorized" if isinstance(error, ProxyException) else "Forbidden" + assert message == (expected if native_available or media_type == "application/json" else "") diff --git a/tests/unit/proxy/test_proxy_reject_logging.py b/tests/unit/proxy/test_proxy_reject_logging.py index eb5c5a52f0a..d5a3acb2cd7 100644 --- a/tests/unit/proxy/test_proxy_reject_logging.py +++ b/tests/unit/proxy/test_proxy_reject_logging.py @@ -152,6 +152,7 @@ async def test_chat_completion_request_with_redaction(route, body): scope={ "type": "http", "method": "POST", + "path": route, "headers": [(b"content-type", b"application/json")], "query_string": query_params.encode(), } diff --git a/tests/unit/proxy/test_proxy_server.py b/tests/unit/proxy/test_proxy_server.py index 8947da4d9fc..300edc8e435 100644 --- a/tests/unit/proxy/test_proxy_server.py +++ b/tests/unit/proxy/test_proxy_server.py @@ -263,7 +263,7 @@ def test_add_headers_to_request(litellm_key_header_name): "X-Stainless-Header": "Stainless-Value", "anthropic-beta": "beta-value", } - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") request._body = json.dumps({"model": "gpt-3.5-turbo"}).encode("utf-8") request_headers = clean_headers(headers, litellm_key_header_name) @@ -466,7 +466,7 @@ async def test_team_disable_guardrails(mock_acompletion, client_no_auth, monkeyp setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "prisma_client", "hello-world") - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") body = {"metadata": {"guardrails": {"hide_secrets": False}}} @@ -1347,7 +1347,7 @@ async def test_create_team_member_add_team_admin_user_api_key_auth( from starlette.datastructures import URL - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": team_route, "headers": []}) request._url = URL(url=team_route) body = {} diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index 2e34172acfd..aa1403b8db9 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -96,8 +96,22 @@ def client() -> TestClient: return TestClient(app) -def test_501_when_tracing_not_enabled(client): - assert client.post("/v1/traces", content=b"").status_code == 501 +@pytest.mark.parametrize("native_available", [True, False]) +def test_501_when_tracing_not_enabled( + client: TestClient, native_available: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + from google.rpc.status_pb2 import Status + + from litellm.rust_bridge import loader + + if not native_available: + monkeypatch.setattr(loader, "_cached_bridge", None) + response: Final = client.post("/v1/traces", content=b"") + assert response.status_code == 501 + assert response.headers["content-type"] == "application/x-protobuf" + assert Status.FromString(response.content).message == ( + "Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL." if native_available else "" + ) assert client.get("/v1/traces").status_code == 501 @@ -111,7 +125,7 @@ def test_post_protobuf_returns_empty_protobuf(client, receiver): assert response.content == b"" assert response.headers["content-type"] == "application/x-protobuf" kwargs = receiver.ingest.call_args.kwargs - assert kwargs["body"] == b"\x0a\x00" + assert kwargs["body"] is not None assert kwargs["content_type"] == "application/x-protobuf" assert kwargs["content_encoding"] == "gzip" assert kwargs["tenant"].team_id == "team-research" @@ -134,7 +148,9 @@ def test_post_too_large_is_413(client, receiver): receiver.ingest.side_effect = TracingPayloadTooLargeError("OTLP body exceeds 10 bytes") response = client.post("/v1/traces", content=b"x" * 20) assert response.status_code == 413 - assert "exceeds" in response.json()["detail"] + from google.rpc.status_pb2 import Status + + assert "exceeds" in Status.FromString(response.content).message def test_list_traces_passes_scope_window_and_cursor(client, receiver): @@ -222,8 +238,13 @@ def test_view_only_admin_cannot_ingest_traces(client, receiver): receiver.ingest.assert_not_called() -@pytest.mark.parametrize("status_code", [401, 403]) -def test_auth_failure_precedes_disabled_receiver(client: TestClient, status_code: int) -> None: +@pytest.mark.parametrize( + "status_code, field, message", + [(401, "detail", "Invalid API key"), (403, "message", "Not allowed to ingest agent traces")], +) +def test_auth_failure_precedes_disabled_receiver( + client: TestClient, status_code: int, field: str, message: str +) -> None: def unavailable() -> None: return None @@ -234,11 +255,9 @@ def test_auth_failure_precedes_disabled_receiver(client: TestClient, status_code client.app.dependency_overrides[user_api_key_auth] = authenticate client.app.dependency_overrides[tracing_endpoints.provide_receiver] = unavailable - response: Final = client.post("/v1/traces", content=b"{}") + response: Final = client.post("/v1/traces", content=b"{}", headers={"content-type": "application/json"}) assert response.status_code == status_code - assert response.json() == { - "detail": "Invalid API key" if status_code == 401 else "Not allowed to ingest agent traces" - } + assert response.json() == {field: message} def test_disabled_receiver_precedes_read_scope_rejection(client: TestClient) -> None: diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index d12d0a3219c..97d1782b08c 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -115,7 +115,7 @@ import type { ComplexityRouterConfigPayload } from "./add_model/build_complexity import type { AutoRouterPresetsResponse } from "@/lib/autorouter_presets"; import type { VectorStoreIndex } from "@/app/(dashboard)/vector-stores/_components/IndexesTab"; import type { RoutingDecision } from "./view_logs/LogDetailsDrawer/RoutingDecisionCard"; -import type { SpanDetail, Trace, TracePage } from "./view_logs/TraceView/traceTypes"; +import type { SpanDetail, SpanErrorPage, Trace, TracePage } from "./view_logs/TraceView/traceTypes"; import { createApiClient, deriveErrorMessage, @@ -2139,6 +2139,17 @@ export const agentTraceSpanCall = async ( query: { trace_ref: traceRef || undefined }, }); +export const agentTraceSpanErrorCall = async ( + accessToken: string, + traceId: string, + spanId: string, + options: { traceRef?: string; cursor?: string | null }, +): Promise => + apiClient.get(`/v1/traces/${encodeURIComponent(traceId)}/spans/${encodeURIComponent(spanId)}/error`, { + accessToken, + query: { trace_ref: options.traceRef || undefined, cursor: options.cursor || undefined }, + }); + export const adminSpendLogsCall = async (accessToken: string) => { try { const data = await apiClient.get(`/global/spend/logs`, { accessToken }); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx index 9c156d88bd5..86ec3007f72 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx @@ -1,15 +1,17 @@ "use client"; import { useQuery, type UseQueryOptions } from "@tanstack/react-query"; +import { useState } from "react"; import { AlertTriangle } from "lucide-react"; +import { Button } from "@/components/ui/button"; import { cn } from "@/lib/cva.config"; -import { agentTraceSpanCall } from "../../networking"; +import { agentTraceSpanCall, agentTraceSpanErrorCall } from "../../networking"; import { type KeyValue, KeyValueRows, objectEntries } from "./KeyValueRows"; import { Card, MessageCard, Section, ToolResultCard } from "./MessageCard"; import type { ErrorSource } from "./traceTree"; -import type { Span, SpanDetail, TraceMessage, UIContent, UIMessage } from "./traceTypes"; +import type { Span, SpanDetail, SpanErrorPage, TraceMessage, UIContent, UIMessage } from "./traceTypes"; import { errorSource, parseJson, parseMessages, prettyPayload } from "./traceUtils"; const ERROR_SOURCE_LABEL: Record = { tool: "Tool", model: "Model", litellm: "LiteLLM" }; @@ -143,6 +145,58 @@ interface DetailContentProps { span: Span; } +function DiagnosticContent({ accessToken, traceId, traceRef, span }: DetailContentProps) { + const [opened, setOpened] = useState(false); + const [cursor, setCursor] = useState(null); + const queryOptions: UseQueryOptions = { + queryKey: ["agentTraceSpanError", traceId, traceRef, span.span_id, accessToken, cursor], + queryFn: () => agentTraceSpanErrorCall(accessToken, traceId, span.span_id, { traceRef, cursor }), + enabled: opened, + staleTime: Infinity, + gcTime: 0, + retry: false, + }; + const query = useQuery(queryOptions); + return ( +
+ {span.error_truncated &&

Error preview truncated

} + {!opened && ( + + )} + {opened && query.isPending &&

Loading diagnostic…

} + {opened && query.isError && ( +
+ Could not load diagnostic: {query.error.message} + +
+ )} + {opened && query.data && ( + <> + +

+ {cursor ? "Continuation" : "Beginning"} of stored diagnostic ({query.data.total_chars.toLocaleString()}{" "} + characters) +

+ {query.data.next_cursor && ( + + )} + {cursor && ( + + )} + + )} +
+ ); +} + /** Content tab: the error first (if any), then collapsible Input and Output rendered as chat cards. */ export function DetailContent({ accessToken, traceId, traceRef, span }: DetailContentProps) { const detailQuery = useSpanDetail(accessToken, traceId, span.span_id, traceRef); @@ -152,6 +206,15 @@ export function DetailContent({ accessToken, traceId, traceRef, span }: DetailCo return (
+ {span.error && ( + + )} {detailQuery.isLoading &&
Loading span…
} {detailQuery.isError &&
Could not load span: {detailQuery.error.message}
} {detail?.input ? ( diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.integration.test.tsx similarity index 89% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.test.tsx rename to ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.integration.test.tsx index 1adcda0f6cc..be5eefa1546 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.integration.test.tsx @@ -6,14 +6,15 @@ import { renderWithProviders, testQueryClient } from "../../../../tests/test-uti import { DetailPane } from "./DetailPane"; import { absoluteTime, SpanHoverCard, spanFacts } from "./SpanHoverCard"; import type { GroupRowData, SpanRowData } from "./traceTree"; -import type { Span, SpanDetail, Trace } from "./traceTypes"; +import type { Span, SpanDetail, SpanErrorPage, Trace } from "./traceTypes"; vi.mock("../../networking", () => ({ agentTraceSpanCall: vi.fn(), + agentTraceSpanErrorCall: vi.fn(), getProxyBaseUrl: () => "http://proxy.test/", })); -import { agentTraceSpanCall } from "../../networking"; +import { agentTraceSpanCall, agentTraceSpanErrorCall } from "../../networking"; type SpanFields = Partial & Pick; @@ -322,3 +323,31 @@ describe("SpanHoverCard", () => { expect(within(card).getByRole("region", { name: "Tags" })).toHaveTextContent("agent:support_triage_agent"); }); }); + +it("retrieves the retained diagnostic one section at a time", async () => { + const firstPage: SpanErrorPage = { + span_id: "tool1", + message: "First diagnostic section", + total_chars: 100, + next_cursor: "next-section", + }; + const lastPage: SpanErrorPage = { + span_id: "tool1", + message: "Last diagnostic section", + total_chars: 100, + next_cursor: null, + }; + vi.mocked(agentTraceSpanErrorCall).mockResolvedValueOnce(firstPage).mockResolvedValueOnce(lastPage); + renderPane(spanRow({ ...failedTool, error_truncated: true })); + expect(screen.getByText("Error preview truncated")).toBeInTheDocument(); + await userEvent.click(screen.getByRole("button", { name: "View stored diagnostic" })); + expect(await screen.findByText("First diagnostic section")).toBeInTheDocument(); + await userEvent.click(screen.getByRole("button", { name: "Next section" })); + expect(await screen.findByText("Last diagnostic section")).toBeInTheDocument(); + expect(screen.queryByText("First diagnostic section")).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Next section" })).not.toBeInTheDocument(); + expect(agentTraceSpanErrorCall).toHaveBeenLastCalledWith("sk-test", "t1", "tool1", { + traceRef: undefined, + cursor: "next-section", + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts index 4711799bd1a..d3080f8aab4 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts @@ -20,6 +20,7 @@ export interface Span { status: SpanStatus; /** Exception message when status is "error". */ error?: string | null; + error_truncated?: boolean; input_preview: string; model: string | null; input_tokens: number; @@ -120,3 +121,10 @@ export interface TraceMessage { name?: string; tool_calls?: TraceToolCall[]; } + +export interface SpanErrorPage { + span_id: string; + message: string; + total_chars: number; + next_cursor: string | null; +} diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c6bb9be41df..ebc6d0e70cc 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -21779,6 +21779,23 @@ export interface paths { patch?: never; trace?: never; }; + "/v1/traces/{trace_id}/spans/{span_id}/error": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get Agent Trace Span Error */ + get: operations["get_agent_trace_span_error_v1_traces__trace_id__spans__span_id__error_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/v1/unified_access_group": { parameters: { query?: never; @@ -44115,6 +44132,17 @@ export interface components { /** Version */ version?: string; }; + /** SpanErrorPage */ + SpanErrorPage: { + /** Message */ + message: string; + /** Next Cursor */ + next_cursor: string | null; + /** Span Id */ + span_id: string; + /** Total Chars */ + total_chars: number; + }; /** SpendAnalyticsPaginatedResponse */ SpendAnalyticsPaginatedResponse: { metadata?: components["schemas"]["DailySpendMetadata"]; @@ -77335,6 +77363,41 @@ export interface operations { }; }; }; + get_agent_trace_span_error_v1_traces__trace_id__spans__span_id__error_get: { + parameters: { + query?: { + trace_ref?: string; + cursor?: string | null; + }; + header?: never; + path: { + trace_id: string; + span_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["SpanErrorPage"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; list_access_groups_v1_unified_access_group_get: { parameters: { query?: never; From a38fff65601ce67b8ee2c3a38dc14eacdfa646e0 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 1 Oct 2026 13:48:41 -0700 Subject: [PATCH 012/203] fix(proxy): enforce key/team vector_stores allowlist on /v1/rag/query (#43953) * add test case for /rag/query and stronger auth check * style(proxy): ruff format auth_checks.py Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): build rag query vector store ids immutably and test the no-registry path Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): integration coverage for /v1/rag/query vector store allowlist Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): audit cells for /v1/rag/query vector store allowlist Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): consolidate vector store allowlist audit coverage Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): type RAG vector store request body Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Mrinal Chanshetty Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 32 ++- .../test_rag_query_vector_store_allowlist.py | 222 ++++++++++++++++++ ...st_auth_checks_object_access_and_lookup.py | 60 +++++ 3 files changed, 306 insertions(+), 8 deletions(-) create mode 100644 tests/integration/authorization/test_rag_query_vector_store_allowlist.py diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 4b7b8290a30..dbd6f28a183 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -6609,6 +6609,19 @@ def _is_wildcard_pattern(allowed_model_pattern: str) -> bool: return "*" in allowed_model_pattern +def _get_rag_query_vector_store_id(request_body: Mapping[str, object]) -> str | None: + """ + /v1/rag/query carries its vector store in retrieval_config.vector_store_id, + not in vector_store_ids or tools[].vector_store_ids. + """ + retrieval_config: Final = request_body.get("retrieval_config") + if not isinstance(retrieval_config, dict): + return None + + vector_store_id: Final = retrieval_config.get("vector_store_id") + return vector_store_id if isinstance(vector_store_id, str) and vector_store_id else None + + async def vector_store_access_check( request_body: dict, team_object: LiteLLM_TeamTable | None, @@ -6628,13 +6641,16 @@ async def vector_store_access_check( verbose_proxy_logger.debug("Prisma client not found, skipping vector store access check") return True - if litellm.vector_store_registry is None: - verbose_proxy_logger.debug("Vector store registry not found, skipping vector store access check") - return True - - vector_store_ids_to_run: Final = litellm.vector_store_registry.get_vector_store_ids_to_run( - non_default_params=request_body, tools=request_body.get("tools", None) - ) + registry_ids: Final = ( + litellm.vector_store_registry.get_vector_store_ids_to_run( + non_default_params=request_body, tools=request_body.get("tools", None) + ) + if litellm.vector_store_registry is not None + else None + ) or () + rag_vector_store_id: Final = _get_rag_query_vector_store_id(_typed_request_body(request_body)) + rag_ids: Final = (rag_vector_store_id,) if rag_vector_store_id is not None else () + vector_store_ids_to_run: Final = tuple(dict.fromkeys((*registry_ids, *rag_ids))) if not vector_store_ids_to_run: verbose_proxy_logger.debug("Vector store to run not found, skipping vector store access check") return True @@ -6674,7 +6690,7 @@ async def vector_store_access_check( def _can_object_call_vector_stores( object_type: Literal["key", "team", "org"], - vector_store_ids_to_run: list[str], + vector_store_ids_to_run: Sequence[str], object_permissions: _VectorStorePermissionsRow | None, ): """ diff --git a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py new file mode 100644 index 00000000000..896c88c68bb --- /dev/null +++ b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py @@ -0,0 +1,222 @@ +from __future__ import annotations + +import uuid +from collections.abc import Iterator, Mapping +from pathlib import Path +from types import MappingProxyType +from typing import Final, Literal, TypeAlias + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value +from integration._support.process import owned_proxy +from integration.authorization._guardrail_opt_out import upstream_observations +from pydantic import JsonValue + +CONFIG_STORE_ID: Final = "vs_integration_config_store" +PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" +REMOVE_OPENAI_API_BASE: Final = ("OPENAI_API_BASE",) +JsonObject: TypeAlias = dict[str, JsonValue] + + +def _json_array(*values: JsonValue) -> JsonValue: + return [*values] # mutable-ok: request payloads and YAML sequences require list values + + +def _permission_for_stores(*store_ids: str) -> JsonObject: + permission: Final[JsonObject] = {"vector_stores": _json_array(*store_ids)} + return permission + + +def _key_for_scope(scenario: Scenario, model: str, scope: Literal["key", "team"], store_id: str) -> str: + if scope == "key": + return scenario.key(models=_json_array(model), object_permission=_permission_for_stores(store_id)) + team: Final = scenario.team(models=_json_array(model), object_permission=_permission_for_stores(store_id)) + return scenario.key(team_id=team, models=_json_array(model)) + + +def _rag_query_body(model: str, marker: str, store_id: str) -> JsonObject: + body: Final[JsonObject] = { + "model": model, + "messages": _json_array({"role": "user", "content": marker}), + "retrieval_config": {"vector_store_id": store_id, "custom_llm_provider": "openai", "top_k": 1}, + } + return body + + +def _rag_query( + gateway: Gateway, + model: str, + marker: str, + key: str, + *, + store_id: str = CONFIG_STORE_ID, + path: str = "/v1/rag/query", +) -> httpx.Response: + return gateway.request("POST", path, _rag_query_body(model, marker, store_id), key=key) + + +def _searches_for_marker( + gateway: Gateway, marker: str, store_id: str = CONFIG_STORE_ID +) -> tuple[Mapping[str, JsonValue], ...]: + search_path: Final = f"/vector_stores/{store_id}/search" + return tuple( + observation + for observation in upstream_observations(gateway) + if observation["path"] == search_path and marker in str(observation["body"]) + ) + + +def _no_registry_config(directory: Path) -> Path: + config: Final = object_value(yaml.safe_load(PROXY_CONFIG.read_text())) + config_without_registry: Final[Mapping[str, JsonValue]] = MappingProxyType( + {name: value for name, value in config.items() if name != "vector_store_registry"} + ) + yaml_config: Final[JsonObject] = {**config_without_registry, "model_list": _json_array()} + path: Final = directory / "proxy_no_vector_store_registry.yaml" + path.write_text(yaml.safe_dump(yaml_config)) + return path + + +def _openai_environment(gateway: Gateway) -> Mapping[str, str]: + return MappingProxyType({"OPENAI_BASE_URL": gateway.upstream_url, "OPENAI_API_KEY": "synthetic-openai-key"}) + + +@pytest.fixture(scope="module") +def no_registry_gateways(tmp_path_factory: pytest.TempPathFactory) -> Iterator[tuple[Gateway, Gateway]]: + with gateway_from_environment() as upstream_gateway: + directory: Final = tmp_path_factory.mktemp("rag_query_no_registry") + config: Final = _no_registry_config(directory) + with owned_proxy( + upstream_gateway, + directory, + _openai_environment(upstream_gateway), + config=config, + remove_environment=REMOVE_OPENAI_API_BASE, + workers=2, + ) as no_registry_gateway: + yield no_registry_gateway, upstream_gateway + + +@pytest.mark.parametrize( + ("scope", "error_type"), + (("key", "key_vector_store_access_denied"), ("team", "team_vector_store_access_denied")), +) +def test_rag_query_is_denied_when_key_or_team_allowlist_excludes_store( + gateway: Gateway, scope: Literal["key", "team"], error_type: str +) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = _key_for_scope(scenario, model, scope, "vs_some_other_store") + marker: Final = f"lit5610 rag query denied {uuid.uuid4().hex}" + + response: Final = _rag_query(gateway, model, marker, key) + + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == error_type, response.text + assert _searches_for_marker(gateway, marker) == () + + +@pytest.mark.parametrize("scope", ("key", "team")) +def test_rag_query_searches_configured_store_when_allowlist_includes_it( + gateway: Gateway, scope: Literal["key", "team"] +) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = _key_for_scope(scenario, model, scope, CONFIG_STORE_ID) + marker: Final = f"lit5610 rag query allowed {uuid.uuid4().hex}" + + response: Final = _rag_query(gateway, model, marker, key) + assert response.status_code == 200, response.text + + searches: Final = _searches_for_marker(gateway, marker) + assert len(searches) == 1, searches + assert marker in str(object_value(searches[0]["body"])["query"]), searches + + +def test_rag_query_without_key_object_permission_can_search_store(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = scenario.key(models=_json_array(model)) + marker: Final = f"lit5610 rag query no permission {uuid.uuid4().hex}" + + response: Final = _rag_query(gateway, model, marker, key) + assert response.status_code == 200, response.text + + searches: Final = _searches_for_marker(gateway, marker) + assert len(searches) == 1, searches + assert marker in str(object_value(searches[0]["body"])["query"]), searches + + +@pytest.mark.parametrize("scope", ("team", "key")) +def test_no_registry_rag_query_denies_unregistered_store_when_allowlist_excludes( + no_registry_gateways: tuple[Gateway, Gateway], scope: Literal["team", "key"] +) -> None: + no_registry_gateway, upstream_gateway = no_registry_gateways + with no_registry_gateway.scenario() as scenario: + model: Final = scenario.model() + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + key: Final = _key_for_scope(scenario, model, scope, "vs_some_other_store") + marker: Final = f"lit5610 no registry denied {scope} {uuid.uuid4().hex}" + error_type: Final = "team_vector_store_access_denied" if scope == "team" else "key_vector_store_access_denied" + + response: Final = _rag_query(no_registry_gateway, model, marker, key, store_id=store_id) + searches: Final = _searches_for_marker(upstream_gateway, marker, store_id) + + assert response.status_code == 401, f"{response.text}; scripted_upstream_searches={searches!r}" + assert response.json()["error"]["type"] == error_type, response.text + assert searches == () + + +def test_no_registry_rag_query_allows_team_allowlisted_unregistered_store( + no_registry_gateways: tuple[Gateway, Gateway], +) -> None: + no_registry_gateway, upstream_gateway = no_registry_gateways + with no_registry_gateway.scenario() as scenario: + model: Final = scenario.model() + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + key: Final = _key_for_scope(scenario, model, "team", store_id) + marker: Final = f"lit5610 no registry allowed {uuid.uuid4().hex}" + + response: Final = _rag_query(no_registry_gateway, model, marker, key, store_id=store_id) + assert response.status_code == 200, response.text + + searches: Final = _searches_for_marker(upstream_gateway, marker, store_id) + assert len(searches) == 1, searches + assert marker in str(object_value(searches[0]["body"])["query"]), searches + + +def test_chat_completions_top_level_retrieval_config_uses_team_allowlist(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = _key_for_scope(scenario, model, "team", "vs_some_other_store") + marker: Final = f"lit5610 chat top-level retrieval config denied {uuid.uuid4().hex}" + body: Final[JsonObject] = { + "model": model, + "messages": _json_array({"role": "user", "content": marker}), + "retrieval_config": { + "vector_store_id": CONFIG_STORE_ID, + "custom_llm_provider": "openai", + "top_k": 1, + }, + } + + response: Final = gateway.request("POST", "/v1/chat/completions", body, key=key) + observations: Final = upstream_observations(gateway) + + assert response.status_code == 401, f"{response.text}; scripted_upstream_observations={observations!r}" + assert response.json()["error"]["type"] == "team_vector_store_access_denied", response.text + + +def test_rag_query_alias_denies_store_when_team_allowlist_excludes(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = _key_for_scope(scenario, model, "team", "vs_some_other_store") + marker: Final = f"lit5610 rag query alias denied {uuid.uuid4().hex}" + + response: Final = _rag_query(gateway, model, marker, key, path="/rag/query") + + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == "team_vector_store_access_denied", response.text + assert _searches_for_marker(gateway, marker) == () diff --git a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py index 353249dddf0..6c8b6571991 100644 --- a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py +++ b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py @@ -94,6 +94,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( tag_registry_cache_key, ) from litellm.utils import get_utc_datetime +from litellm.vector_stores.vector_store_registry import VectorStoreRegistry def _rendered_log_message(call): @@ -1753,6 +1754,65 @@ async def test_vector_store_access_check_with_team_permissions(): assert exc_info.value.type == ProxyErrorTypes.team_vector_store_access_denied +@pytest.mark.asyncio +@pytest.mark.parametrize( + "requested_vector_store_id,expected_error_type", + [ + ("KBOTHERTEAM99", ProxyErrorTypes.team_vector_store_access_denied), + ("KBALLOWED123", None), + ], +) +@pytest.mark.parametrize("vector_store_registry", [VectorStoreRegistry(), None], ids=["registry", "no-registry"]) +async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query( + requested_vector_store_id: str, + expected_error_type: ProxyErrorTypes | None, + vector_store_registry: VectorStoreRegistry | None, +): + """ + /v1/rag/query carries its vector store in retrieval_config.vector_store_id, + not in tools[].vector_store_ids. The team allowlist must apply either way. + """ + request_body = { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "what is in this KB?"}], + "retrieval_config": { + "vector_store_id": requested_vector_store_id, + "custom_llm_provider": "bedrock", + }, + } + valid_token = UserAPIKeyAuth(token="team-test-token", object_permission_id=None) + + team_object = MagicMock() + team_object.object_permission_id = "team-permission" + + mock_prisma_client = MagicMock() + team_permissions = MagicMock() + team_permissions.vector_stores = ["KBALLOWED123"] + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=team_permissions) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.vector_store_registry", vector_store_registry), + ): + if expected_error_type is None: + result = await vector_store_access_check( + request_body=request_body, + team_object=team_object, + valid_token=valid_token, + ) + assert result is True + return + + with pytest.raises(ProxyException) as exc_info: + await vector_store_access_check( + request_body=request_body, + team_object=team_object, + valid_token=valid_token, + ) + + assert exc_info.value.type == expected_error_type + + def test_can_object_call_model_with_alias(): """Test that can_object_call_model works with model aliases""" from litellm import Router From f88ac7424d9c7a5d510c2aae76cc56d42da97622 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Thu, 1 Oct 2026 13:50:55 -0700 Subject: [PATCH 013/203] feat(lens): move traces and setup into Lens (#44068) * feat(lens): move traces and setup into Lens * fix(lens): refresh trace readiness and preserve loaded traces --- .../_components/LensView.integration.test.tsx | 87 ++++- .../(dashboard)/lens/_components/LensView.tsx | 32 +- .../lens/_components/LensWelcome.tsx | 48 ++- .../src/app/(dashboard)/lens/page.test.tsx | 63 ++++ .../src/app/(dashboard)/lens/page.tsx | 50 ++- .../src/components/leftnav.test.tsx | 1 + .../src/components/leftnav.tsx | 2 - .../view_logs/TraceView/AgentTracesPage.tsx | 5 +- .../TraceView/AgentTracesSection.test.tsx | 103 +++++- .../TraceView/AgentTracesSection.tsx | 54 ++- .../TraceView/TracingSetupCard.test.tsx | 77 +++-- .../view_logs/TraceView/TracingSetupCard.tsx | 324 ++++++++++-------- .../view_logs/TraceView/useAgentTraces.ts | 17 +- .../src/components/view_logs/index.test.tsx | 8 +- .../src/components/view_logs/index.tsx | 17 +- 15 files changed, 641 insertions(+), 247 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/lens/page.test.tsx diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensView.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensView.integration.test.tsx index 179de1f6643..0d9344fdd14 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensView.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensView.integration.test.tsx @@ -1,8 +1,10 @@ -import { screen, within } from "@testing-library/react"; +import { act, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, testQueryClient } from "@/../tests/test-utils"; +import { ApiError } from "@/lib/http/client"; import { apiClient } from "@/components/networking"; +import { LIVE_TAIL_INTERVAL_MS } from "@/components/view_logs/log_filter_logic"; import { LensView } from "./LensView"; import { nextCheckStatus, type Lens, type Finding } from "./lensData"; @@ -217,12 +219,16 @@ it("runs saved settings immediately without opening setup", async () => { it("guides a first-time administrator into worker connection and lens setup", async () => { testQueryClient.clear(); vi.mocked(apiClient.get).mockImplementation(async (path) => - path === "/lens" ? { lenses: [], workers: [], tracing_enabled: true } : { data: [] }, + path === "/lens" ? { lenses: [], workers: [], tracing_enabled: true } : { data: [{ trace_id: "first-trace" }] }, ); const user = userEvent.setup(); renderWithProviders(); const guide = within(await screen.findByRole("region", { name: "Understand what your agents are doing" })); - expect(guide.getByRole("link", { name: "View logs" })).toHaveAttribute("href", "/ui/logs/"); + expect(apiClient.get).toHaveBeenCalledWith("/v1/traces", { accessToken: "test", query: { start_ms: 0 } }); + expect(guide.getByRole("link", { name: "View traces" })).toHaveAttribute( + "href", + expect.stringMatching(/^\/ui\/lens\/?\?tab=traces$/), + ); await user.click(guide.getByRole("button", { name: "Connect analyzer" })); const connection = within(await screen.findByRole("dialog", { name: "Set up Lens analysis" })); expect(connection.getByRole("button", { name: "Generate setup command" })).toBeVisible(); @@ -298,3 +304,78 @@ it("reads request content from the beginning after its abbreviated preview", asy await user.click(screen.getByRole("button", { name: "Previous section" })); expect(await screen.findByText("Abbreviated preview")).toBeVisible(); }); + +it.each([false, true])( + "directs a new user to traces when tracing_enabled=%s and there are no traces", + async (enabled) => { + testQueryClient.clear(); + vi.mocked(apiClient.get).mockImplementation(async (path) => + path === "/lens" ? { lenses: [], workers: [], tracing_enabled: enabled } : { data: [] }, + ); + renderWithProviders(); + expect(await screen.findByRole("heading", { name: "Set up traces to start running investigations" })).toBeVisible(); + expect(screen.getByRole("link", { name: "Set up traces" })).toHaveAttribute( + "href", + expect.stringMatching(/^\/ui\/lens\/?\?tab=traces$/), + ); + expect(screen.queryByRole("button", { name: "Set up your first lens" })).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Set up analysis" })).not.toBeInTheDocument(); + }, +); + +it("enables first-lens setup when a trace arrives without leaving Investigations", async () => { + testQueryClient.clear(); + const traceCheck = vi.fn().mockResolvedValue({ data: [] }); + vi.mocked(apiClient.get).mockImplementation(async (path) => + path === "/lens" ? { lenses: [], workers: [], tracing_enabled: true } : traceCheck(), + ); + vi.useFakeTimers(); + try { + const view = renderWithProviders(); + await act(async () => vi.advanceTimersByTimeAsync(50)); + expect(screen.getByRole("link", { name: "Set up traces" })).toBeVisible(); + + traceCheck.mockResolvedValue({ data: [{ trace_id: "first-trace" }] }); + await act(async () => vi.advanceTimersByTimeAsync(LIVE_TAIL_INTERVAL_MS)); + expect(screen.getByRole("button", { name: "Set up your first lens" })).toBeVisible(); + expect(screen.queryByRole("link", { name: "Set up traces" })).not.toBeInTheDocument(); + + const completedChecks = traceCheck.mock.calls.length; + await act(async () => vi.advanceTimersByTimeAsync(LIVE_TAIL_INTERVAL_MS * 2)); + expect(traceCheck).toHaveBeenCalledTimes(completedChecks); + view.unmount(); + } finally { + vi.useRealTimers(); + } +}); + +it("allows retrying a failed trace readiness check without treating it as an empty account", async () => { + testQueryClient.clear(); + const traceCheck = vi + .fn() + .mockRejectedValueOnce(new ApiError("Trace storage unavailable", 503, {})) + .mockResolvedValue({ data: [] }); + vi.mocked(apiClient.get).mockImplementation(async (path) => { + if (path === "/lens") return { lenses: [], workers: [], tracing_enabled: true }; + if (path === "/v1/traces") return traceCheck(); + return { data: [] }; + }); + const user = userEvent.setup(); + renderWithProviders(); + expect(await screen.findByRole("alert")).toHaveTextContent("Could not check traces. Trace storage unavailable"); + expect(screen.queryByRole("link", { name: "Set up traces" })).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: "Retry" })); + expect(await screen.findByRole("link", { name: "Set up traces" })).toBeVisible(); +}); + +it("keeps saved investigations accessible when tracing is disabled", async () => { + testQueryClient.clear(); + vi.mocked(apiClient.get).mockImplementation(async (path) => { + if (path === "/lens") return { lenses: [lens], workers: [], tracing_enabled: false }; + if (path === "/lens/lens/runs") return lens.jobs; + return { data: [] }; + }); + renderWithProviders(); + expect(await screen.findByText(issue.title)).toBeVisible(); + expect(screen.queryByRole("link", { name: "Set up traces" })).not.toBeInTheDocument(); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensView.tsx index dacd93310fd..7a31e0773d6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensView.tsx @@ -3,18 +3,7 @@ import type { components } from "@/lib/http/schema"; import { useState } from "react"; import { useQuery, useQueryClient } from "@tanstack/react-query"; -import { - Aperture, - ArrowUpRight, - CheckCircle2, - Circle, - Info, - Layers3, - Pause, - Play, - Plus, - Settings2, -} from "lucide-react"; +import { ArrowUpRight, CheckCircle2, Circle, Info, Layers3, Pause, Play, Plus, Settings2 } from "lucide-react"; import { Button } from "@/components/ui/button"; import { Tabs, TabsList, TabsTrigger, TabsContent } from "@/components/ui/tabs"; import { Sheet, SheetContent, SheetHeader, SheetTitle, SheetDescription } from "@/components/ui/sheet"; @@ -186,18 +175,9 @@ export function LensView({ accessToken, readOnly = false }: { accessToken: strin }; return ( -
-
-
-
-
-

- Understand your agent activity. Find patterns worth acting on. -

-
- {!readOnly && ( +
+
+ {!readOnly && !showEmpty && (
-
+ ); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensWelcome.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensWelcome.tsx index 0b92aa9018e..317e9c92747 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensWelcome.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensWelcome.tsx @@ -1,18 +1,58 @@ +import Link from "next/link"; +import { isTracingNotEnabled, useTraceAvailability } from "@/components/view_logs/TraceView/useAgentTraces"; import { Aperture, ArrowUpRight, CheckCircle2 } from "lucide-react"; import { Button } from "@/components/ui/button"; import { uiHref } from "@/utils/uiHref"; export function LensWelcome({ + accessToken, + tracingEnabled, connected, readOnly, onConnect, onCreate, }: { + accessToken: string; + tracingEnabled: boolean; connected: boolean; readOnly: boolean; onConnect: () => void; onCreate: () => void; }) { + const traces = useTraceAvailability(accessToken, tracingEnabled); + if (tracingEnabled && traces.isPending) { + return ( +

+ Checking for traces… +

+ ); + } + if (traces.error && !isTracingNotEnabled(traces.error)) { + return ( +
+

Could not check traces. {traces.error.message}

+ +
+ ); + } + if (!tracingEnabled || !traces.data || isTracingNotEnabled(traces.error)) { + return ( +
+

Set up traces to start running investigations

+ {tracingEnabled && !traces.error && ( +

No agent traces received yet.

+ )} + + Set up traces
+ ); + } return (
@@ -34,12 +74,12 @@ export function LensWelcome({ Use the agent traces or LLM requests already in LiteLLM. Lens needs their inputs and outputs to understand what happened.

- - View logs + View traces