From 615ed7900f3bcfd09b324c38c5123eb6ed4769da Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 05:44:21 -0700 Subject: [PATCH 001/139] test(mcp): fire MCP client test deadlines on conditions instead of wall-clock time (#44166) * test(mcp): fire MCP client test deadlines on conditions instead of wall-clock time Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): bound MCP client test failure paths independently of the triggered deadline Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): measure outer-deadline cleanup budget from cancellation, not setup 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> --- .../test_mcp_client.py | 82 ++++++++++++++----- 1 file changed, 61 insertions(+), 21 deletions(-) diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 504219a64e1..508ab447326 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -3,8 +3,9 @@ import base64 import importlib import json import os +import selectors import sys -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Callable from pathlib import Path from types import ModuleType from typing import Final @@ -84,8 +85,8 @@ class _MockTransportClient(MCPClient): class _ManualClockLoop(asyncio.SelectorEventLoop): """An event loop whose clock moves only when the test advances it, so timeouts fire on test-controlled conditions""" - def __init__(self) -> None: - super().__init__() + def __init__(self, selector: selectors.BaseSelector | None = None) -> None: + super().__init__(selector) self._now = 0.0 def time(self) -> float: @@ -95,6 +96,28 @@ class _ManualClockLoop(asyncio.SelectorEventLoop): self._now += seconds +class _AutojumpSelector(selectors.DefaultSelector): + def __init__(self, advance: Callable[[float], None]) -> None: + super().__init__() + self._advance = advance + + def select(self, timeout: float | None = None) -> list[tuple[selectors.SelectorKey, int]]: + ready: Final = super().select(0) + if ready or timeout == 0: + return ready + if timeout is None: + return super().select(None) + self._advance(timeout) + return [] + + +class _AutojumpClockLoop(_ManualClockLoop): + """A manual-clock loop that jumps to the next timer only once no callback or I/O event is left to run""" + + def __init__(self) -> None: + super().__init__(_AutojumpSelector(self.advance)) + + class _FakeExceptionGroup(Exception): """Duck-typed stand-in for an anyio/builtin ExceptionGroup. @@ -1890,16 +1913,26 @@ async def test_transport_parsing_failure_is_preserved(transport: MCPTransport, f ) -@pytest.mark.asyncio -async def test_sse_read_failure_is_preserved() -> None: - client: Final = MCPClient(server_url="https://example.com/sse", transport_type=MCPTransport.sse, timeout=0.2) - with pytest.raises(httpx2.ReadError, match="secret-read-error"): - await asyncio.wait_for( - client._execute_session_operation( - _diagnostic_transport(MCPTransport.sse, "io-error", "tools/list"), lambda session: session.list_tools() - ), - timeout=3, - ) +def test_sse_read_failure_is_preserved() -> None: + loop: Final = _AutojumpClockLoop() + + async def run() -> None: + client: Final = MCPClient(server_url="https://example.com/sse", transport_type=MCPTransport.sse, timeout=0.2) + with pytest.raises(httpx2.ReadError, match="secret-read-error"): + await asyncio.wait_for( + client._execute_session_operation( + _diagnostic_transport(MCPTransport.sse, "io-error", "tools/list"), + lambda session: session.list_tools(), + ), + timeout=3, + ) + + try: + loop.run_until_complete(run()) + finally: + loop.run_until_complete(loop.shutdown_asyncgens()) + loop.run_until_complete(loop.shutdown_default_executor()) + loop.close() @pytest.mark.asyncio @@ -2638,9 +2671,12 @@ def test_public_mcp_import_preserves_incompatible_sdk_error() -> None: @pytest.mark.parametrize("grouped", (False, True)) @pytest.mark.parametrize("raise_on_error", (False, True)) @pytest.mark.parametrize("termination", ("ok", "failure", "hang")) -async def test_outer_deadline_delivers_session_termination(termination: str, grouped: bool, raise_on_error: bool) -> None: +async def test_outer_deadline_delivers_session_termination( + termination: str, grouped: bool, raise_on_error: bool +) -> None: deleted: Final = asyncio.Event() started: Final = asyncio.Event() + caller_deadline: Final[asyncio.Future[anyio.CancelScope]] = asyncio.get_running_loop().create_future() async def respond(request: httpx2.Request) -> httpx2.Response: await anyio.lowlevel.checkpoint() @@ -2672,26 +2708,30 @@ async def test_outer_deadline_delivers_session_termination(termination: str, gro if payload.method == "tools/list": return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"tools": []}}) started.set() + caller_deadline.result().deadline = anyio.current_time() await anyio.sleep_forever() raise AssertionError("cancelled request resumed") client: Final = _MockTransportClient(respond, server_url="https://example.com/mcp", timeout=30) - async def invoke(): - with anyio.fail_after(0.2): - pending: Final = client.call_tool(CallToolRequestParams(name="slow", arguments={}), raise_on_error=raise_on_error) + async def invoke() -> None: + with anyio.fail_after(None) as deadline: + caller_deadline.set_result(deadline) + pending: Final = client.call_tool( + CallToolRequestParams(name="slow", arguments={}), raise_on_error=raise_on_error + ) if grouped: await asyncio.gather(pending) else: await pending - before: Final = anyio.current_time() - with pytest.raises(TimeoutError): - await invoke() + with anyio.fail_after(20): + with pytest.raises(TimeoutError): + await invoke() assert started.is_set() assert deleted.is_set(), "Cancellation must deliver DELETE before returning to the caller" - assert anyio.current_time() - before < 6.5 + assert anyio.current_time() - caller_deadline.result().deadline < 6.5 assert await client.list_tools(raise_on_error=True) == [] From a5fef4b4e68963640c3062d464509129ec8863c6 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 09:26:19 -0700 Subject: [PATCH 002/139] fix(cost-map): add OpenAI TTS and GPT-5.x deprecation dates (#44175) * fix(cost-map): add OpenAI TTS and GPT-5.x deprecation dates Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(cost-map): credit OpenAI TTS deprecation dates from #44113 Co-authored-by: Ben Langfeld <210221+benlangfeld@users.noreply.github.com> 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> Co-authored-by: Ben Langfeld <210221+benlangfeld@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 13 ++++++++++--- model_prices_and_context_window.json | 13 ++++++++++--- 2 files changed, 20 insertions(+), 6 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 3b9956ad1a6..6a548e7a82d 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -32815,6 +32815,7 @@ ] }, "gpt-4o-mini-tts-2025-03-20": { + "deprecation_date": "2027-01-06", "input_cost_per_token": 6e-07, "litellm_provider": "openai", "mode": "audio_speech", @@ -33514,7 +33515,8 @@ "output_cost_per_token_flex": 5e-06, "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": false, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "deprecation_date": "2027-04-01" }, "gpt-5.1-codex-mini": { "cache_read_input_token_cost": 2.5e-08, @@ -35311,7 +35313,8 @@ "default_reasoning_effort": "none", "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "deprecation_date": "2027-04-01" }, "gpt-5.4-nano-2026-03-17": { "cache_read_input_token_cost": 2e-08, @@ -35610,7 +35613,8 @@ "supports_web_search": true, "supports_none_reasoning_effort": false, "supports_xhigh_reasoning_effort": false, - "supports_minimal_reasoning_effort": true + "supports_minimal_reasoning_effort": true, + "deprecation_date": "2027-04-01" }, "gpt-5.3-chat-latest": { "cache_read_input_token_cost": 1.75e-07, @@ -46869,6 +46873,7 @@ "supports_vision": true }, "tts-1": { + "deprecation_date": "2027-01-06", "input_cost_per_character": 1.5e-05, "litellm_provider": "openai", "mode": "audio_speech", @@ -46878,6 +46883,7 @@ ] }, "tts-1-hd": { + "deprecation_date": "2027-01-06", "input_cost_per_character": 3e-05, "litellm_provider": "openai", "mode": "audio_speech", @@ -56760,6 +56766,7 @@ ] }, "gpt-4o-mini-tts-2025-12-15": { + "deprecation_date": "2027-01-06", "input_cost_per_token": 6e-07, "litellm_provider": "openai", "mode": "audio_speech", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 3b9956ad1a6..6a548e7a82d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -32815,6 +32815,7 @@ ] }, "gpt-4o-mini-tts-2025-03-20": { + "deprecation_date": "2027-01-06", "input_cost_per_token": 6e-07, "litellm_provider": "openai", "mode": "audio_speech", @@ -33514,7 +33515,8 @@ "output_cost_per_token_flex": 5e-06, "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": false, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "deprecation_date": "2027-04-01" }, "gpt-5.1-codex-mini": { "cache_read_input_token_cost": 2.5e-08, @@ -35311,7 +35313,8 @@ "default_reasoning_effort": "none", "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "deprecation_date": "2027-04-01" }, "gpt-5.4-nano-2026-03-17": { "cache_read_input_token_cost": 2e-08, @@ -35610,7 +35613,8 @@ "supports_web_search": true, "supports_none_reasoning_effort": false, "supports_xhigh_reasoning_effort": false, - "supports_minimal_reasoning_effort": true + "supports_minimal_reasoning_effort": true, + "deprecation_date": "2027-04-01" }, "gpt-5.3-chat-latest": { "cache_read_input_token_cost": 1.75e-07, @@ -46869,6 +46873,7 @@ "supports_vision": true }, "tts-1": { + "deprecation_date": "2027-01-06", "input_cost_per_character": 1.5e-05, "litellm_provider": "openai", "mode": "audio_speech", @@ -46878,6 +46883,7 @@ ] }, "tts-1-hd": { + "deprecation_date": "2027-01-06", "input_cost_per_character": 3e-05, "litellm_provider": "openai", "mode": "audio_speech", @@ -56760,6 +56766,7 @@ ] }, "gpt-4o-mini-tts-2025-12-15": { + "deprecation_date": "2027-01-06", "input_cost_per_token": 6e-07, "litellm_provider": "openai", "mode": "audio_speech", From 276fc9c63a37fa86f2ea497ab85c4ec34af9f16a Mon Sep 17 00:00:00 2001 From: yujonglee Date: Fri, 2 Oct 2026 09:31:26 -0700 Subject: [PATCH 003/139] fix(tracing): unify ClickHouse storage configuration (#43941) * fix(tracing): use ClickHouse URL for reads by default * fix(tracing): unify ClickHouse storage configuration * fix(tracing): update dashboard setup copy for one URL * test(tracing): make tests/unit/tracing a package Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(config): drop legacy string tracing store variant Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(tracing): own ClickHouse defaults in constants and reject unset env references Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(tracing): use raw regex patterns in config tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(tracing): read ClickHouse env defaults when tracing config resolves Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): split audit log query guard to fit condition-chain budget 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> --- deploy/lens/README.md | 15 ++- docker/docker-compose.tracing.yml | 1 - docker/tracing-config.yaml | 5 +- litellm-rust/crates/config/src/lib.rs | 5 +- litellm-rust/crates/config/src/settings.rs | 44 +++++++ litellm-rust/crates/config/tests/config.rs | 50 +++++++- litellm-rust/crates/python-bridge/src/lib.rs | 3 +- .../crates/python-bridge/src/routes/traces.rs | 65 +++++----- .../crates/storage-clickhouse/README.md | 2 +- .../crates/storage-clickhouse/src/lib.rs | 12 +- .../storage-clickhouse/tests/connection.rs | 27 ++-- .../migrations/0002_otel_traces_ttl.sql | 2 +- .../migrations/0004_agent_traces_ttl.sql | 2 +- .../traces/migrations/0007_spend_logs_ttl.sql | 2 +- litellm-rust/crates/traces/src/config.rs | 25 ++++ litellm-rust/crates/traces/src/lib.rs | 2 + litellm-rust/crates/traces/src/schema.rs | 26 ++-- .../crates/traces/tests/migrations.rs | 56 ++++----- .../crates/traces/tests/query_access.rs | 2 +- litellm/constants.py | 4 +- .../clickhouse/clickhouse_batch_logger.py | 7 +- litellm/integrations/clickhouse/schema.py | 4 +- litellm/proxy/proxy_server.py | 13 +- litellm/proxy/tracing_runtime.py | 9 +- litellm/rust_bridge/_native.pyi | 14 ++- litellm/rust_bridge/traces.py | 31 ++++- litellm/tracing/config.py | 84 +++++++++++++ litellm/tracing/receiver.py | 30 ++--- scripts/run_tracing_proxy_local.sh | 21 +++- tests/test_litellm_rust/test_traces.py | 80 +++++++++--- .../proxy/proxy_server/test_proxy_config.py | 24 ++++ tests/unit/tracing/__init__.py | 0 tests/unit/tracing/test_config.py | 116 ++++++++++++++++++ .../components/view_logs/AuditLogsPanel.tsx | 3 +- .../TraceView/AgentTracesSection.test.tsx | 5 +- .../TraceView/TracingSetupCard.test.tsx | 3 +- .../view_logs/TraceView/TracingSetupCard.tsx | 11 +- 37 files changed, 614 insertions(+), 191 deletions(-) create mode 100644 litellm-rust/crates/traces/src/config.rs create mode 100644 litellm/tracing/config.py create mode 100644 tests/unit/tracing/__init__.py create mode 100644 tests/unit/tracing/test_config.py diff --git a/deploy/lens/README.md b/deploy/lens/README.md index 7a80bafa59e..315d072b501 100644 --- a/deploy/lens/README.md +++ b/deploy/lens/README.md @@ -4,7 +4,20 @@ Lens reviews recorded activity and saves evidence-linked findings in the LiteLLM ## Start a worker -Upgrade your existing LiteLLM proxy to a release that includes Lens with PostgreSQL, agent tracing (`general_settings.tracing: {store: clickhouse}`), and ClickHouse configured through `CLICKHOUSE_URL` and a separate SELECT-only `CLICKHOUSE_READER_URL`. Enable the ClickHouse callback and request/response logging to analyze LLM requests. Lens can only inspect content you actually retain +Upgrade your existing LiteLLM proxy to a release that includes Lens with PostgreSQL and agent tracing. Configure one ClickHouse URL for trace writes, bounded reads, and Lens queries: + +```yaml +general_settings: + tracing: + store: + type: clickhouse + url: os.environ/CLICKHOUSE_URL + retention_days: 14 +``` + +The URL, database, and retention settings can also come from `CLICKHOUSE_URL`, `CLICKHOUSE_DATABASE`, and `AGENT_TRACING_RETENTION_DAYS` when omitted from YAML. A YAML value wins when both are set. The database defaults to `litellm`. `retention_days` defaults to 14 and applies to both traces and spend logs + +Retention changes require a proxy restart. ClickHouse removes expired rows during background merges, not immediately at startup. Enable request/response logging to analyze LLM requests. Lens can only inspect content you actually retain In Lens, click **Set up analysis**, choose an existing virtual key or **Create worker key**, then **Generate setup command**. The LiteLLM address is filled in for you; change it only if the server running Docker needs a different network address. Copy the command and run it on your server. The dialog changes to **Analyzer connected** when the container checks in diff --git a/docker/docker-compose.tracing.yml b/docker/docker-compose.tracing.yml index b39fc8f4561..4f87e49eb27 100644 --- a/docker/docker-compose.tracing.yml +++ b/docker/docker-compose.tracing.yml @@ -12,7 +12,6 @@ services: DATABASE_URL: postgresql://litellm:litellm@db:5432/litellm STORE_MODEL_IN_DB: "True" CLICKHOUSE_URL: http://default:local-tracing@clickhouse:8123 - CLICKHOUSE_READER_URL: http://default:local-tracing@clickhouse:8123 CLICKHOUSE_DATABASE: litellm OPENAI_API_KEY: ${OPENAI_API_KEY:-} volumes: diff --git a/docker/tracing-config.yaml b/docker/tracing-config.yaml index 03637cfa9fb..d8e3759641f 100644 --- a/docker/tracing-config.yaml +++ b/docker/tracing-config.yaml @@ -7,4 +7,7 @@ model_list: general_settings: master_key: os.environ/LITELLM_MASTER_KEY tracing: - store: clickhouse + store: + type: clickhouse + url: os.environ/CLICKHOUSE_URL + retention_days: 14 diff --git a/litellm-rust/crates/config/src/lib.rs b/litellm-rust/crates/config/src/lib.rs index e7010941c3d..fed60ab1a4f 100644 --- a/litellm-rust/crates/config/src/lib.rs +++ b/litellm-rust/crates/config/src/lib.rs @@ -12,7 +12,10 @@ use serde::Deserialize; pub use error::Error; pub use mcp::{McpAuth, McpServer, McpTransport}; pub use model::{LiteLlmParams, Model}; -pub use settings::{GeneralSettings, LiteLlmSettings, RouterSettings}; +pub use settings::{ + ClickHouseStoreSettings, GeneralSettings, LiteLlmSettings, RouterSettings, TracingSettings, + TracingStoreSettings, +}; pub use value::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value}; #[derive(Clone, Default, Deserialize)] diff --git a/litellm-rust/crates/config/src/settings.rs b/litellm-rust/crates/config/src/settings.rs index b6b97475eda..b1ead35e55b 100644 --- a/litellm-rust/crates/config/src/settings.rs +++ b/litellm-rust/crates/config/src/settings.rs @@ -5,6 +5,47 @@ use serde::Deserialize; use crate::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value}; +#[derive(Clone, Debug, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum TracingStoreKind { + Clickhouse, +} + +#[derive(Clone, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ClickHouseStoreSettings { + #[serde(rename = "type")] + pub kind: TracingStoreKind, + pub url: Option, + pub database: Option, + pub retention_days: Option, +} + +impl fmt::Debug for ClickHouseStoreSettings { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ClickHouseStoreSettings") + .field("kind", &self.kind) + .field("database", &self.database) + .field("retention_days", &self.retention_days) + .finish() + } +} + +#[derive(Clone, Debug, Deserialize)] +#[serde(untagged)] +pub enum TracingStoreSettings { + ClickHouse(ClickHouseStoreSettings), +} + +#[derive(Clone, Default, Debug, Deserialize)] +#[serde(default)] +pub struct TracingSettings { + pub store: Option, + #[serde(flatten)] + pub additional_fields: AdditionalFields, +} + #[derive(Clone, Deserialize)] #[serde(default)] pub struct GeneralSettings { @@ -14,6 +55,7 @@ pub struct GeneralSettings { pub admission_queue_timeout_seconds: f64, pub master_key: Option, pub database_url: Option, + pub tracing: Option, pub database_connection_pool_limit: Option, pub database_connection_timeout: Option, pub database_connect_timeout: Option, @@ -50,6 +92,7 @@ impl Default for GeneralSettings { admission_queue_timeout_seconds: 1.0, master_key: None, database_url: None, + tracing: None, database_connection_pool_limit: Some(10), database_connection_timeout: Some(60.0), database_connect_timeout: None, @@ -97,6 +140,7 @@ impl fmt::Debug for GeneralSettings { ) .field("master_key", &self.master_key) .field("database_url", &self.database_url) + .field("tracing", &self.tracing) .field("store_model_in_db", &self.store_model_in_db) .field("additional_fields", &self.additional_fields.keys()) .finish_non_exhaustive() diff --git a/litellm-rust/crates/config/tests/config.rs b/litellm-rust/crates/config/tests/config.rs index 447aa9e1d2a..ab9f3403a01 100644 --- a/litellm-rust/crates/config/tests/config.rs +++ b/litellm-rust/crates/config/tests/config.rs @@ -1,4 +1,4 @@ -use litellm_config::{Config, Error, Flag, NumberOrString}; +use litellm_config::{Config, Error, Flag, NumberOrString, TracingStoreSettings}; use rstest::{fixture, rstest}; use tempfile::TempDir; @@ -113,6 +113,54 @@ fn missing_general_settings_has_no_master_key() { assert!(config.general_settings.master_key.is_none()); } +#[test] +fn tracing_settings_are_typed_and_redact_the_url() { + let config = Config::from_yaml( + "general_settings:\n tracing:\n store:\n type: clickhouse\n url: https://writer:password@example.com\n database: analytics\n retention_days: 7\n", + ) + .unwrap(); + let tracing = config.general_settings.tracing.as_ref().unwrap(); + let Some(TracingStoreSettings::ClickHouse(store)) = tracing.store.as_ref() else { + panic!("expected ClickHouse tracing store") + }; + assert_eq!( + store.url.as_ref().unwrap().expose(), + "https://writer:password@example.com" + ); + assert_eq!(store.database.as_deref(), Some("analytics")); + assert_eq!(store.retention_days, Some(NumberOrString::Number(7.0))); + assert!(!format!("{config:?}").contains("password")); +} + +#[test] +fn tracing_settings_accept_environment_references() { + let config = Config::from_yaml( + "general_settings:\n tracing:\n store:\n type: clickhouse\n url: os.environ/CLICKHOUSE_URL\n retention_days: os.environ/RETENTION_DAYS\n", + ) + .unwrap(); + let Some(TracingStoreSettings::ClickHouse(store)) = + config.general_settings.tracing.unwrap().store + else { + panic!("expected ClickHouse tracing store") + }; + assert_eq!( + store.retention_days, + Some(NumberOrString::String( + "os.environ/RETENTION_DAYS".to_owned() + )) + ); +} + +#[test] +fn tracing_settings_reject_string_store() { + assert!(Config::from_yaml("general_settings:\n tracing:\n store: clickhouse\n").is_err()); +} + +#[test] +fn tracing_settings_reject_removed_reader_configuration() { + assert!(Config::from_yaml("general_settings:\n tracing:\n store:\n type: clickhouse\n reader_url: http://localhost:8123\n").is_err()); +} + #[rstest] fn empty_config_matches_python_defaults() { let config = Config::from_yaml("{}").unwrap(); diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 6e091f0594f..67a34e6fda5 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -45,7 +45,7 @@ mod _native { use crate::routes::token_counter::TokenCounter; #[pymodule_export] use crate::routes::traces::{ - NativeTraceStorage, trace_decode_otlp, trace_encode_error, + NativeTraceConfig, NativeTraceStorage, trace_decode_otlp, trace_encode_error, trace_normalized_field_definitions, }; #[cfg(feature = "huggingface")] @@ -112,6 +112,7 @@ mod tests { "aresponses", "ResponsesWebSocketConnection", "NativeDiagnosticProcessor", + "NativeTraceConfig", "NativeTraceStorage", "trace_decode_otlp", "trace_encode_error", diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 0ff5fe98f55..f81fc6a6d75 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -2,9 +2,9 @@ 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, QueryAccessError, QueryReaders, QueryScope, ReadQuery, Shared, + Config, Error, InsertTable, Parameter, QueryAccessError, QueryReaders, QueryScope, ReadQuery, + Shared, }; use prost::Message; use pyo3::{ @@ -63,48 +63,49 @@ fn map_query_access_error(error: QueryAccessError) -> PyErr { } } +#[pyclass(frozen)] +pub struct NativeTraceConfig { + inner: Config, +} + +#[pymethods] +impl NativeTraceConfig { + #[new] + fn new(database: String, url: &str, retention_days: u32) -> PyResult { + Ok(Self { + inner: Config::new(database, url, retention_days).map_err(map_error)?, + }) + } +} + #[pyclass] pub struct NativeTraceStorage { - storage: Storage, + config: Config, query_readers: QueryReaders, } #[pymethods] impl NativeTraceStorage { #[new] - #[pyo3(signature = (database, url, reader_url = None))] - fn new(database: String, url: &str, reader_url: Option<&str>) -> PyResult { - litellm_traces::schema_statements(&database, 1, 1).map_err(map_error)?; - let storage = Storage::new(database, url, reader_url).map_err(map_error)?; + fn new(config: PyRef<'_, NativeTraceConfig>) -> PyResult { Ok(Self { query_readers: QueryReaders::new( - storage.writer().clone(), - storage.database().to_owned(), + config.inner.storage().writer().clone(), + config.inner.storage().database().to_owned(), ), - storage, + config: config.inner.clone(), }) } - fn ensure_schema<'py>( - &self, - py: Python<'py>, - trace_retention_days: u32, - spend_log_retention_days: u32, - ) -> PyResult> { + fn ensure_schema<'py>(&self, py: Python<'py>) -> PyResult> { let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; - let connection = self.storage.writer().clone(); - let database = self.storage.database().to_owned(); + let connection = self.config.storage().writer().clone(); + let database = self.config.storage().database().to_owned(); + let retention_days = self.config.retention_days(); crate::execution::run_async( py, async move { - litellm_traces::ensure_schema( - &client, - &connection, - &database, - trace_retention_days, - spend_log_retention_days, - ) - .await + litellm_traces::ensure_schema(&client, &connection, &database, retention_days).await }, map_error, ) @@ -118,8 +119,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.storage.writer().clone(); - let database = self.storage.database().to_owned(); + let connection = self.config.storage().writer().clone(); + let database = self.config.storage().database().to_owned(); crate::execution::run_async( py, async move { @@ -186,9 +187,7 @@ impl NativeTraceStorage { >, ) -> PyResult> { let query = litellm_traces::LensQuery::parse(name).map_err(map_error)?; - let connection = self.storage.reader().cloned().ok_or_else(|| { - PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL") - })?; + let connection = self.config.storage().reader().clone(); let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; crate::execution::run_async( py, @@ -209,9 +208,7 @@ impl NativeTraceStorage { >, ) -> PyResult> { let query = ReadQuery::parse(query).map_err(map_error)?; - let connection = self.storage.reader().cloned().ok_or_else(|| { - PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL") - })?; + let connection = self.config.storage().reader().clone(); let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; crate::execution::run_async( py, diff --git a/litellm-rust/crates/storage-clickhouse/README.md b/litellm-rust/crates/storage-clickhouse/README.md index 7c4f86e4589..676260f6043 100644 --- a/litellm-rust/crates/storage-clickhouse/README.md +++ b/litellm-rust/crates/storage-clickhouse/README.md @@ -1,5 +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 +`litellm-storage-clickhouse` exports `Storage`, a writer and bounded reader derived from one ClickHouse URL and 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/lib.rs b/litellm-rust/crates/storage-clickhouse/src/lib.rs index d11ee9d5cde..6a6a957d456 100644 --- a/litellm-rust/crates/storage-clickhouse/src/lib.rs +++ b/litellm-rust/crates/storage-clickhouse/src/lib.rs @@ -89,19 +89,17 @@ impl Connection { pub struct Storage { database: String, writer: Connection, - reader: Option, + reader: Connection, } impl Storage { - pub fn new(database: String, url: &str, reader_url: Option<&str>) -> Result { + pub fn new(database: String, url: &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()?, + reader: Connection::reader(url, &database)?, database, }) } @@ -114,8 +112,8 @@ impl Storage { &self.writer } - pub fn reader(&self) -> Option<&Connection> { - self.reader.as_ref() + pub fn reader(&self) -> &Connection { + &self.reader } } diff --git a/litellm-rust/crates/storage-clickhouse/tests/connection.rs b/litellm-rust/crates/storage-clickhouse/tests/connection.rs index 0874b693249..e371718259a 100644 --- a/litellm-rust/crates/storage-clickhouse/tests/connection.rs +++ b/litellm-rust/crates/storage-clickhouse/tests/connection.rs @@ -10,25 +10,30 @@ 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"); +#[test] +fn storage_uses_one_url_for_writes_and_bounded_reads() { + let storage = + Storage::new("litellm".to_owned(), "http://localhost:8123").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); + assert_eq!(storage.reader().url().port(), Some(8123)); + assert_eq!( + storage + .reader() + .url() + .query_pairs() + .find(|(key, _)| key == "database") + .unwrap() + .1, + "litellm" + ); } #[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()); + assert!(Storage::new(database.to_owned(), "http://localhost:8123").is_err()); } diff --git a/litellm-rust/crates/traces/migrations/0002_otel_traces_ttl.sql b/litellm-rust/crates/traces/migrations/0002_otel_traces_ttl.sql index 4ac597b8902..7402634b7e1 100644 --- a/litellm-rust/crates/traces/migrations/0002_otel_traces_ttl.sql +++ b/litellm-rust/crates/traces/migrations/0002_otel_traces_ttl.sql @@ -1 +1 @@ -ALTER TABLE {database}.otel_traces MODIFY TTL toDateTime(Timestamp) + INTERVAL {trace_retention_days} DAY +ALTER TABLE {database}.otel_traces MODIFY TTL toDateTime(Timestamp) + INTERVAL {retention_days} DAY diff --git a/litellm-rust/crates/traces/migrations/0004_agent_traces_ttl.sql b/litellm-rust/crates/traces/migrations/0004_agent_traces_ttl.sql index 8681f0622a4..70147f95d0e 100644 --- a/litellm-rust/crates/traces/migrations/0004_agent_traces_ttl.sql +++ b/litellm-rust/crates/traces/migrations/0004_agent_traces_ttl.sql @@ -1 +1 @@ -ALTER TABLE {database}.agent_traces_by_key MODIFY TTL toDateTime(StartTs) + INTERVAL {trace_retention_days} DAY +ALTER TABLE {database}.agent_traces_by_key MODIFY TTL toDateTime(StartTs) + INTERVAL {retention_days} DAY diff --git a/litellm-rust/crates/traces/migrations/0007_spend_logs_ttl.sql b/litellm-rust/crates/traces/migrations/0007_spend_logs_ttl.sql index 131573927ac..d9b1a2403b4 100644 --- a/litellm-rust/crates/traces/migrations/0007_spend_logs_ttl.sql +++ b/litellm-rust/crates/traces/migrations/0007_spend_logs_ttl.sql @@ -1 +1 @@ -ALTER TABLE {database}.spend_logs MODIFY TTL toDateTime(start_time) + INTERVAL {spend_log_retention_days} DAY +ALTER TABLE {database}.spend_logs MODIFY TTL toDateTime(start_time) + INTERVAL {retention_days} DAY diff --git a/litellm-rust/crates/traces/src/config.rs b/litellm-rust/crates/traces/src/config.rs new file mode 100644 index 00000000000..ec88fca2e3f --- /dev/null +++ b/litellm-rust/crates/traces/src/config.rs @@ -0,0 +1,25 @@ +use litellm_storage_clickhouse::{Error, Storage}; + +#[derive(Clone)] +pub struct Config { + storage: Storage, + retention_days: u32, +} + +impl Config { + pub fn new(database: String, url: &str, retention_days: u32) -> Result { + crate::schema_statements(&database, retention_days)?; + Ok(Self { + storage: Storage::new(database, url)?, + retention_days, + }) + } + + pub fn storage(&self) -> &Storage { + &self.storage + } + + pub fn retention_days(&self) -> u32 { + self.retention_days + } +} diff --git a/litellm-rust/crates/traces/src/lib.rs b/litellm-rust/crates/traces/src/lib.rs index 9e01d10f73e..05b56c7ea46 100644 --- a/litellm-rust/crates/traces/src/lib.rs +++ b/litellm-rust/crates/traces/src/lib.rs @@ -1,3 +1,4 @@ +mod config; mod error; mod insert; mod normalize; @@ -8,6 +9,7 @@ mod schema; mod shared; mod sql; +pub use config::Config; pub use error::{DecodeError, QueryAccessError}; pub use insert::{InsertRow, InsertTable, encode_rows, insert_rows, insert_shared_rows}; pub use litellm_storage_clickhouse::{Connection, Error, Parameter, execute_read}; diff --git a/litellm-rust/crates/traces/src/schema.rs b/litellm-rust/crates/traces/src/schema.rs index cf11c564931..c52db1523cc 100644 --- a/litellm-rust/crates/traces/src/schema.rs +++ b/litellm-rust/crates/traces/src/schema.rs @@ -9,17 +9,12 @@ const SCHEMA_REQUEST_TIMEOUT: Duration = Duration::from_secs(30); const MIGRATIONS: &[Migration] = litellm_migrate::migrate!("migrations"); -pub fn schema_statements( - database: &str, - trace_retention_days: u32, - spend_log_retention_days: u32, -) -> Result, Error> { +pub fn schema_statements(database: &str, retention_days: u32) -> Result, Error> { if database.is_empty() || !database .bytes() .all(|c| c.is_ascii_alphanumeric() || c == b'_') - || trace_retention_days == 0 - || spend_log_retention_days == 0 + || retention_days == 0 { return Err(Error::InvalidSchema); } @@ -30,11 +25,7 @@ pub fn schema_statements( migration .sql .replace("{database}", &database) - .replace("{trace_retention_days}", &trace_retention_days.to_string()) - .replace( - "{spend_log_retention_days}", - &spend_log_retention_days.to_string(), - ) + .replace("{retention_days}", &retention_days.to_string()) })) .collect(), ) @@ -44,15 +35,13 @@ pub async fn ensure_schema( client: &Client, connection: &Connection, database: &str, - trace_retention_days: u32, - spend_log_retention_days: u32, + retention_days: u32, ) -> Result<(), Error> { ensure_schema_with_timeout( client, connection, database, - trace_retention_days, - spend_log_retention_days, + retention_days, SCHEMA_REQUEST_TIMEOUT, ) .await @@ -62,11 +51,10 @@ async fn ensure_schema_with_timeout( client: &Client, connection: &Connection, database: &str, - trace_retention_days: u32, - spend_log_retention_days: u32, + retention_days: u32, request_timeout: Duration, ) -> Result<(), Error> { - for statement in schema_statements(database, trace_retention_days, spend_log_retention_days)? { + for statement in schema_statements(database, retention_days)? { let response = client .post(connection.url().clone()) .timeout(request_timeout) diff --git a/litellm-rust/crates/traces/tests/migrations.rs b/litellm-rust/crates/traces/tests/migrations.rs index d93083b052c..ccefa8b0b9f 100644 --- a/litellm-rust/crates/traces/tests/migrations.rs +++ b/litellm-rust/crates/traces/tests/migrations.rs @@ -109,8 +109,8 @@ async fn schema_supports_span_rollups_and_spend_joins( ) -> TestResult { let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; let span = serde_json::from_value(serde_json::json!({ "Timestamp": timestamp, "TraceId": "trace-1", "SpanId": "span-1", "ParentSpanId": "", @@ -221,7 +221,6 @@ async fn normalized_fields_match_clickhouse_catalog( &Connection::writer(&database.url)?, "trace_test", 7, - 14, ) .await?; let catalog = read_json(&database, "SELECT name, type FROM system.columns WHERE database = 'trace_test' AND table = 'otel_traces'").await?; @@ -257,7 +256,7 @@ async fn insert_rejects_unknown_columns_even_if_url_requests_skipping_them( "{}?input_format_skip_unknown_fields=1", database.url ))?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let row = BTreeMap::from([ ( "Timestamp".to_owned(), @@ -291,7 +290,7 @@ async fn retried_trace_insert_does_not_inflate_rollup( ) -> TestResult { let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let row: BTreeMap = serde_json::from_value(serde_json::json!({ "Timestamp": time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64, "TraceId": "retried-trace", "SpanId": "span-1", "ParentSpanId": "", @@ -326,7 +325,7 @@ async fn keyed_rollup_keeps_same_trace_ids_separate_by_api_key( ) -> TestResult { let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; let rows = vec![ serde_json::from_value(serde_json::json!({ @@ -370,7 +369,7 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields( ) -> TestResult { let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let day_start = time::OffsetDateTime::now_utc() .replace_time(time::Time::MIDNIGHT) .unix_timestamp_nanos() as i64; @@ -417,7 +416,7 @@ async fn spend_deduplication_preserves_subsecond_requests_and_retries( ) -> TestResult { let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let now_ms = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64 / 1_000_000; let base_start_time = now_ms / 1000 * 1000; let first_start_time = base_start_time + 100; @@ -469,7 +468,7 @@ async fn retention_changes_materialize_existing_rows_and_remain_idempotent( ) -> TestResult { let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 30, 30).await?; + ensure_schema(&database.client, &writer, "trace_test", 30).await?; let tables = read_json( &database, "SELECT name FROM system.tables WHERE database = 'trace_test' \ @@ -499,7 +498,7 @@ async fn retention_changes_materialize_existing_rows_and_remain_idempotent( insert_rows(&database, "otel_traces", vec![span]).await?; insert_rows(&database, "spend_logs", vec![spend]).await?; assert_eq!(table_rows(&database, "agent_traces_by_key").await?, 1); - ensure_schema(&database.client, &writer, "trace_test", 14, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 14).await?; let deadline = tokio::time::Instant::now() + Duration::from_secs(60); loop { let response = read_json( @@ -531,7 +530,7 @@ async fn retention_changes_materialize_existing_rows_and_remain_idempotent( assert_eq!(table_rows(&database, "agent_traces_by_key").await?, 0); assert_eq!(table_rows(&database, "spend_logs").await?, 0); let mutation_count = mutation_rows(&database).await?; - ensure_schema(&database.client, &writer, "trace_test", 14, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 14).await?; assert_eq!(mutation_rows(&database).await?, mutation_count); Ok(()) } @@ -550,7 +549,7 @@ async fn schema_statement_timeout_maps_to_transport_error() -> TestResult { let writer = Connection::writer(&url)?; let result = tokio::time::timeout( Duration::from_secs(35), - ensure_schema(&client, &writer, "trace_test", 7, 14), + ensure_schema(&client, &writer, "trace_test", 7), ) .await; server.abort(); @@ -559,16 +558,11 @@ async fn schema_statement_timeout_maps_to_transport_error() -> TestResult { } #[rstest] -#[case::empty("", 7, 14)] -#[case::sql("db; DROP DATABASE default", 7, 14)] -#[case::trace_retention("traces", 0, 14)] -#[case::spend_retention("traces", 7, 0)] -fn schema_rejects_invalid_configuration( - #[case] database: &str, - #[case] traces: u32, - #[case] spend: u32, -) { - assert!(schema_statements(database, traces, spend).is_err()); +#[case::empty("", 7)] +#[case::sql("db; DROP DATABASE default", 7)] +#[case::retention("traces", 0)] +fn schema_rejects_invalid_configuration(#[case] database: &str, #[case] retention_days: u32) { + assert!(schema_statements(database, retention_days).is_err()); } #[rstest] @@ -579,7 +573,7 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate( use litellm_traces::{LensQuery, Parameter}; let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; for (key, text) in [("one", "timeout"), ("two", "success")] { insert_rows(&database, "otel_traces", vec![serde_json::from_value(serde_json::json!({ @@ -686,7 +680,7 @@ async fn lens_request_sample_does_not_trust_caller_tags( use litellm_traces::{LensQuery, Parameter}; let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64 / 1_000_000; for (id, internal) in [("external", false), ("internal", true)] { let row = serde_json::from_value(serde_json::json!({ @@ -755,7 +749,6 @@ async fn lens_selection_pages_without_losing_or_repeating_runs( &Connection::writer(&database.url)?, "trace_test", 7, - 14, ) .await?; execute_write(&database, "INSERT INTO trace_test.spend_logs (request_id,team_id,start_time,end_time) SELECT toString(number),'team',now64(3)-INTERVAL 5 MINUTE,now64(3)-INTERVAL 5 MINUTE FROM numbers(1001)").await?; @@ -836,7 +829,6 @@ async fn lens_content_keeps_output_visible_after_long_input( &Connection::writer(&database.url)?, "trace_test", 7, - 14, ) .await?; insert_rows(&database, "spend_logs", vec![serde_json::from_value(serde_json::json!({ @@ -902,7 +894,7 @@ async fn trace_error_previews_preserve_paginated_diagnostics( ) -> TestResult { let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; let rows = (0..span_count) .map(|index| { @@ -992,7 +984,7 @@ async fn duplicate_span_preview_matches_diagnostic( ) -> TestResult { let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; let message = "a".repeat(200); let rows = [ @@ -1041,7 +1033,7 @@ fn schema_includes_every_migration_file() -> TestResult { .filter_map(|entry| entry.ok()) .filter(|entry| entry.path().extension().is_some_and(|ext| ext == "sql")) .count(); - assert_eq!(schema_statements("trace_test", 7, 14)?.len(), 1 + files); + assert_eq!(schema_statements("trace_test", 7)?.len(), 1 + files); Ok(()) } @@ -1053,7 +1045,7 @@ async fn lens_agent_discovery_and_selection_preserve_scope( use litellm_traces::LensQuery; let database = database.await?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; for (team, key, trace, agent, span, parent) in [ ("alpha", "one", "research", "research_agent", "root", ""), @@ -1160,7 +1152,7 @@ async fn query_help_discovers_live_schema_and_runs_its_examples( ) -> TestResult { let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; execute_write(&database, "CREATE USER help_reader").await?; for table in ["otel_traces", "agent_traces_by_key", "spend_logs"] { execute_write( @@ -1371,7 +1363,7 @@ async fn query_help_preserves_schema_and_guide_when_discovery_hits_reader_limits ) -> TestResult { let database = database?; let writer = Connection::writer(&database.url)?; - ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; execute_write( &database, "CREATE USER help_reader SETTINGS max_rows_to_read = 1", diff --git a/litellm-rust/crates/traces/tests/query_access.rs b/litellm-rust/crates/traces/tests/query_access.rs index d6f44075e3b..b71bcce614e 100644 --- a/litellm-rust/crates/traces/tests/query_access.rs +++ b/litellm-rust/crates/traces/tests/query_access.rs @@ -33,7 +33,7 @@ async fn database() -> Result> { container.get_host_port_ipv4(8123).await? ))?; let client = Client::no_redirect_for_test(); - ensure_schema(&client, &writer, "trace_test", 7, 7).await?; + ensure_schema(&client, &writer, "trace_test", 7).await?; for sql in [ "INSERT INTO trace_test.otel_traces (TeamId, ApiKeyHash, TraceId, SpanId, Timestamp, SpanAttributes) VALUES ('team-a', 'key-a1', 'shared-trace', 'a1', now(), map('visible', 'a')), ('team-a', 'key-a2', 'shared-trace', 'a2', now(), map('visible', 'a')), ('team-b', 'key-b', 'shared-trace', 'b', now(), map('secret-b', 'b'))", "INSERT INTO trace_test.spend_logs (team_id, api_key, request_id, start_time, end_time, metadata) VALUES ('team-a', 'key-a1', 'a1', now(), now(), '{\"visible\":1}'), ('team-a', 'key-a2', 'a2', now(), now(), '{\"visible\":1}'), ('team-b', 'key-b', 'b', now(), now(), '{\"secret_b\":1}')", diff --git a/litellm/constants.py b/litellm/constants.py index af4d1268c03..18fc6aa7e74 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -50,8 +50,8 @@ CLICKHOUSE_BATCH_SIZE: Final = get_env_int("CLICKHOUSE_BATCH_SIZE", 10_000) CLICKHOUSE_FLUSH_INTERVAL_SECONDS: Final = float(os.getenv("CLICKHOUSE_FLUSH_INTERVAL_SECONDS", "1.0")) CLICKHOUSE_MAX_BUFFERED_ROWS: Final = get_env_int("CLICKHOUSE_MAX_BUFFERED_ROWS", 200_000) 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) +DEFAULT_CLICKHOUSE_DATABASE: Final = "litellm" +DEFAULT_AGENT_TRACING_RETENTION_DAYS: Final = 14 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) diff --git a/litellm/integrations/clickhouse/clickhouse_batch_logger.py b/litellm/integrations/clickhouse/clickhouse_batch_logger.py index ac782ffebb2..2290e6520b0 100644 --- a/litellm/integrations/clickhouse/clickhouse_batch_logger.py +++ b/litellm/integrations/clickhouse/clickhouse_batch_logger.py @@ -9,7 +9,6 @@ gzip JSONEachRow insert, either every `CLICKHOUSE_FLUSH_INTERVAL_SECONDS` or as """ import asyncio -import os from collections.abc import Mapping, Sequence from contextlib import suppress from typing import Any, ClassVar, Final @@ -23,13 +22,11 @@ from litellm.constants import ( ) from litellm.integrations.custom_batch_logger import CustomBatchLogger from litellm.rust_bridge.traces import ClickHouseStorage +from litellm.tracing.config import trace_storage_config def clickhouse_storage_from_env() -> ClickHouseStorage: - return ClickHouseStorage( - database=os.getenv("CLICKHOUSE_DATABASE", "litellm"), - url=os.getenv("CLICKHOUSE_URL", ""), - ) + return ClickHouseStorage(trace_storage_config({})) class ClickHouseBatchLogger(CustomBatchLogger): diff --git a/litellm/integrations/clickhouse/schema.py b/litellm/integrations/clickhouse/schema.py index 5bf2b21cda5..adf538b0f6f 100644 --- a/litellm/integrations/clickhouse/schema.py +++ b/litellm/integrations/clickhouse/schema.py @@ -7,5 +7,5 @@ AGENT_TRACES_BY_KEY_TABLE: Final = "agent_traces_by_key" SPEND_LOGS_TABLE: Final = "spend_logs" -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) +async def ensure_schema(storage: ClickHouseStorage) -> None: + await storage.ensure_schema() diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 21570142500..5d47cceee02 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -858,6 +858,7 @@ from litellm.secret_managers.main import ( secret_manager_would_be_consulted, str_to_bool, ) +from litellm.tracing.config import is_clickhouse_tracing_enabled from litellm.types.integrations.slack_alerting import AlertType, SlackAlertingArgs from litellm.types.llms.anthropic import ( AnthropicMessagesRequest, @@ -1567,11 +1568,15 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState register_scheduled_sync(scheduler) - tracing_settings: Final = general_settings.get("tracing") - tracing_enabled: Final = TypeAdapter(bool).validate_python( - isinstance(tracing_settings, dict) and tracing_settings.get("store") == "clickhouse" + tracing_settings: Final = cast( # cast-ok: Pydantic validates the legacy untyped settings value + dict[str, object] | None, + TypeAdapter(dict[str, object] | None).validate_python(general_settings.get("tracing")), ) - async with manage_tracing(enabled=tracing_enabled) as receiver: + tracing_enabled: Final = is_clickhouse_tracing_enabled(tracing_settings) + async with manage_tracing( + enabled=tracing_enabled, + settings=tracing_settings, + ) as receiver: state: Final[ProxyLifespanState] = {"tracing_receiver": receiver} yield state diff --git a/litellm/proxy/tracing_runtime.py b/litellm/proxy/tracing_runtime.py index 0b706d66a40..730a620cf70 100644 --- a/litellm/proxy/tracing_runtime.py +++ b/litellm/proxy/tracing_runtime.py @@ -1,4 +1,4 @@ -from collections.abc import AsyncGenerator, Callable +from collections.abc import AsyncGenerator, Callable, Mapping from contextlib import asynccontextmanager from typing import Final @@ -44,9 +44,12 @@ async def _start_receiver(factory: Callable[[], TraceReceiver]) -> TraceReceiver @asynccontextmanager async def manage_tracing( - enabled: bool, receiver_factory: Callable[[], TraceReceiver] = TraceReceiver.from_env + enabled: bool, + receiver_factory: Callable[[], TraceReceiver] | None = None, + settings: Mapping[str, object] | None = None, ) -> AsyncGenerator[TraceReceiver | None, None]: - tracing: Final = await _start_receiver(receiver_factory) if enabled else None + factory: Final = receiver_factory or (lambda: TraceReceiver.from_settings(settings or {})) + tracing: Final = await _start_receiver(factory) if enabled else None if tracing is None: yield tracing return diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 5d23b048e86..0f70c559840 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -25,10 +25,19 @@ def trace_decode_otlp(body: bytes, content_type: str | None) -> list[DecodedSpan def trace_encode_error(message: str) -> bytes: ... def trace_normalized_field_definitions() -> list[dict[str, str]]: ... +@final +class NativeTraceConfig: + def __new__( + cls, + database: str, + url: str, + retention_days: int, + ) -> NativeTraceConfig: ... + @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 __new__(cls, config: NativeTraceConfig) -> NativeTraceStorage: ... + def ensure_schema(self) -> Future[None]: ... def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> Future[None]: ... def query_sql(self, sql: str, scope: QueryScope, secret: str) -> Future[str]: ... def query_help(self, scope: QueryScope, secret: str) -> Future[str]: ... @@ -329,6 +338,7 @@ __all__ = [ "ForkedAfterNativeRuntimeStarted", "HuggingFaceEncoding", "NativeDiagnosticProcessor", + "NativeTraceConfig", "NativeTraceStorage", "ProcessReservedForForking", "ResponsesWebSocketConnection", diff --git a/litellm/rust_bridge/traces.py b/litellm/rust_bridge/traces.py index afc1620a480..1e579edb73e 100644 --- a/litellm/rust_bridge/traces.py +++ b/litellm/rust_bridge/traces.py @@ -1,4 +1,5 @@ from collections.abc import Awaitable, Mapping, Sequence +from dataclasses import dataclass from types import MappingProxyType from typing import Final, Literal, Protocol, TypedDict, cast @@ -77,9 +78,9 @@ QueryScope = AdminQueryScope | TeamQueryScope | KeyQueryScope class NativeStore(Protocol): - def __init__(self, database: str, url: str, reader_url: str | None = None) -> None: ... + def __init__(self, config: "NativeConfig") -> None: ... - def ensure_schema(self, trace_retention_days: int, spend_log_retention_days: int) -> Awaitable[None]: ... + def ensure_schema(self) -> Awaitable[None]: ... def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> Awaitable[None]: ... @@ -93,6 +94,7 @@ class NativeStore(Protocol): class NativeTraces(Protocol): + NativeTraceConfig: type["NativeConfig"] NativeTraceStorage: type[NativeStore] def trace_decode_otlp( @@ -115,6 +117,17 @@ QUERY_PARAMETERS: Final = TypeAdapter(dict[str, str | int | list[str]]) _FIELD_DEFINITIONS_ADAPTER: Final = TypeAdapter(tuple[NormalizedFieldDefinition, ...]) +class NativeConfig(Protocol): + def __init__(self, database: str, url: str, retention_days: int) -> None: ... + + +@dataclass(frozen=True, slots=True, repr=False) +class TraceStorageConfig: + url: str + database: str = "litellm" + retention_days: int = 14 + + def _native() -> NativeTraces: native: Final = get_native_bridge() if native is None: @@ -143,11 +156,17 @@ def encode_error(message: str) -> bytes: class ClickHouseStorage: - def __init__(self, database: str, url: str, reader_url: str | None = None) -> None: - self._native: Final = _native().NativeTraceStorage(database, url, reader_url) + def __init__(self, config: TraceStorageConfig) -> None: + native: Final = _native() + validated: Final = native.NativeTraceConfig( + config.database, + config.url, + config.retention_days, + ) + self._native: Final = native.NativeTraceStorage(validated) - async def ensure_schema(self, trace_retention_days: int, spend_log_retention_days: int) -> None: - await self._native.ensure_schema(trace_retention_days, spend_log_retention_days) + async def ensure_schema(self) -> None: + await self._native.ensure_schema() async def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> None: await self._native.insert_rows(table, rows) diff --git a/litellm/tracing/config.py b/litellm/tracing/config.py new file mode 100644 index 00000000000..16f6b0a5975 --- /dev/null +++ b/litellm/tracing/config.py @@ -0,0 +1,84 @@ +import os +from collections.abc import Mapping +from typing import Final + +from pydantic import TypeAdapter + +from litellm.constants import DEFAULT_AGENT_TRACING_RETENTION_DAYS, DEFAULT_CLICKHOUSE_DATABASE +from litellm.rust_bridge.traces import TraceStorageConfig + +STORE_SETTINGS: Final = TypeAdapter(dict[str, object]) + + +def is_clickhouse_tracing_enabled(settings: object) -> bool: + if not isinstance(settings, Mapping): + return False + typed_settings: Final = STORE_SETTINGS.validate_python(settings) + store: Final = typed_settings.get("store") + if not isinstance(store, Mapping): + return False + return STORE_SETTINGS.validate_python(store).get("type") == "clickhouse" + + +def _value(settings: Mapping[str, object], field: str, environ: Mapping[str, str], default: object) -> object: + if field not in settings: + return default + supplied: Final = settings[field] + resolved: Final = ( + environ.get(supplied.removeprefix("os.environ/")) + if isinstance(supplied, str) and supplied.startswith("os.environ/") + else supplied + ) + if resolved is None: + raise ValueError(f"tracing.store.{field} is set but resolved to no value") + return resolved + + +def _retention_days(value: object) -> int: + if isinstance(value, bool) or not isinstance(value, (int, str)): + raise ValueError("tracing.store.retention_days must be a positive integer") + try: + days: Final = int(value) + except ValueError as error: + raise ValueError("tracing.store.retention_days must be a positive integer") from error + if not 0 < days <= 2**32 - 1: + raise ValueError("tracing.store.retention_days must be a positive integer") + return days + + +def _clickhouse_store(settings: Mapping[str, object]) -> Mapping[str, object]: + raw_store: Final = settings.get("store") + if raw_store is None: + return {} + if isinstance(raw_store, Mapping): + store: Final = STORE_SETTINGS.validate_python(raw_store) + if store.get("type") == "clickhouse": + return store + raise ValueError("tracing.store.type must be clickhouse") + + +def trace_storage_config(settings: Mapping[str, object], environ: Mapping[str, str] = os.environ) -> TraceStorageConfig: + store: Final = _clickhouse_store(settings) + unknown: Final = store.keys() - {"type", "url", "database", "retention_days"} + if unknown: + raise ValueError(f"unsupported tracing.store settings: {', '.join(sorted(unknown))}") + url: Final = _value(store, "url", environ, environ.get("CLICKHOUSE_URL")) + database: Final = _value( + store, "database", environ, environ.get("CLICKHOUSE_DATABASE", DEFAULT_CLICKHOUSE_DATABASE) + ) + if not isinstance(url, str) or not url: + raise ValueError("tracing.store.url or CLICKHOUSE_URL is required") + if not isinstance(database, str): + raise ValueError("tracing.store.database must be a string") + return TraceStorageConfig( + url=url, + database=database, + retention_days=_retention_days( + _value( + store, + "retention_days", + environ, + environ.get("AGENT_TRACING_RETENTION_DAYS", DEFAULT_AGENT_TRACING_RETENTION_DAYS), + ) + ), + ) diff --git a/litellm/tracing/receiver.py b/litellm/tracing/receiver.py index 1fae2361572..9164e0722a3 100644 --- a/litellm/tracing/receiver.py +++ b/litellm/tracing/receiver.py @@ -13,21 +13,15 @@ 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_MAX_CONCURRENT_INGESTS, -) -from litellm.integrations.clickhouse.schema import ensure_schema +from litellm.constants import OTLP_MAX_BODY_BYTES, OTLP_MAX_CONCURRENT_INGESTS from litellm.rust_bridge.traces import ClickHouseStorage +from litellm.tracing.config import trace_storage_config from litellm.tracing.decode import OTLPPayloadTooLargeError, decode_otlp from litellm.tracing.store import TraceStore from litellm.tracing.types import ( @@ -103,22 +97,14 @@ class TraceReceiver: @classmethod def from_env(cls) -> "TraceReceiver": - return cls( - store=TraceStore( - ClickHouseStorage( - database=os.getenv("CLICKHOUSE_DATABASE", "litellm"), - url=os.environ["CLICKHOUSE_URL"], - reader_url=os.getenv("CLICKHOUSE_READER_URL", os.environ["CLICKHOUSE_URL"]), - ) - ) - ) + return cls.from_settings({}) + + @classmethod + def from_settings(cls, settings: Mapping[str, object]) -> "TraceReceiver": + return cls(store=TraceStore(ClickHouseStorage(trace_storage_config(settings)))) async def start(self) -> None: - await ensure_schema( - self.store.storage, - trace_retention_days=AGENT_TRACING_RETENTION_DAYS, - spend_log_retention_days=AGENT_TRACING_SPEND_LOG_RETENTION_DAYS, - ) + await self.store.storage.ensure_schema() async def ingest( self, diff --git a/scripts/run_tracing_proxy_local.sh b/scripts/run_tracing_proxy_local.sh index fd48590bf93..8c44d4c783d 100755 --- a/scripts/run_tracing_proxy_local.sh +++ b/scripts/run_tracing_proxy_local.sh @@ -13,11 +13,18 @@ VIRTUAL_ENV="$repo_root/.venv" uvx --from maturin==1.15.0 maturin develop \ config_file="$(mktemp "${TMPDIR:-/tmp}/litellm-tracing-local.XXXXXX.yaml")" trap 'rm -f "$config_file"' EXIT cat > "$config_file" <<'EOF' -model_list: [] +model_list: + - model_name: claude-sonnet + litellm_params: + model: anthropic/claude-sonnet-5-5 + api_key: os.environ/ANTHROPIC_API_KEY general_settings: master_key: os.environ/LITELLM_MASTER_KEY tracing: - store: clickhouse + store: + type: clickhouse + url: os.environ/CLICKHOUSE_URL + retention_days: 14 EOF export LITELLM_MASTER_KEY=sk-local-tracing @@ -25,10 +32,16 @@ export LITELLM_SALT_KEY=sk-local-tracing-salt-key export DATABASE_URL=postgresql://litellm:litellm@127.0.0.1:15432/litellm export STORE_MODEL_IN_DB=True export CLICKHOUSE_URL=http://default:local-tracing@127.0.0.1:18123 -export CLICKHOUSE_READER_URL="$CLICKHOUSE_URL" export CLICKHOUSE_DATABASE=litellm export LITELLM_LOCAL_MODEL_COST_MAP=True -printf 'Proxy: http://127.0.0.1:4002/ui\nMaster key: %s\n' "$LITELLM_MASTER_KEY" +( + cd "$repo_root/ui/litellm-dashboard" + "$repo_root/scripts/with_dashboard_node.sh" npm ci + NEXT_PUBLIC_BASE_URL= "$repo_root/scripts/with_dashboard_node.sh" npm run build +) +export LITELLM_UI_PATH="$repo_root/ui/litellm-dashboard/out" + +printf 'Dashboard: http://127.0.0.1:4002/ui/\nProxy: http://127.0.0.1:4002\nMaster key: %s\n' "$LITELLM_MASTER_KEY" "$repo_root/.venv/bin/python" litellm/proxy/proxy_cli.py \ --config "$config_file" --host 127.0.0.1 --port 4002 diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index 5a09578661e..efc0a44f311 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -8,21 +8,31 @@ from urllib.parse import parse_qs, urlsplit import pytest -from litellm.rust_bridge._native import NativeTraceStorage, trace_decode_otlp -from litellm.rust_bridge.traces import ClickHouseStorage, NormalizedSpan, normalized_field_definitions +from litellm.rust_bridge._native import NativeTraceConfig, NativeTraceStorage, trace_decode_otlp +from litellm.rust_bridge.traces import ( + ClickHouseStorage, + NormalizedSpan, + TraceStorageConfig, + normalized_field_definitions, +) from litellm.tracing import Tenant, TraceReceiver, TracingPayloadTooLargeError from litellm.tracing.decode import decode_otlp from litellm.tracing.store import TraceStore +from litellm.tracing.types import TraceScope from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec pytestmark = pytest.mark.requires_rust_extension +def _native_storage(database: str, url: str, retention_days: int = 14) -> NativeTraceStorage: + return NativeTraceStorage(NativeTraceConfig(database, url, retention_days)) + + @pytest.mark.asyncio async def test_trace_reader_projects_connection_and_parameters(recording_server: RecordingServer) -> None: recording_server.enqueue(ResponseSpec(body={"data": [{"trace_id": "trace-1"}]})) - reader_url: Final = recording_server.base_url.replace("http://", "http://reader:p%40ss%2Fword%25@") - storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, reader_url + "?database=wrong") + url: Final = recording_server.base_url.replace("http://", "http://reader:p%40ss%2Fword%25@") + storage: Final = _native_storage("trace_test", url + "?database=wrong") rows: Final = json.loads(await storage.query("trace_spans", {"trace_id": "trace-1"})) request: Final = recording_server.requests[0] parameters: Final = parse_qs(urlsplit(request.path).query) @@ -39,7 +49,7 @@ async def test_trace_reader_projects_connection_and_parameters(recording_server: @pytest.mark.asyncio async def test_trace_reader_rejects_success_status_with_embedded_error(recording_server: RecordingServer) -> None: recording_server.enqueue(ResponseSpec(body={"data": [], "exception": "query failed"})) - storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, recording_server.base_url) + storage: Final = _native_storage("trace_test", recording_server.base_url) with pytest.raises(RuntimeError, match="invalid or failed JSON"): await storage.query("trace_spans", {}) @@ -47,7 +57,7 @@ async def test_trace_reader_rejects_success_status_with_embedded_error(recording @pytest.mark.asyncio async def test_reader_rejects_arbitrary_sql_before_sending(recording_server: RecordingServer) -> None: recording_server.expected_requests = 0 - storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, recording_server.base_url) + storage: Final = _native_storage("trace_test", recording_server.base_url) with pytest.raises(ValueError, match="unknown ClickHouse read query"): await storage.query("SELECT 1", {}) @@ -55,14 +65,44 @@ async def test_reader_rejects_arbitrary_sql_before_sending(recording_server: Rec @pytest.mark.asyncio async def test_schema_binding_rejects_invalid_database() -> None: with pytest.raises(ValueError, match=r"database.*retention"): - NativeTraceStorage("db; DROP DATABASE default", "http://localhost:8123") + NativeTraceConfig("db; DROP DATABASE default", "http://localhost:8123", 14) @pytest.mark.asyncio async def test_schema_binding_rejects_non_positive_retention() -> None: - storage: Final = NativeTraceStorage("traces", "http://localhost:8123") with pytest.raises(ValueError, match=r"database.*retention"): - await storage.ensure_schema(0, 14) + NativeTraceConfig("traces", "http://localhost:8123", 0) + + +def test_invalid_url_error_does_not_expose_credentials() -> None: + with pytest.raises(RuntimeError, match="invalid ClickHouse HTTP URL") as error: + NativeTraceConfig("traces", "secret://writer:password@example.com", 7) + assert "password" not in str(error.value) + + +@pytest.mark.asyncio +async def test_from_env_reads_with_clickhouse_url( + recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch +) -> None: + recording_server.enqueue(ResponseSpec(body={"data": []})) + monkeypatch.setenv("CLICKHOUSE_URL", recording_server.base_url) + monkeypatch.delenv("CLICKHOUSE_READER_URL", raising=False) + scope: Final[TraceScope] = {"team_ids": (), "api_key_hash": ""} + page: Final = await TraceReceiver.from_env().list_traces(scope, 0, 1) + assert page == {"data": (), "next_cursor": None} + assert len(recording_server.requests) == 1 + + +@pytest.mark.asyncio +async def test_schema_setup_uses_configured_retention(recording_server: RecordingServer) -> None: + recording_server.expected_requests = 8 + storage: Final = _native_storage("trace_test", recording_server.base_url, 7) + await storage.ensure_schema() + ttl_statements: Final = tuple( + request.raw_body for request in recording_server.requests if b"MODIFY TTL" in request.raw_body + ) + assert len(ttl_statements) == 3 + assert all(b"INTERVAL 7 DAY" in statement for statement in ttl_statements) @pytest.mark.asyncio @@ -73,9 +113,9 @@ async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement recording_server.enqueue(ResponseSpec(body="")) recording_server.enqueue(ResponseSpec(status=403, body="denied")) writer_url: Final = recording_server.base_url.replace("http://", "http://writer:p%40ss%2Fword%25@") - storage: Final = NativeTraceStorage("trace_test", writer_url + "?database=wrong&readonly=1") + storage: Final = _native_storage("trace_test", writer_url + "?database=wrong&readonly=1", 7) with pytest.raises(RuntimeError, match="schema setup failed with HTTP status 403"): - await storage.ensure_schema(7, 14) + await storage.ensure_schema() assert len(recording_server.requests) == 2 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") @@ -89,7 +129,7 @@ async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement @pytest.mark.asyncio async def test_insert_encodes_and_sends_rows(recording_server: RecordingServer) -> None: recording_server.enqueue(ResponseSpec(body="")) - storage: Final = NativeTraceStorage("trace_test", recording_server.base_url) + storage: Final = _native_storage("trace_test", recording_server.base_url) before: Final = time.time_ns() // 1_000_000 await storage.insert_rows("otel_traces", [{"Timestamp": 1_234_567_890, "Input": "hello", "EngineReceivedMs": -1}]) after: Final = time.time_ns() // 1_000_000 @@ -169,7 +209,9 @@ def test_normalized_field_contract_matches_decoded_rust_span() -> None: @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))) + receiver: Final = TraceReceiver( + TraceStore(ClickHouseStorage(TraceStorageConfig(recording_server.base_url, "trace_test"))) + ) 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) @@ -186,7 +228,9 @@ async def test_resource_fanout_reaches_insert_with_identical_values(recording_se 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))) + receiver: Final = TraceReceiver( + TraceStore(ClickHouseStorage(TraceStorageConfig(recording_server.base_url, "trace_test"))) + ) with pytest.raises(TracingPayloadTooLargeError, match="encoded size limit"): await receiver.ingest(body, "application/json", None, Tenant("team-a", "key-a")) assert recording_server.requests == [] @@ -194,7 +238,7 @@ async def test_shared_resource_still_hits_insert_limit_before_transport(recordin @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) + storage: Final = ClickHouseStorage(TraceStorageConfig(recording_server.base_url, "trace_test")) invalid: Final = object() with pytest.raises(ValueError, match=type(invalid).__name__): await storage.insert_rows("otel_traces", [{"ResourceAttributes": invalid}]) @@ -225,7 +269,7 @@ def test_trace_sql_endpoint_executes_for_admin_and_preserves_clickhouse_envelope for _ in range(11): recording_server.enqueue(ResponseSpec(body="")) recording_server.enqueue(ResponseSpec(body=envelope)) - storage: Final = ClickHouseStorage("trace_test", recording_server.base_url, recording_server.base_url) + storage: Final = ClickHouseStorage(TraceStorageConfig(recording_server.base_url, "trace_test")) app: Final = FastAPI() app.include_router(router) app.dependency_overrides[provide_trace_query_secret] = lambda: "test-master-secret" @@ -260,7 +304,7 @@ def test_trace_help_endpoint_runs_native_schema_and_metadata_discovery(recording {"data": [{"key": "custom.resource"}]}, ): recording_server.enqueue(ResponseSpec(body=response)) - storage: Final = ClickHouseStorage("trace_test", recording_server.base_url, recording_server.base_url) + storage: Final = ClickHouseStorage(TraceStorageConfig(recording_server.base_url, "trace_test")) app: Final = FastAPI() app.include_router(router) app.dependency_overrides[provide_trace_query_secret] = lambda: "test-master-secret" @@ -299,7 +343,7 @@ def test_trace_sql_endpoint_distinguishes_query_errors_from_reader_failures( recording_server.enqueue(ResponseSpec(status=clickhouse_status, body=b"ClickHouse rejected the query")) envelope: Final = {"meta": [{"name": "answer", "type": "UInt8"}], "data": [{"answer": 42}], "rows": 1} recording_server.enqueue(ResponseSpec(body=envelope)) - storage: Final = ClickHouseStorage("trace_test", recording_server.base_url, recording_server.base_url) + storage: Final = ClickHouseStorage(TraceStorageConfig(recording_server.base_url, "trace_test")) app: Final = FastAPI() app.include_router(router) app.dependency_overrides[provide_trace_query_secret] = lambda: "test-master-secret" diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index 9785bdd5e32..51ecc210119 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -40,10 +40,34 @@ from litellm.proxy.proxy_server import ( validate_deployment_complexity_router_placement, validate_deployment_max_agentic_loops, ) +from litellm.tracing.config import trace_storage_config from .conftest import normalize +@pytest.mark.asyncio +async def test_proxy_config_loads_tracing_url_and_retention_from_yaml(tmp_path, monkeypatch) -> None: + config_file: Final = tmp_path / "tracing.yaml" + config_file.write_text( + "model_list: []\ngeneral_settings:\n tracing:\n store:\n" + " type: clickhouse\n url: os.environ/TRACING_TEST_URL\n" + " database: analytics\n retention_days: 7\n" + ) + monkeypatch.setenv("TRACING_TEST_URL", "http://localhost:8123") + monkeypatch.setenv("CLICKHOUSE_URL", "http://unused:8123") + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + + _, _, settings = await ProxyConfig().load_config(router=None, config_file_path=str(config_file)) + tracing = trace_storage_config(settings["tracing"]) + assert (tracing.url, tracing.database, tracing.retention_days) == ( + "http://localhost:8123", + "analytics", + 7, + ) + + @pytest.mark.asyncio @pytest.mark.parametrize("shutdown_error", [False, True]) async def test_tracing_config_automatically_logs_spend_without_callback_setting(shutdown_error: bool) -> None: diff --git a/tests/unit/tracing/__init__.py b/tests/unit/tracing/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/tracing/test_config.py b/tests/unit/tracing/test_config.py new file mode 100644 index 00000000000..030c4247c62 --- /dev/null +++ b/tests/unit/tracing/test_config.py @@ -0,0 +1,116 @@ +import pytest + +from litellm import constants +from litellm.tracing.config import is_clickhouse_tracing_enabled, trace_storage_config + + +@pytest.mark.parametrize( + ("settings", "enabled"), + [ + ({"store": "clickhouse"}, False), + ({"store": {"type": "clickhouse"}}, True), + ({"store": {"type": "other"}}, False), + (None, False), + ], +) +def test_clickhouse_tracing_enablement(settings: object, enabled: bool) -> None: + assert is_clickhouse_tracing_enabled(settings) is enabled + + +def test_yaml_values_override_defaults_and_resolve_nested_references() -> None: + config = trace_storage_config( + { + "store": { + "type": "clickhouse", + "url": "os.environ/TRACING_URL", + "database": "os.environ/TRACING_DATABASE", + "retention_days": "os.environ/TRACING_RETENTION_DAYS", + }, + }, + { + "TRACING_URL": "https://writer:password@clickhouse.example:8443", + "TRACING_DATABASE": "analytics", + "TRACING_RETENTION_DAYS": "7", + "CLICKHOUSE_URL": "https://other.example:8443", + }, + ) + assert config.url == "https://writer:password@clickhouse.example:8443" + assert config.database == "analytics" + assert config.retention_days == 7 + assert "password" not in repr(config) + + +def test_omitted_fields_use_environment() -> None: + config = trace_storage_config( + {}, + { + "CLICKHOUSE_URL": "http://localhost:8123", + "CLICKHOUSE_DATABASE": "env_database", + "AGENT_TRACING_RETENTION_DAYS": "11", + }, + ) + assert (config.url, config.database, config.retention_days) == ("http://localhost:8123", "env_database", 11) + + +def test_environment_is_read_when_config_is_resolved(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("CLICKHOUSE_URL", "http://localhost:8123") + monkeypatch.setenv("CLICKHOUSE_DATABASE", "late_database") + monkeypatch.setenv("AGENT_TRACING_RETENTION_DAYS", "9") + config = trace_storage_config({}) + assert (config.database, config.retention_days) == ("late_database", 9) + + +def test_omitted_fields_without_environment_use_constant_defaults() -> None: + config = trace_storage_config({}, {"CLICKHOUSE_URL": "http://localhost:8123"}) + assert (config.database, config.retention_days) == ( + constants.DEFAULT_CLICKHOUSE_DATABASE, + constants.DEFAULT_AGENT_TRACING_RETENTION_DAYS, + ) + assert (config.database, config.retention_days) == ("litellm", 14) + + +@pytest.mark.parametrize("field", ["url", "database", "retention_days"]) +def test_unset_environment_reference_does_not_fall_back(field: str) -> None: + store: dict[str, object] = {"type": "clickhouse", "url": "http://localhost:8123", field: "os.environ/MISSING"} + with pytest.raises(ValueError, match=rf"tracing.store.{field} is set but resolved to no value") as error: + trace_storage_config({"store": store}, {"CLICKHOUSE_URL": "http://fallback:8123"}) + assert "MISSING" not in str(error.value) + + +@pytest.mark.parametrize("store", ["clickhouse", {"type": "other"}]) +def test_non_clickhouse_store_is_rejected(store: object) -> None: + with pytest.raises(ValueError, match=r"tracing\.store\.type must be clickhouse"): + trace_storage_config({"store": store}, {"CLICKHOUSE_URL": "http://localhost:8123"}) + + +def test_non_string_database_is_rejected() -> None: + with pytest.raises(ValueError, match=r"tracing\.store\.database must be a string"): + trace_storage_config({"store": {"type": "clickhouse", "url": "http://localhost:8123", "database": 1}}, {}) + + +@pytest.mark.parametrize("value", [0, -1, True, "not-a-number", 2**32]) +def test_invalid_retention_is_rejected(value: object) -> None: + with pytest.raises(ValueError, match=r"tracing.store.retention_days must be a positive integer"): + trace_storage_config( + {"store": {"type": "clickhouse", "url": "http://localhost:8123", "retention_days": value}}, {} + ) + + +def test_missing_url_is_rejected() -> None: + with pytest.raises(ValueError, match=r"tracing.store.url or CLICKHOUSE_URL is required"): + trace_storage_config({"store": {"type": "clickhouse"}}, {}) + + +def test_legacy_reader_and_split_retention_fields_are_rejected() -> None: + with pytest.raises(ValueError, match="reader_url, trace_retention_days"): + trace_storage_config( + { + "store": { + "type": "clickhouse", + "url": "http://localhost:8123", + "reader_url": "http://localhost:8124", + "trace_retention_days": 30, + } + }, + {}, + ) diff --git a/ui/litellm-dashboard/src/components/view_logs/AuditLogsPanel.tsx b/ui/litellm-dashboard/src/components/view_logs/AuditLogsPanel.tsx index 5ce4ca6053f..abc00b75884 100644 --- a/ui/litellm-dashboard/src/components/view_logs/AuditLogsPanel.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/AuditLogsPanel.tsx @@ -53,7 +53,8 @@ export default function AuditLogsPanel({ return typeof entry?.value === "string" && entry.value.trim() ? entry.value.trim() : undefined; }; - const canQueryAuditLogs = !!accessToken && !!token && !!userRole && !!userID && isActive && premiumUser; + const hasSession = [accessToken, token, userRole, userID].every(Boolean); + const canQueryAuditLogs = hasSession && isActive && premiumUser; const query = useQuery({ queryKey: ["audit_logs", pagination.pageIndex, pagination.pageSize, columnFilters, searchTerm], diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx index 595bd7684d5..5db34ef9788 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx @@ -116,7 +116,8 @@ describe("AgentTracesSection", () => { const card = await screen.findByTestId("tracing-setup-card"); expect(card).toHaveTextContent("Tracing is not enabled"); - expect(card).toHaveTextContent("store: clickhouse"); + expect(card).toHaveTextContent("type: clickhouse"); + expect(card).toHaveTextContent("url: os.environ/CLICKHOUSE_URL"); expect(screen.getByRole("button", { name: "Check setup" })).toBeEnabled(); expect(card).not.toHaveTextContent(/langsmith/i); expect(card).toHaveTextContent("ClickHouse and proxy setup"); @@ -225,7 +226,7 @@ describe("AgentTracesSection", () => { const card = await screen.findByTestId("tracing-setup-card"); expect(card).toHaveTextContent("Tracing is not enabled"); - expect(card).toHaveTextContent("CLICKHOUSE_READER_URL"); + expect(card).toHaveTextContent("url: os.environ/CLICKHOUSE_URL"); }); it("lists every run with its input, counts and failed column", async () => { diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.test.tsx index 027870ab1ae..65e5e9229ff 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.test.tsx @@ -182,7 +182,8 @@ describe("TracingSetupCard", () => { const onCheck = vi.fn(); const { card } = renderCard({ detail: "Agent tracing is not enabled", onCheck }); expect(screen.getByRole("heading", { name: "Enable tracing" })).toBeVisible(); - expect(card).toHaveTextContent("store: clickhouse"); + expect(card).toHaveTextContent("type: clickhouse"); + expect(card).toHaveTextContent("url: os.environ/CLICKHOUSE_URL"); expect(screen.queryByRole("combobox", { name: "Your agent framework" })).not.toBeInTheDocument(); expect(screen.queryByRole("button", { name: "Send a test trace" })).not.toBeInTheDocument(); await user.click(screen.getByRole("button", { name: "Check setup" })); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.tsx index 5ac5c988efe..de776da0617 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.tsx @@ -231,9 +231,10 @@ export const otlpEndpoints = (proxyUrl: string): readonly (readonly [string, str export const PROXY_CONFIG_SNIPPET = [ "general_settings:", " tracing:", - " store: clickhouse", - "", - "# env: CLICKHOUSE_URL (writer) and CLICKHOUSE_READER_URL (read-only user)", + " store:", + " type: clickhouse", + " url: os.environ/CLICKHOUSE_URL", + " retention_days: 14", ].join("\n"); function CodeBlock({ @@ -575,8 +576,8 @@ function EnableTracing({ checked, checking, onCheck }: { checked: boolean; check <>

- Set your ClickHouse writer and read-only reader URLs, add this to config.yaml, then restart the proxy. Ask - your proxy administrator if you don’t manage this deployment. + Set your ClickHouse URL, add this to config.yaml, then restart the proxy. Ask your proxy administrator if you + don’t manage this deployment.

config.yaml} /> Date: Fri, 2 Oct 2026 09:37:37 -0700 Subject: [PATCH 004/139] feat(docker): one-command quickstart that starts the gateway, Postgres, and the admin UI (#43673) * feat(docker): one-command quickstart that starts the gateway, Postgres, and the admin UI scripts/quickstart.sh downloads the quickstart compose file into ~/litellm-gateway, generates the master key, salt key, and a random Postgres password into .env, picks a free port, starts the stack, waits for it to be ready, and prints where to log in. It asks at most two questions (install folder, open the browser) and asks nothing without a terminal, under CI or Claude Code, or with --yes The compose file reads POSTGRES_PASSWORD and LITELLM_PORT from .env and falls back to the current values, so existing installs keep working unchanged * fix(quickstart): address review findings on reinstall, ports, gitignore, and binding Stop with instructions instead of generating a new password when a database volume from an earlier install is still there, since Postgres keeps the original password Keep port 4000 for an existing .env that has no saved port, and only search for a free port on fresh installs Only write the catch-all .gitignore into a folder the script created, and warn instead of writing into a folder that already existed Add LITELLM_BIND to the compose port mapping. It is empty by default, so existing installs keep "4000:4000", and the script sets it to 127.0.0.1: so new installs listen on this machine only * fix(quickstart): keep generated .env out of git in an existing repository folder When LITELLM_DIR is inside a git repository that does not ignore .env, add //.env to the clone's local exclude list (.git/info/exclude) instead of only warning. Tracked files, including .gitignore, are not touched * fix(quickstart): ignore an inherited LITELLM_BIND and keep .env ignored outside git Clear LITELLM_BIND from the environment before starting Compose, like the keys and project name, so .env decides the bind address and new installs stay on 127.0.0.1 In an existing folder that is not inside a git repository, add a single .env line to its .gitignore (creating it if needed, without a duplicate), so the keys stay out of commits if the folder later becomes a repository * fix(quickstart): keep an exported LITELLM_BIND for an .env without one, and show a read-first install The bind address now follows .env only when .env sets it, which every install this script creates does. For an older .env without a bind line, a LITELLM_BIND exported in the shell is kept, so an intentional 127.0.0.1: is not dropped The header shows how to download and read the script before running it --- docker/docker-compose.quickstart.yml | 8 +- scripts/quickstart.sh | 314 +++++++++++++++++++++++++++ 2 files changed, 319 insertions(+), 3 deletions(-) create mode 100755 scripts/quickstart.sh diff --git a/docker/docker-compose.quickstart.yml b/docker/docker-compose.quickstart.yml index 11631603a72..a1d47e323ff 100644 --- a/docker/docker-compose.quickstart.yml +++ b/docker/docker-compose.quickstart.yml @@ -13,11 +13,13 @@ services: litellm: image: docker.litellm.ai/berriai/litellm:main-stable ports: - - "4000:4000" + # LITELLM_BIND is empty by default, so this stays "4000:4000". The quickstart + # script sets it to "127.0.0.1:" so new installs listen on this machine only. + - "${LITELLM_BIND:-}${LITELLM_PORT:-4000}:4000" environment: LITELLM_MASTER_KEY: ${LITELLM_MASTER_KEY:?set it in .env - see the header of this file} LITELLM_SALT_KEY: ${LITELLM_SALT_KEY:?set it in .env - see the header of this file} - DATABASE_URL: postgresql://litellm:litellm@db:5432/litellm + DATABASE_URL: postgresql://litellm:${POSTGRES_PASSWORD:-litellm}@db:5432/litellm STORE_MODEL_IN_DB: "True" depends_on: db: @@ -27,7 +29,7 @@ services: image: postgres:16 environment: POSTGRES_USER: litellm - POSTGRES_PASSWORD: litellm + POSTGRES_PASSWORD: ${POSTGRES_PASSWORD:-litellm} POSTGRES_DB: litellm healthcheck: test: ["CMD-SHELL", "pg_isready -U litellm"] diff --git a/scripts/quickstart.sh b/scripts/quickstart.sh new file mode 100755 index 00000000000..469f8f2a37f --- /dev/null +++ b/scripts/quickstart.sh @@ -0,0 +1,314 @@ +#!/bin/sh +# LiteLLM Gateway quickstart: the gateway, Postgres, and the admin UI in one command. +# curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/main/scripts/quickstart.sh | sh +# +# To read it before running it: +# curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/main/scripts/quickstart.sh -o quickstart.sh +# less quickstart.sh +# sh quickstart.sh +# +# Asks at most two questions (where to keep the files, and whether to open the +# admin UI), each with a default you accept by pressing Enter. It asks nothing +# when there is no terminal, under CI or Claude Code, or when run with --yes. +# +# --yes, -y no questions: install to ~/litellm-gateway, don't open a browser +# LITELLM_DIR folder to install into (skips the folder question) +# LITELLM_PORT port for the gateway (default 4000, or the next free one) +# +# New installs listen on this machine only (127.0.0.1). To reach the gateway +# from other machines, remove LITELLM_BIND from .env and put it behind TLS. +# +# Keys and the database password are random (openssl rand), written only to +# .env with permissions 600, and never printed. Needs Docker with Compose v2. +# Everything runs inside main(), so a partial download runs nothing. +set -eu + +COMPOSE_URL="${LITELLM_COMPOSE_URL:-https://raw.githubusercontent.com/BerriAI/litellm/main/docker/docker-compose.quickstart.yml}" + +# ---------------------------------------------------------------- terminal + +INTERACTIVE=0 # a person is at a terminal we can ask +ARROWS=0 # that terminal supports the arrow-key menu +STTY_SAVED="" +POINTER='>' + +detect_terminal() { + # Piped from curl, stdin is the script itself, so questions go to /dev/tty. + if (exec /dev/null && [ "${TERM:-dumb}" != "dumb" ]; then + INTERACTIVE=1 + if STTY_SAVED="$(stty -g /dev/null)" && [ -n "$STTY_SAVED" ]; then + ARROWS=1 + fi + fi + case "${LC_ALL:-${LC_CTYPE:-${LANG:-}}}" in + *UTF-8* | *utf-8* | *UTF8* | *utf8*) POINTER='❯' ;; + esac +} + +restore_terminal() { + if [ -n "$STTY_SAVED" ]; then + stty "$STTY_SAVED" /dev/null || true + printf '\033[?25h' >/dev/tty 2>/dev/null || true + fi +} + +on_interrupt() { + restore_terminal + printf '\nCancelled.\n' >&2 + exit 130 +} + +read_key() { + # One keypress in raw mode. Enter comes back empty (command substitution + # drops the newline); arrows come back as "up" or "down". + k="$(dd bs=1 count=1 2>/dev/null sets CHOICE to the 1-based pick. +menu() { + question="$1" + CHOICE="$2" + shift 2 + count=$# + if [ "$INTERACTIVE" != 1 ]; then return 0; fi + + printf '\n%s\n' "$question" >/dev/tty + if [ "$ARROWS" = 1 ]; then + trap on_interrupt INT TERM + stty -icanon -echo min 1 time 0 /dev/tty + first=1 + while :; do + [ "$first" = 1 ] || printf '\033[%sA' "$count" >/dev/tty + first=0 + i=1 + for opt in "$@"; do + if [ "$i" = "$CHOICE" ]; then + printf '\033[2K \033[1;36m%s %s\033[0m\n' "$POINTER" "$opt" >/dev/tty + else + printf '\033[2K %s\n' "$opt" >/dev/tty + fi + i=$((i + 1)) + done + key="$(read_key)" + case "$key" in + up | k) [ "$CHOICE" -gt 1 ] && CHOICE=$((CHOICE - 1)) ;; + down | j) [ "$CHOICE" -lt "$count" ] && CHOICE=$((CHOICE + 1)) ;; + [1-9]) [ "$key" -le "$count" ] && CHOICE="$key" ;; + '' | "$(printf '\r')") break ;; + esac + done + restore_terminal + trap - INT TERM + else + i=1 + for opt in "$@"; do + printf ' %s) %s\n' "$i" "$opt" >/dev/tty + i=$((i + 1)) + done + printf 'Choose [%s]: ' "$CHOICE" >/dev/tty + answer="" + read -r answer .gitignore + elif command -v git >/dev/null 2>&1 && git rev-parse --is-inside-work-tree >/dev/null 2>&1 && + ! git check-ignore -q .env 2>/dev/null; then + # In a folder that already existed, such as a repository root, leave the + # tracked .gitignore alone and add only .env to this clone's local exclude + # list, so the generated keys cannot be committed. + exclude="$(git rev-parse --git-path info/exclude)" + mkdir -p "$(dirname "$exclude")" + exclude="$(cd "$(dirname "$exclude")" && pwd)/exclude" + printf '/%s.env\n' "$(git rev-parse --show-prefix)" >>"$exclude" + echo "Added .env to this repository's local git exclude list ($exclude), so your keys stay out of commits." + elif ! git rev-parse --is-inside-work-tree >/dev/null 2>&1; then + # An existing folder outside git: ignore only .env, so it stays out of + # commits if the folder becomes a repository later. + if ! grep -qxF '.env' .gitignore 2>/dev/null; then + # Start on a new line if the file does not end with one. + if [ -s .gitignore ] && [ -n "$(tail -c 1 .gitignore)" ]; then printf '\n' >>.gitignore; fi + printf '.env\n' >>.gitignore + fi + fi +} + +pick_port() { + saved="" + [ -f .env ] && saved="$(sed -n 's/^LITELLM_PORT=//p' .env | tail -n 1)" + if [ -n "${LITELLM_PORT:-}" ]; then + PORT="$LITELLM_PORT" + elif [ -n "$saved" ]; then + PORT="$saved" + elif [ -f .env ]; then + # An existing install without a saved port runs on the compose default. + PORT=4000 + else + PORT=4000 + while ! port_free "$PORT"; do + PORT=$((PORT + 1)) + if [ "$PORT" -gt 4099 ]; then + echo "Ports 4000 to 4099 are all in use. Set LITELLM_PORT to a free port and run this again." >&2 + exit 1 + fi + done + [ "$PORT" = 4000 ] || echo "Port 4000 is in use, so LiteLLM will use $PORT." + fi + export LITELLM_PORT="$PORT" +} + +# Docker names containers and the database volume after the project, so an +# install outside the home folder gets its own name and never shares a +# database with another litellm-gateway folder. +check_new_install() { + project=litellm-gateway + [ "$DIR" = "$HOME/litellm-gateway" ] || project="litellm-gateway-$(printf '%s' "$DIR" | cksum | cut -d ' ' -f 1)" + # Postgres keeps the password it was created with, so a new password over an + # old database volume would lock the gateway out. Stop and explain instead. + if docker volume inspect "${project}_postgres_data" >/dev/null 2>&1; then + cat >&2 </dev/null 2>&1; then + open "$url" >/dev/null 2>&1 || true + elif command -v xdg-open >/dev/null 2>&1; then + xdg-open "$url" >/dev/null 2>&1 || true + fi +} + +main() { + NO_QUESTIONS=0 + for arg in "$@"; do + case "$arg" in + -y | --yes) NO_QUESTIONS=1 ;; + *) echo "Unknown option: $arg" >&2; exit 1 ;; + esac + done + + detect_terminal + # Agents and CI get the defaults even inside a terminal, so nothing waits on a keypress. + if [ "$NO_QUESTIONS" = 1 ] || [ -n "${CI:-}" ] || [ -n "${CLAUDECODE:-}" ]; then INTERACTIVE=0; fi + trap restore_terminal EXIT + + if ! command -v docker >/dev/null 2>&1; then + cat >&2 <<'EOF' +Docker is not installed. The LiteLLM Gateway runs in Docker alongside a Postgres database. + + Install Docker, then run this again: https://docs.docker.com/get-docker/ + Or deploy in one click (Railway or Render): https://docs.litellm.ai/docs/proxy/docker_quick_start + Only need to call models from Python? pip install litellm +EOF + exit 1 + fi + docker compose version >/dev/null 2>&1 || { echo "Docker Compose v2 ('docker compose') is required." >&2; exit 1; } + docker info >/dev/null 2>&1 || { echo "Docker is installed but not running. Start it and run this again." >&2; exit 1; } + command -v openssl >/dev/null 2>&1 || { echo "openssl is required to generate keys." >&2; exit 1; } + + echo "LiteLLM quickstart" + pick_folder + [ -f .env ] || check_new_install + curl -fsSL -o docker-compose.quickstart.yml "$COMPOSE_URL" + pick_port + + if [ -f .env ]; then + echo "Reusing $DIR/.env, so existing keys and data keep working." + else + (umask 077 && printf 'LITELLM_MASTER_KEY=sk-%s\nLITELLM_SALT_KEY=sk-%s\nPOSTGRES_PASSWORD=%s\nLITELLM_PORT=%s\nLITELLM_BIND=127.0.0.1:\nCOMPOSE_PROJECT_NAME=%s\n' \ + "$(openssl rand -hex 32)" "$(openssl rand -hex 32)" "$(openssl rand -hex 24)" "$PORT" "$project" >.env) + echo "Generated $DIR/.env with your master key, salt key, and database password. Keep this file." + fi + + # Compose prefers values already set in the shell over .env, so drop any + # inherited ones: .env stays the only source for keys and the project name. + unset LITELLM_MASTER_KEY LITELLM_SALT_KEY POSTGRES_PASSWORD COMPOSE_PROJECT_NAME + # The bind address follows .env when .env sets it (every install this script + # creates does). For an older .env without it, a value exported in the shell + # is kept, so an intentional LITELLM_BIND=127.0.0.1: is not dropped. + if grep -q '^LITELLM_BIND=' .env; then unset LITELLM_BIND; fi + + echo "Starting LiteLLM and Postgres (the first run downloads the images)..." + docker compose -f docker-compose.quickstart.yml up -d + + i=0 + until curl -fsS "http://127.0.0.1:$PORT/health/readiness" >/dev/null 2>&1; do + i=$((i + 1)) + if [ "$i" -gt 90 ]; then + echo "The gateway did not become ready in 3 minutes. Check: cd $DIR && docker compose -f docker-compose.quickstart.yml logs litellm" >&2 + exit 1 + fi + sleep 2 + done + + echo + echo "LiteLLM is running." + echo " Admin UI: http://localhost:$PORT/ui" + echo " Username: admin" + echo " Password: the LITELLM_MASTER_KEY value in $DIR/.env" + echo " Next: in the UI, open Models + Endpoints > Add Model and paste a provider API key" + echo " Stop it: cd $DIR && docker compose -f docker-compose.quickstart.yml down" + + open_browser "http://localhost:$PORT/ui" +} + +main "$@" From 71788fe1c5db9276b3fe06a73c81075603b78e92 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Fri, 2 Oct 2026 09:59:19 -0700 Subject: [PATCH 005/139] fix(lens): simplify the example investigation preview (#44123) * fix(lens): simplify the example investigation preview * test(lens): cover example preview interactions --- .../lens/_components/LensOverview.tsx | 86 ++++++++----------- .../_components/LensView.integration.test.tsx | 36 ++++++++ 2 files changed, 73 insertions(+), 49 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensOverview.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensOverview.tsx index 414f9eb41fd..2e2a0318d85 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensOverview.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensOverview.tsx @@ -75,62 +75,50 @@ export function InvestigationExample({ onClose }: { onClose: () => void }) { if (!open) onClose(); }} > - -
- - Example investigation - Does the support agent recover when a tool fails? - -
-
-
- + + + Example investigation + +
+
+ + 3 of 20 conversations
-

- Failed lookups leave customers without an answer +

+ Failed lookups leave customers without answers

-

- The agent repeats the same failed order lookup, then ends the conversation without an answer or a handoff. -

-
+ + The agent retries the same failed order lookup, then ends the conversation without an answer or a handoff. + +
+
+

After three failed lookups, the agent replies:

-
“I will check that for you.”
-
- - See the trace - -
    -
  1. - Customer · Where is my order? -
  2. -
  3. - Order lookup · Service unavailable -
  4. -
  5. - Two retries · Same error, no new information -
  6. -
  7. - Agent · I will check that for you. Conversation - ends. -
  8. -
-
+
“I will check that for you.”
-
-

What to change

-

- After repeated failures, explain the problem and offer a handoff instead of retrying. -

-
-
-
-

Illustrative data, not your agent’s results

- +
+ + See the trace + +
    +
  1. + Customer · Where is my order? +
  2. +
  3. + Order lookup · Service unavailable +
  4. +
  5. + Two retries · Same error, no new information +
  6. +
  7. + Agent · I will check that for you. Conversation + ends. +
  8. +
+
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 2e4f1dcc1bc..a3ec7877223 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 @@ -231,6 +231,42 @@ it("runs saved settings immediately without opening setup", async () => { expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); }); +it("lets a new user inspect example evidence and return to setup without starting an investigation", async () => { + window.history.replaceState({}, "", "/lens/"); + testQueryClient.clear(); + vi.mocked(apiClient.get).mockImplementation(async (path) => { + if (path === "/lens") return { lenses: [], workers: [], tracing_enabled: false }; + return { traces: false, requests: false }; + }); + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(await screen.findByRole("button", { name: "View an example" })); + const example = within(screen.getByRole("dialog", { name: "Example investigation" })); + expect(example.getByRole("heading", { name: "Failed lookups leave customers without answers" })).toBeVisible(); + const evidence = example.getByText(/Where is my order\?/); + expect(evidence).not.toBeVisible(); + await user.click(example.getByText("See the trace")); + expect(evidence).toBeVisible(); + expect(example.getByText(/Service unavailable/)).toBeVisible(); + await user.click(example.getByText("See the trace")); + expect(evidence).not.toBeVisible(); + + await user.click(example.getByRole("button", { name: "Close" })); + await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); + expect(screen.getByRole("link", { name: "Set up traces" })).toBeVisible(); + expect(screen.getByRole("button", { name: "Connect worker" })).toBeDisabled(); + expect(screen.getByRole("button", { name: "New investigation" })).toBeDisabled(); + + await user.click(screen.getByRole("button", { name: "View an example" })); + expect(screen.getByRole("dialog", { name: "Example investigation" })).toBeVisible(); + expect(screen.getByText(/Where is my order\?/)).not.toBeVisible(); + await user.keyboard("{Escape}"); + await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); + expect(screen.getByRole("button", { name: "View an example" })).toBeVisible(); + expect(apiClient.post).not.toHaveBeenCalled(); +}); + it("guides a first-time administrator into worker connection and lens setup", async () => { window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); From ebfec956d26b8eab3269cbae86f148187282e488 Mon Sep 17 00:00:00 2001 From: Bernedotcom2312 <88426124+Bernedotcom2312@users.noreply.github.com> Date: Fri, 2 Oct 2026 19:13:52 +0200 Subject: [PATCH 006/139] chore(helm): drop migrationJob values the chart never reads (#42141) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `migrationJob.retries` and `migrationJob.disableSchemaUpdate` are declared in values.yaml but referenced by no template, no test and no README row. Setting either changes nothing about the rendered Job. `disableSchemaUpdate` is the misleading one: its comment promises "the job will exit with code 0", but the Job hardcodes DISABLE_SCHEMA_UPDATE=false and renders it after envVars/extraEnvVars precisely so nothing can turn the migration off — that ordering is what #12809 fixed. An operator who reads values.yaml, sets the flag and watches migrations run anyway has no way to tell the knob is inert. `migrationJob.enabled: false` is the supported way to skip the Job, and the componentized chart in helm/litellm already ships a migrationJob block with neither key. `retries` is simply dead: Jobs retry through `backoffLimit`, which the chart does render. Removing values keys is backward compatible — Helm ignores user values that no template consumes, so existing releases setting either key keep working. Adds a test pinning the override: with envVars.DISABLE_SCHEMA_UPDATE="true" the Job's last env entry is still DISABLE_SCHEMA_UPDATE=false, so the last-wins ordering cannot regress and the key cannot quietly come back as a chart value. Verified by mutation: flipping the hardcoded value and moving the entry above the envVars loop each fail the suite. Co-authored-by: Claude Opus 5 Co-authored-by: ryan-crabbe-berri --- .../tests/migrations-job_tests.yaml | 18 ++++++++++++++++++ helm/litellm-helm/values.yaml | 2 -- 2 files changed, 18 insertions(+), 2 deletions(-) diff --git a/helm/litellm-helm/tests/migrations-job_tests.yaml b/helm/litellm-helm/tests/migrations-job_tests.yaml index 1fe545636d4..dd4276ac60f 100644 --- a/helm/litellm-helm/tests/migrations-job_tests.yaml +++ b/helm/litellm-helm/tests/migrations-job_tests.yaml @@ -112,6 +112,24 @@ tests: name: CUSTOM_VAR value: "custom_value" + - it: should override a user-supplied DISABLE_SCHEMA_UPDATE so the Job always migrates + template: migrations-job.yaml + set: + envVars: + DISABLE_SCHEMA_UPDATE: "true" + migrationJob: + enabled: true + asserts: + # The Job is what owns the schema, so it renders its own + # DISABLE_SCHEMA_UPDATE=false after envVars and extraEnvVars. Kubernetes + # takes the last value for a duplicated name, so the user's "true" cannot + # leave the schema unmigrated. Skipping migrations is migrationJob.enabled. + - equal: + path: spec.template.spec.containers[0].env[-1] + value: + name: DISABLE_SCHEMA_UPDATE + value: "false" + - it: should not include DATABASE_URL when deployStandalone is false template: migrations-job.yaml set: diff --git a/helm/litellm-helm/values.yaml b/helm/litellm-helm/values.yaml index fcee331a5aa..03d2a66a2b5 100644 --- a/helm/litellm-helm/values.yaml +++ b/helm/litellm-helm/values.yaml @@ -545,7 +545,6 @@ redis: # Prisma migration job settings migrationJob: enabled: true # Enable or disable the schema migration Job - retries: 3 # Number of retries for the Job in case of failure backoffLimit: 4 # Backoff limit for Job restarts # Wall-clock budget for the whole Job, shared across every `backoffLimit` # retry rather than granted per attempt. Without it a migration that blocks @@ -554,7 +553,6 @@ migrationJob: # stop reconciling the whole chart until someone deletes the Job by hand. # Set to null to opt out and restore the unbounded behaviour. activeDeadlineSeconds: 1800 - disableSchemaUpdate: false # Skip schema migrations for specific environments. When True, the job will exit with code 0. # Optional service account for the migration job. # Only used when migrationJob.hooks.helm.enabled=true and serviceAccount.create=true. # In that case, pre-install/pre-upgrade hooks run before normal resources, so this defaults to "default". From 19da81579b3fa40bea39fb41ca3379d9da32d83e Mon Sep 17 00:00:00 2001 From: Jim Aldon D'Souza Date: Fri, 2 Oct 2026 10:15:41 -0700 Subject: [PATCH 007/139] fix(ui): register tencent in the Add Model provider dropdown (#40924) * fix(ui): register tencent in the Add Model provider dropdown The Add Model provider dropdown is driven by the proxy's /public/providers/fields endpoint, which serves provider_create_fields.json. Tencent was frozen in the test's ADD_MODEL_UNLISTED_PROVIDERS set, so it never appeared in the dropdown. Add a Tencent entry (optional api_base + required api_key, matching TENCENT_API_BASE/TENCENT_API_KEY) and unfreeze it in the backend test. Register Tencent in the UI Providers enum, provider_map, and placeholder map so the dropdown resolves the display name and model placeholder. * fix(tencent): drop test docstring to satisfy comment policy * fix(ui): bundle the Tencent Cloud logo for the provider dropdown --- .../provider_create_fields.json | 28 +++++++++++++++++++ .../public_endpoints/test_public_endpoints.py | 26 ++++++++++++++++- .../public/assets/logos/tencent.svg | 6 ++++ .../components/provider_info_helpers.test.tsx | 9 ++++++ .../src/components/provider_info_helpers.tsx | 5 ++++ 5 files changed, 73 insertions(+), 1 deletion(-) create mode 100644 ui/litellm-dashboard/public/assets/logos/tencent.svg diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index 11d2ff61b95..6e96d6ad0ec 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -3161,6 +3161,34 @@ ], "default_model_placeholder": "soniox/stt-async-v5" }, + { + "provider": "Tencent", + "provider_display_name": "Tencent", + "litellm_provider": "tencent", + "credential_fields": [ + { + "key": "api_base", + "label": "API Base", + "placeholder": "https://tokenhub-intl.tencentcloudmaas.com/v1", + "tooltip": null, + "required": false, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "api_key", + "label": "API Key", + "placeholder": null, + "tooltip": null, + "required": true, + "field_type": "password", + "options": null, + "default_value": null + } + ], + "default_model_placeholder": "tencent/deepseek-v4-pro" + }, { "provider": "TEXT_COMPLETION_CODESTRAL", "provider_display_name": "Text-Completion-Codestral", diff --git a/tests/unit/proxy/public_endpoints/test_public_endpoints.py b/tests/unit/proxy/public_endpoints/test_public_endpoints.py index 18839a65d62..be309a67d58 100644 --- a/tests/unit/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/unit/proxy/public_endpoints/test_public_endpoints.py @@ -377,6 +377,31 @@ def test_chatgpt_provider_fields(): assert chatgpt["credential_fields"] == [] +def test_tencent_provider_fields(): + app_instance = FastAPI() + app_instance.include_router(router) + test_client = TestClient(app_instance) + + response = test_client.get("/public/providers/fields") + assert response.status_code == 200 + providers = response.json() + + tencent = next((p for p in providers if p["provider"] == "Tencent"), None) + assert tencent is not None, "Tencent provider entry not found" + + assert tencent["provider_display_name"] == "Tencent" + assert tencent["litellm_provider"] == LlmProviders.TENCENT.value + assert tencent["default_model_placeholder"].startswith("tencent/") + + fields_by_key = {f["key"]: f for f in tencent["credential_fields"]} + + assert fields_by_key["api_key"]["required"] is True + assert fields_by_key["api_key"]["field_type"] == "password" + + assert fields_by_key["api_base"]["field_type"] == "text" + assert fields_by_key["api_base"]["required"] is False + + ADD_MODEL_UNLISTED_PROVIDERS: Final = frozenset( { "a2a", @@ -412,7 +437,6 @@ ADD_MODEL_UNLISTED_PROVIDERS: Final = frozenset( "scaleway", "stability", "synthetic", - "tencent", "tensormesh", "text-completion-inception", "transcribe", diff --git a/ui/litellm-dashboard/public/assets/logos/tencent.svg b/ui/litellm-dashboard/public/assets/logos/tencent.svg new file mode 100644 index 00000000000..ee43c71f4d5 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/tencent.svg @@ -0,0 +1,6 @@ + + Tencent Cloud + + + + diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx index 29ce2d4865a..7cfdaf3275d 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx @@ -209,6 +209,11 @@ describe("provider_info_helpers", () => { const { logo } = getProviderLogoAndName("openai"); expect(logo).toContain("openai_small"); }); + + it("should resolve the Tencent provider to its bundled logo", () => { + const { logo } = getProviderLogoAndName("tencent"); + expect(logo).toContain("tencent"); + }); }); describe("getPlaceholder", () => { @@ -309,6 +314,10 @@ describe("provider_info_helpers", () => { expect(getPlaceholder("CHATGPT")).toBe("chatgpt/gpt-5.4"); }); + it("should return a tencent/ placeholder for the Tencent provider", () => { + expect(getPlaceholder(Providers.Tencent)).toBe("tencent/deepseek-v4-pro"); + }); + it("should return default gpt-3.5-turbo placeholder for unknown provider", () => { expect(getPlaceholder("UnknownProvider" as any)).toBe("gpt-3.5-turbo"); }); diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx index b0e33338bab..5ea693bea10 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx @@ -55,6 +55,7 @@ import sapLogo from "../../public/assets/logos/sap.png"; import scxAiLogo from "../../public/assets/logos/scx_ai.svg"; import snowflakeLogo from "../../public/assets/logos/snowflake.svg"; import sonioxLogo from "../../public/assets/logos/soniox.svg"; +import tencentLogo from "../../public/assets/logos/tencent.svg"; import togetheraiLogo from "../../public/assets/logos/togetherai.svg"; import topazLogo from "../../public/assets/logos/topaz.svg"; import v0Logo from "../../public/assets/logos/v0.svg"; @@ -167,6 +168,7 @@ export enum Providers { Snowflake = "Snowflake", Soniox = "Soniox", TEXT_COMPLETION_CODESTRAL = "Text-Completion-Codestral", + Tencent = "Tencent", TogetherAI = "TogetherAI", TOPAZ = "Topaz", Triton = "Triton", @@ -286,6 +288,7 @@ export const provider_map: Record = { Snowflake: "snowflake", Soniox: "soniox", TEXT_COMPLETION_CODESTRAL: "text-completion-codestral", + Tencent: "tencent", TogetherAI: "together_ai", TOPAZ: "topaz", Triton: "triton", @@ -384,6 +387,7 @@ export const providerLogoMap: Partial> = { [Providers.SCX_AI]: scxAiLogo.src, [Providers.Snowflake]: snowflakeLogo.src, [Providers.Soniox]: sonioxLogo.src, + [Providers.Tencent]: tencentLogo.src, [Providers.TEXT_COMPLETION_CODESTRAL]: mistralLogo.src, [Providers.TogetherAI]: togetheraiLogo.src, [Providers.TOPAZ]: topazLogo.src, @@ -453,6 +457,7 @@ const providerPlaceholderMap: Partial> = { [Providers.Sail]: "sail/openai/gpt-oss-120b", [Providers.SCX_AI]: "scx-ai/GLM-5.2", [Providers.Snowflake]: "snowflake/mistral-7b", + [Providers.Tencent]: "tencent/deepseek-v4-pro", [Providers.Vertex_AI]: "gemini-pro", [Providers.VolcEngine]: "volcengine/", [Providers.Voyage]: "voyage/", From 6c2ede00ac897fd9edd2a1f60e5733859f3ca308 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 10:18:50 -0700 Subject: [PATCH 008/139] test: remove 130 legacy tests owned by stronger unit proofs (#44157) * test: remove 130 legacy tests owned by stronger unit proofs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test: list router _embedding and _aembedding as covered via public embedding calls Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yuneng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/audio_tests/test_audio_speech.py | 10 - tests/audio_tests/test_whisper.py | 11 - .../router_code_coverage.py | 2 + tests/image_gen_tests/test_image_edits.py | 60 +-- .../test_azure_responses_api.py | 3 + .../test_openai_responses_api.py | 20 +- .../test_google_interactions_integration.py | 27 -- .../test_anthropic_completion.py | 152 +------ tests/llm_translation/test_azure_ai.py | 2 +- tests/llm_translation/test_azure_o_series.py | 4 + tests/llm_translation/test_azure_openai.py | 15 - .../test_bedrock_completion.py | 112 +---- tests/llm_translation/test_bedrock_gpt_oss.py | 2 + .../test_bedrock_invoke_tests.py | 29 ++ tests/llm_translation/test_bedrock_llama.py | 4 + .../llm_translation/test_bedrock_moonshot.py | 2 + .../llm_translation/test_bedrock_nova_json.py | 8 + tests/llm_translation/test_gemini.py | 10 + tests/llm_translation/test_groq.py | 4 + tests/llm_translation/test_openai.py | 12 +- tests/llm_translation/test_openai_o1.py | 31 +- tests/llm_translation/test_together_ai.py | 8 + tests/llm_translation/test_xai.py | 27 +- tests/local_testing/test_acooldowns_router.py | 106 ----- .../test_amazing_vertex_completion.py | 220 ---------- tests/local_testing/test_arize_ai.py | 20 - tests/local_testing/test_async_fn.py | 37 -- tests/local_testing/test_completion.py | 384 ------------------ .../test_function_call_parsing.py | 2 +- tests/local_testing/test_function_calling.py | 149 +------ .../local_testing/test_lowest_cost_routing.py | 32 -- tests/local_testing/test_router.py | 249 ------------ tests/local_testing/test_streaming.py | 22 - tests/local_testing/test_timeout.py | 29 -- .../test_e2e_openai_responses_api.py | 14 - .../test_anthropic_messages_passthrough.py | 70 +--- tests/test_openai_endpoints.py | 33 -- .../unified_google_tests/base_google_test.py | 40 -- .../test_google_ai_studio.py | 2 + .../test_litellm_responses_bridge.py | 2 + 40 files changed, 133 insertions(+), 1833 deletions(-) diff --git a/tests/audio_tests/test_audio_speech.py b/tests/audio_tests/test_audio_speech.py index 2559e4704d0..77e7fab3f00 100644 --- a/tests/audio_tests/test_audio_speech.py +++ b/tests/audio_tests/test_audio_speech.py @@ -320,16 +320,6 @@ def test_audio_speech_cost_calc(): assert standard_logging_payload["response_cost"] > 0 -def test_audio_speech_gemini(): - result = litellm.speech( - model="gemini/gemini-2.5-flash-preview-tts", - input="the quick brown fox jumped over the lazy dogs", - api_key=os.getenv("GEMINI_API_KEY"), - ) - - print(result) - - @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=1) async def test_azure_ava_tts_async(): diff --git a/tests/audio_tests/test_whisper.py b/tests/audio_tests/test_whisper.py index 5302a78b94a..0509999e9f4 100644 --- a/tests/audio_tests/test_whisper.py +++ b/tests/audio_tests/test_whisper.py @@ -138,17 +138,6 @@ async def test_whisper_log_pre_call(): mock_log_pre_call.assert_called_once() -@pytest.mark.asyncio -async def test_gpt_4o_transcribe(): - from litellm.litellm_core_utils.litellm_logging import Logging - from datetime import datetime - from unittest.mock import patch, MagicMock - - await litellm.atranscription( - model="openai/gpt-4o-transcribe", file=_audio_file(), response_format="json" - ) - - @pytest.mark.asyncio async def test_gpt_4o_transcribe_model_mapping(): """Test that GPT-4o transcription models are correctly mapped and not hardcoded to whisper-1""" diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index df149f6c56a..8d7c1e140d2 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -91,6 +91,8 @@ ignored_function_names = [ "_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py "_apply_updated_routing_strategy_args", # Tested via update_settings in test_lowest_latency.py (file lacks "router" in name) "arm_routing_read_prefetch", # Tested in tests/unit/caching/test_request_redis_batch_pre_call.py (file lacks "router" in name) + "_embedding", + "_aembedding", ] diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py index 36fd65ba71b..ff0cf7e3075 100644 --- a/tests/image_gen_tests/test_image_edits.py +++ b/tests/image_gen_tests/test_image_edits.py @@ -128,6 +128,8 @@ class TestOpenAIImageEditGPTImage1(BaseLLMImageEditTest): Concrete implementation of BaseLLMImageEditTest for OpenAI image edits. """ + test_openai_image_edit_litellm_sdk = None + def get_base_image_edit_call_args(self) -> dict: """Return base call args for OpenAI image edit""" return { @@ -622,64 +624,6 @@ def test_recraft_image_edit_config(): assert files[0][1][2] == "image/png" # Content type -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.flaky(retries=3, delay=2) -@pytest.mark.asyncio -async def test_multiple_vs_single_image_edit(sync_mode): - """Test that both single and multiple image editing work correctly""" - from litellm import image_edit, aimage_edit - - litellm._turn_on_debug() - - try: - prompt = "Add a soft blue tint to the image(s)" - - # Test single image - if sync_mode: - single_result = image_edit( - prompt=prompt, - model="gpt-image-1", - image=_make_single_test_image(), - ) - else: - single_result = await aimage_edit( - prompt=prompt, - model="gpt-image-1", - image=_make_single_test_image(), - ) - - print("Single image result:", single_result) - ImageResponse.model_validate(single_result) - - # Test multiple images - if sync_mode: - multiple_result = image_edit( - prompt=prompt, - model="gpt-image-1", - image=_make_test_images(), - ) - else: - multiple_result = await aimage_edit( - prompt=prompt, - model="gpt-image-1", - image=_make_test_images(), - ) - - print("Multiple images result:", multiple_result) - ImageResponse.model_validate(multiple_result) - - # Both should return valid responses - assert single_result is not None - assert multiple_result is not None - assert single_result.data is not None - assert multiple_result.data is not None - assert len(single_result.data) > 0 - assert len(multiple_result.data) > 0 - - except litellm.ContentPolicyViolationError as e: - pytest.skip(f"Content policy violation: {e}") - - @pytest.mark.flaky(retries=3, delay=2) @pytest.mark.asyncio async def test_multiple_image_edit_with_different_formats(): diff --git a/tests/llm_responses_api_testing/test_azure_responses_api.py b/tests/llm_responses_api_testing/test_azure_responses_api.py index 6f1bb440341..1ec7bafd1ad 100644 --- a/tests/llm_responses_api_testing/test_azure_responses_api.py +++ b/tests/llm_responses_api_testing/test_azure_responses_api.py @@ -18,6 +18,9 @@ from base_responses_api import BaseResponsesAPITest class TestAzureResponsesAPITest(BaseResponsesAPITest): + test_multiturn_responses_api = None + test_responses_api_with_tool_calls = None + def get_base_completion_call_args(self): return { "model": "azure/gpt-4.1-mini", diff --git a/tests/llm_responses_api_testing/test_openai_responses_api.py b/tests/llm_responses_api_testing/test_openai_responses_api.py index 373d61a367c..051eb7494b2 100644 --- a/tests/llm_responses_api_testing/test_openai_responses_api.py +++ b/tests/llm_responses_api_testing/test_openai_responses_api.py @@ -23,6 +23,8 @@ from base_responses_api import BaseResponsesAPITest, validate_responses_api_resp class TestOpenAIResponsesAPITest(BaseResponsesAPITest): + test_responses_api_with_tool_calls = None + def get_base_completion_call_args(self): return { "model": "openai/gpt-5.5", @@ -1597,24 +1599,6 @@ async def test_openai_gpt5_reasoning_effort_parameter(): print("Response:", json.dumps(response, indent=4, default=str)) -@pytest.mark.asyncio -@pytest.mark.parametrize("stream", [True, False]) -async def test_basic_openai_responses_with_websearch(stream): - litellm._turn_on_debug() - request_model = "gpt-5.5" - response = await litellm.aresponses( - model=request_model, - stream=stream, - input="hi", - tools=[{"type": "web_search", "search_context_size": "low"}], - ) - if stream: - async for chunk in response: - print("chunk=", json.dumps(chunk, indent=4, default=str)) - else: - print("response=", json.dumps(response, indent=4, default=str)) - - @pytest.mark.asyncio async def test_openai_responses_api_token_limit_error(): """ diff --git a/tests/llm_translation/interactions/test_google_interactions_integration.py b/tests/llm_translation/interactions/test_google_interactions_integration.py index 10e6cf86e6d..23e83e97c5f 100644 --- a/tests/llm_translation/interactions/test_google_interactions_integration.py +++ b/tests/llm_translation/interactions/test_google_interactions_integration.py @@ -163,33 +163,6 @@ class TestGoogleInteractionsStreaming: class TestGoogleInteractionsMultiTurn: """Tests for multi-turn conversations using Step[] input.""" - def test_multi_turn_conversation(self, api_key): - """Test a multi-turn conversation per OpenAPI spec (Step[] format).""" - response = interactions.create( - model="gemini/gemini-2.5-flash", - input=[ - { - "type": "user_input", - "content": [{"type": "text", "text": "My name is Alice."}], - }, - { - "type": "model_output", - "content": [ - {"type": "text", "text": "Hello Alice! Nice to meet you."} - ], - }, - { - "type": "user_input", - "content": [{"type": "text", "text": "What is my name?"}], - }, - ], - api_key=api_key, - ) - - assert response is not None - print(f"Multi-turn response: {response}") - - class TestGoogleInteractionsAgent: """Tests for agent interactions (per OpenAPI spec).""" diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py index 8c55014955f..396cd74f75a 100644 --- a/tests/llm_translation/test_anthropic_completion.py +++ b/tests/llm_translation/test_anthropic_completion.py @@ -557,6 +557,15 @@ class TestAnthropicCompletion(BaseLLMChatTest, BaseAnthropicChatTest): except litellm.InternalServerError: pytest.skip("Model is overloaded") + @pytest.mark.parametrize("sync_mode", [True]) + @pytest.mark.asyncio + async def test_pdf_handling(self, pdf_messages, sync_mode): + await super().test_pdf_handling(pdf_messages, sync_mode) + test_content_list_handling = None + test_image_url = None + test_image_url_string = None + test_web_search = None + def test_convert_tool_response_to_message_with_values(): """Test converting a tool response with 'values' key to a message""" @@ -910,37 +919,6 @@ def test_map_stop_sequences(stop_input, expected_output, drop_params): assert result == expected_output -@pytest.mark.asyncio -async def test_anthropic_structured_output(): - """ - Test the _transform_response_for_structured_output - - Relevant Issue: https://github.com/BerriAI/litellm/issues/8291 - """ - from litellm import acompletion - - args = { - "model": "claude-sonnet-4-5-20250929", - "seed": 3015206306868917280, - "stop": None, - "messages": [ - { - "role": "system", - "content": 'You are a hello world agent.\nAlways respond in the following valid JSON format: {\n "response": "response",\n}\n', - }, - {"role": "user", "content": "Respond with hello world"}, - ], - "temperature": 0, - "response_format": {"type": "json_object"}, - "drop_params": True, - } - - response = await acompletion(**args) - assert response is not None - - print(response) - - def test_anthropic_citations_api(): """ Test the citations API @@ -1139,92 +1117,6 @@ def test_anthropic_custom_headers(): assert "computer-use-2025-01-24" in headers["anthropic-beta"] -@pytest.mark.parametrize( - "model", - [ - "anthropic/claude-sonnet-4-5-20250929", - # "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - ], -) -def test_anthropic_thinking_in_assistant_message(model): - litellm._turn_on_debug() - params = { - "model": model, - "messages": [ - { - "role": "assistant", - "content": [ - { - "type": "thinking", - "thinking": "Alright, let's get started with resolving this issue about implementing `__len__` for the Pipeline class in scikit-learn.\n\n### Step 1: Explore the repository structure to understand it better\n\nLet's first explore the repository structure to get a better understanding of the codebase:\n\n\nls -la /workspace/scikit-learn__scikit-learn__0.21\n\n\nThis will give us an overview of the top-level directories and files. From there, we can delve deeper into relevant parts of the code.\n\nLet's look at the main organization of the repository first, then we'll locate the Pipeline class, understand how it's currently implemented, and then add the `__len__` method to it.\n\nNext, let's look at the pipeline module, as the issue mentions `pipe[:len(pipe)]`, indicating that there is indexing support for Pipeline, but it lacks a `__len__` method.\n\nSince the issue mentions that Pipeline supports indexing but not `__len__`, I'll check the Pipeline class to understand its implementation.\n\nThe Pipeline class is likely to be in sklearn/pipeline.py or within a pipeline directory. Let's look for it:\n\n\nfind /workspace/scikit-learn__scikit-learn__0.21 -name \"*pipeline*\" -type f | grep -v __pycache__ | grep -v \"\\.pyc$\"\n\n\nThis will help us locate the file containing the Pipeline class.\n\nOnce we've located the file, we'll examine it to understand:\n1. How indexing is currently implemented\n2. Where we need to add the `__len__` method\n3. How the length of a pipeline should be determined\n\nThen, we'll implement the `__len__` method, test our changes to ensure they fix the issue, and verify that our changes don't break existing functionality.\n\nLet's begin by exploring the repository.\n\n### Step 2: Reproduce the error\n\nAfter exploring the repository, I'll create a script to reproduce the error described in the issue. This script will use the code provided in the \"Steps/Code to Reproduce\" section of the issue description.\n\nThen I'll run the script to confirm the error occurs as described, which will help us understand what needs to be fixed.\n\n### Step 3: Fix the issue\n\nOnce I've understood the problem and located the relevant files, I'll modify the code to add the `__len__` method to the Pipeline class. This method should return the number of steps in the pipeline, which is consistent with what a user would expect when calling `len(pipe)`.\n\n### Step 4: Test the fix\n\nAfter implementing the fix, I'll rerun the reproduction script to verify that the error is resolved and `len(pipe)` works correctly.\n\n### Step 5: Consider edge cases\n\nI'll think about potential edge cases like:\n- Empty pipelines\n- Very long pipelines\n- Any other special cases where the `__len__` implementation might behave unexpectedly\n\nI'll add tests for these cases to ensure our implementation handles them correctly.\n\n### Step 6: Comprehensive testing\n\nI'll run the existing tests to ensure our changes don't break any existing functionality, and I'll also write additional tests to specifically test our new `__len__` method.\n\nLet's start with exploring the repository.\n\nLet's start by exploring the repository structure:\n\n\nls -la /workspace/scikit-learn__scikit-learn__0.21\n\n\nNow, let's find the Pipeline class:\n\n\nfind /workspace/scikit-learn__scikit-learn__0.21 -name \"*pipeline*\" -type f | grep -v __pycache__ | grep -v \"\\.pyc$\"\n\n\nBased on the repository structure, let's look at the pipeline implementation:\n\n\ncat /workspace/scikit-learn__scikit-learn__0.21/sklearn/pipeline.py\n\n\nLet's check if there are tests for the Pipeline class:\n\n\nfind /workspace/scikit-learn__scikit-learn__0.21 -name \"*pipeline*test*\" -type f | grep -v __pycache__ | grep -v \"\\.pyc$\"\n\n\nNow, let's create a script to reproduce the error. I'll create a file in the workspace directory and add the code from the issue description:\n\n\ncat > /workspace/reproduce_pipeline_len_error.py << 'EOF'\nfrom sklearn import svm\nfrom sklearn.datasets import samples_generator\nfrom sklearn.feature_selection import SelectKBest\nfrom sklearn.feature_selection import f_regression\nfrom sklearn.pipeline import Pipeline\n\n# generate some data to play with\nX, y = samples_generator.make_classification(\n n_informative=5, n_redundant=0, random_state=42)\n\nanova_filter = SelectKBest(f_regression, k=5)\nclf = svm.SVC(kernel='linear')\npipe = Pipeline([('anova', anova_filter), ('svc', clf)])\n\nprint(\"Trying to get length of pipeline...\")\ntry:\n length = len(pipe)\n print(f\"Length of pipeline: {length}\")\nexcept Exception as e:\n print(f\"Error: {type(e).__name__}: {e}\")\nEOF\n\n\nLet's run this to confirm the error:\n\n\npython /workspace/reproduce_pipeline_len_error.py\n\n\nNow, let's implement the `__len__` method in the Pipeline class. First, I'll examine the Pipeline class to understand how it's structured and where to add the `__len__` method.\n\nFrom the error we saw above and looking at the pipeline.py file, we should add a `__len__` method to the Pipeline class.\n\nLet's modify the pipeline.py file:\n\n\ncp /workspace/scikit-learn__scikit-learn__0.21/sklearn/pipeline.py /workspace/pipeline.py.backup\n\n\nNow, let's edit the file to add the `__len__` method:\n\n\nsed -n -e '/class Pipeline/,/^class /p' /workspace/scikit-learn__scikit-learn__0.21/sklearn/pipeline.py\n\n\nThis will show us the Pipeline class definition so we can understand where to add the `__len__` method.\n\nBased on what we see, I'll add the `__len__` method to the Pipeline class:\n\n\nsed -i -e '/def __getitem__.*:/a\\\n def __len__(self):\\\n \"\"\"Returns the length of the Pipeline.\\\n\\\n Returns\\\n -------\\\n length : int\\\n The number of steps in the pipeline.\\\n \"\"\"\\\n return len(self.steps)\n' /workspace/scikit-learn__scikit-learn__0.21/sklearn/pipeline.py\n\n\nLet's verify that the changes were applied correctly:\n\n\ngrep -A 10 \"__len__\" /workspace/scikit-learn__scikit-learn__0.21/sklearn/pipeline.py\n\n\nNow, let's run our script again to see if the error is fixed:\n\n\npython /workspace/reproduce_pipeline_len_error.py\n\n\nIf the error is fixed, let's test some edge cases by modifying our script:\n\n\ncat > /workspace/test_pipeline_len.py << 'EOF'\nfrom sklearn.pipeline import Pipeline\nfrom sklearn import svm\nfrom sklearn.feature_selection import SelectKBest, f_regression\nfrom sklearn.datasets import samples_generator\nfrom sklearn.preprocessing import StandardScaler\n\n# Test case 1: Normal pipeline\nanova_filter = SelectKBest(f_regression, k=5)\nclf = svm.SVC(kernel='linear')\npipe = Pipeline([('anova', anova_filter), ('svc', clf)])\nprint(f\"Length of normal pipeline: {len(pipe)}\")\n\n# Test case 2: Empty pipeline\nempty_pipe = Pipeline([])\nprint(f\"Length of empty pipeline: {len(empty_pipe)}\")\n\n# Test case 3: Long pipeline\nX, y = samples_generator.make_classification(\n n_informative=5, n_redundant=0, random_state=42)\n\nsteps = []\nfor i in range(10):\n steps.append((f'scaler_{i}', StandardScaler()))\nlong_pipe = Pipeline(steps)\nprint(f\"Length of long pipeline: {len(long_pipe)}\")\n\n# Test case 4: Slicing with len\npipe = Pipeline([('anova', anova_filter), ('svc', clf)])\nsliced_pipe = pipe[:len(pipe)]\nprint(f\"Original pipe steps: {[name for name, _ in pipe.steps]}\")\nprint(f\"Sliced pipe steps: {[name for name, _ in sliced_pipe.steps]}\")\nEOF\n\n\nLet's run this test script:\n\n\npython /workspace/test_pipeline_len.py\n\n\nLet's also run the tests for the pipeline module to make sure we haven't broken anything:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\npython -m pytest sklearn/tests/test_pipeline.py -v\n\n\nAfter testing the changes, let's update our script once more to test the edge case where we use `len(pipe)` in slicing:\n\n\ncat >> /workspace/test_pipeline_len.py << 'EOF'\n\n# Test case 5: Various slicing operations with len\npipe = Pipeline([\n ('step1', StandardScaler()),\n ('step2', SelectKBest(f_regression, k=5)),\n ('step3', svm.SVC(kernel='linear'))\n])\nprint(f\"\\nOriginal pipe length: {len(pipe)}\")\n\n# Slice from beginning to end\nslice1 = pipe[:len(pipe)]\nprint(f\"pipe[:len(pipe)] steps: {[name for name, _ in slice1.steps]}\")\n\n# Slice from beginning to middle\nslice2 = pipe[:len(pipe)//2]\nprint(f\"pipe[:len(pipe)//2] steps: {[name for name, _ in slice2.steps]}\")\n\n# Slice from middle to end\nslice3 = pipe[len(pipe)//2:]\nprint(f\"pipe[len(pipe)//2:] steps: {[name for name, _ in slice3.steps]}\")\n\n# Slice a subset in the middle\nif len(pipe) >= 3:\n slice4 = pipe[1:len(pipe)-1]\n print(f\"pipe[1:len(pipe)-1] steps: {[name for name, _ in slice4.steps]}\")\nEOF\n\n\nLet's run this updated script:\n\n\npython /workspace/test_pipeline_len.py\n\n\nLet's also check all the tests in the sklearn directory to ensure our changes don't break anything:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\npython -m pytest -xvs sklearn/tests/test_pipeline.py\n\n\nBefore we conclude, let's verify that our changes to scikit-learn match the expected functionality:\n\n1. We have implemented the `__len__` method in the Pipeline class\n2. The method returns the number of steps in the pipeline\n3. We have tested that it works for normal, empty, and long pipelines\n4. We have also tested that slicing operations using `len(pipe)` work correctly\n\nSince the issue description mentioned that the repository already supports indexing but not `__len__`, we didn't have to implement indexing. We only needed to add the `__len__` method to make `len(pipe)` work, which we've done.\n\nOur changes are minimal and focused on solving the specific issue described.\n\nFurthermore, to ensure we've done a thorough job, let's run all sklearn tests to make sure we haven't broken anything:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\npython -m pytest sklearn/tests/\n\n\nLet's check the full implementation of the Pipeline to understand where `__len__` fits in compared to other methods like `__getitem__`:\n\n\nhead -n 300 /workspace/scikit-learn__scikit-learn__0.21/sklearn/pipeline.py\n\n\nLet's more specifically check the surrounding methods to ensure our addition fits well with the rest of the class:\n\n\ngrep -A 20 \"__getitem__\" /workspace/scikit-learn__scikit-learn__0.21/sklearn/pipeline.py\n\n\nNow, let's also test our changes against the base commit a62775e99f2a5ea3d51db7160fad783f6cd8a4c5, as mentioned in the requirements:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\ngit status\n\n\nHmm, we don't have access to the git repository directly. Let's focus on thoroughly testing our changes then.\n\nLet's run the full test suite for the sklearn module to make sure our changes don't break any existing functionality:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\npython -m pytest -xvs sklearn/tests/test_pipeline.py\n\n\nGreat, now let's create a more comprehensive test script to ensure our `__len__` implementation works in various scenarios:\n\n\ncat > /workspace/comprehensive_pipeline_len_test.py << 'EOF'\n\"\"\"\nComprehensive tests for the __len__ method of Pipeline.\n\"\"\"\nfrom sklearn.pipeline import Pipeline, FeatureUnion\nfrom sklearn import svm\nfrom sklearn.feature_selection import SelectKBest, f_regression\nfrom sklearn.datasets import samples_generator\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.decomposition import PCA\n\n# Test case 1: Normal pipeline\nprint(\"Test case 1: Normal pipeline\")\nanova_filter = SelectKBest(f_regression, k=5)\nclf = svm.SVC(kernel='linear')\npipe = Pipeline([('anova', anova_filter), ('svc', clf)])\nprint(f\"Length of normal pipeline: {len(pipe)}\")\nassert len(pipe) == 2, \"Length of normal pipeline should be 2\"\n\n# Test case 2: Empty pipeline\nprint(\"\\nTest case 2: Empty pipeline\")\nempty_pipe = Pipeline([])\nprint(f\"Length of empty pipeline: {len(empty_pipe)}\")\nassert len(empty_pipe) == 0, \"Length of empty pipeline should be 0\"\n\n# Test case 3: Long pipeline\nprint(\"\\nTest case 3: Long pipeline\")\nX, y = samples_generator.make_classification(\n n_informative=5, n_redundant=0, random_state=42)\n\nsteps = []\nfor i in range(10):\n steps.append((f'scaler_{i}', StandardScaler()))\nlong_pipe = Pipeline(steps)\nprint(f\"Length of long pipeline: {len(long_pipe)}\")\nassert len(long_pipe) == 10, \"Length of long pipeline should be 10\"\n\n# Test case 4: Pipeline with FeatureUnion\nprint(\"\\nTest case 4: Pipeline with FeatureUnion\")\nunion = FeatureUnion([\n ('pca', PCA(n_components=1)),\n ('select', SelectKBest(k=1))\n])\npipe_with_union = Pipeline([\n ('scaler', StandardScaler()),\n ('union', union),\n ('svc', svm.SVC(kernel='linear'))\n])\nprint(f\"Length of pipeline with FeatureUnion: {len(pipe_with_union)}\")\nassert len(pipe_with_union) == 3, \"Length of pipeline with FeatureUnion should be 3\"\n\n# Test case 5: Various slicing operations with len\nprint(\"\\nTest case 5: Various slicing operations with len\")\npipe = Pipeline([\n ('step1', StandardScaler()),\n ('step2', SelectKBest(f_regression, k=5)),\n ('step3', svm.SVC(kernel='linear'))\n])\nprint(f\"Original pipe length: {len(pipe)}\")\nassert len(pipe) == 3, \"Original pipe length should be 3\"\n\n# Slice from beginning to end\nslice1 = pipe[:len(pipe)]\nprint(f\"pipe[:len(pipe)] steps: {[name for name, _ in slice1.steps]}\")\nassert len(slice1) == 3, \"Length of pipe[:len(pipe)] should be 3\"\n\n# Slice from beginning to middle\nslice2 = pipe[:len(pipe)//2]\nprint(f\"pipe[:len(pipe)//2] steps: {[name for name, _ in slice2.steps]}\")\nassert len(slice2) == 1, \"Length of pipe[:len(pipe)//2] should be 1\"\n\n# Slice from middle to end\nslice3 = pipe[len(pipe)//2:]\nprint(f\"pipe[len(pipe)//2:] steps: {[name for name, _ in slice3.steps]}\")\nassert len(slice3) == 2, \"Length of pipe[len(pipe)//2:] should be 2\"\n\n# Slice a subset in the middle\nslice4 = pipe[1:len(pipe)-1]\nprint(f\"pipe[1:len(pipe)-1] steps: {[name for name, _ in slice4.steps]}\")\nassert len(slice4) == 1, \"Length of pipe[1:len(pipe)-1] should be 1\"\n\nprint(\"\\nAll tests passed!\")\nEOF\n\n\nLet's run this comprehensive test:\n\n\npython /workspace/comprehensive_pipeline_len_test.py\n\n\nLet's also run specific tests from the scikit-learn test suite that relate to Pipeline indexing to ensure our changes integrate well:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\npython -m pytest -xvs sklearn/tests/test_pipeline.py::test_pipeline_indexing\n\n\nOh, there might not be a specific test named `test_pipeline_indexing`. Let's look at the available tests in the test_pipeline.py file:\n\n\ngrep -r \"def test_\" /workspace/scikit-learn__scikit-learn__0.21/sklearn/tests/test_pipeline.py\n\n\nLet's run the tests that seem related to pipeline indexing or slicing:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\npython -m pytest -xvs sklearn/tests/test_pipeline.py::test_pipeline_indexing || python -m pytest -xvs sklearn/tests/test_pipeline.py | grep -i \"index\"\n\n\nSince we might not have direct access to specific test names, let's examine the pipeline test file to identify relevant tests:\n\n\ncat /workspace/scikit-learn__scikit-learn__0.21/sklearn/tests/test_pipeline.py | grep -i \"index\\|slicing\\|getitem\"\n\n\nBased on our findings, let's run specific tests that are relevant to our changes:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\npython -m pytest -xvs sklearn/tests/test_pipeline.py::test_pipeline_indexing || python -m pytest -xvs sklearn/tests/test_pipeline.py\n\n\nLet's check if the test_pipeline.py file contains any tests for slicing functionality:\n\n\ncat /workspace/scikit-learn__scikit-learn__0.21/sklearn/tests/test_pipeline.py | grep -A 20 -B 20 \"slicing\\|__getitem__\"\n\n\nLet's modify our comprehensive test script to include more edge cases and ensure it aligns with any existing tests:\n\n\ncat >> /workspace/comprehensive_pipeline_len_test.py << 'EOF'\n\n# Test case 6: Testing on pipeline with make_pipeline\nprint(\"\\nTest case 6: Testing on pipeline with make_pipeline\")\nfrom sklearn.pipeline import make_pipeline\n\npipe = make_pipeline(StandardScaler(), PCA(n_components=2), SelectKBest(k=1))\nprint(f\"Length of make_pipeline: {len(pipe)}\")\nassert len(pipe) == 3, \"Length of make_pipeline should be 3\"\n\n# Test case 7: Testing on nested pipelines\nprint(\"\\nTest case 7: Testing on nested pipelines\")\ninner_pipe = Pipeline([('scaler', StandardScaler()), ('pca', PCA(n_components=2))])\nouter_pipe = Pipeline([('inner', inner_pipe), ('svc', svm.SVC())])\nprint(f\"Length of outer pipeline: {len(outer_pipe)}\")\nassert len(outer_pipe) == 2, \"Length of outer pipeline should be 2\"\n\n# Test case 8: Testing __len__ with negative indexing\nprint(\"\\nTest case 8: Testing __len__ with negative indexing\")\npipe = Pipeline([\n ('step1', StandardScaler()),\n ('step2', PCA(n_components=2)),\n ('step3', SelectKBest(k=1))\n])\nneg_slice = pipe[-len(pipe):]\nprint(f\"pipe[-len(pipe):] steps: {[name for name, _ in neg_slice.steps]}\")\nassert len(neg_slice) == 3, \"Length of pipe[-len(pipe):] should be 3\"\n\nneg_slice2 = pipe[:-len(pipe)]\nprint(f\"pipe[:-len(pipe)] steps: {[name for name, _ in neg_slice2.steps]}\")\nassert len(neg_slice2) == 0, \"Length of pipe[:-len(pipe)] should be 0\"\n\n# Test case 9: Testing complex slicing expressions with len\nprint(\"\\nTest case 9: Testing complex slicing expressions with len\")\npipe = Pipeline([\n ('step1', StandardScaler()),\n ('step2', PCA(n_components=2)),\n ('step3', SelectKBest(k=1)),\n ('step4', svm.SVC())\n])\ncomplex_slice = pipe[len(pipe)//4:len(pipe)//4*3]\nprint(f\"pipe[len(pipe)//4:len(pipe)//4*3] steps: {[name for name, _ in complex_slice.steps]}\")\nassert len(complex_slice) == 2, \"Length of complex slice should be 2\"\n\nprint(\"\\nAll additional tests passed!\")\nEOF\n\n\nLet's run the updated comprehensive test:\n\n\npython /workspace/comprehensive_pipeline_len_test.py\n\n\nLet's also verify that our implementation doesn't cause any regressions in the scikit-learn test suite:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\npython -m pytest -xvs sklearn/tests/test_pipeline.py\n\n\nNow, let's also run another specific test related to pipeline functionality:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\npython -m pytest -xvs sklearn/tests/test_metaestimators.py | grep -i \"pipeline\"\n\n\nLet's get a deeper understanding of how the Pipeline class is used in tests by examining test files containing pipeline tests:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\ngrep -r \"Pipeline\" --include=\"*.py\" sklearn/tests/ | head -n 20\n\n\nLet's also look for any tests related to pipeline length or slicing:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\ngrep -r \"len(\" --include=\"*.py\" sklearn/tests/ | grep -i \"pipeline\" | head -n 20\n\n\nWith all our tests passing, it appears we've successfully implemented the `__len__` method for the Pipeline class, and it behaves as expected in various scenarios.\n\nLet's summarize what we've done:\n\n1. We located the Pipeline class in sklearn/pipeline.py\n2. We identified that it was missing a `__len__` method, which was causing the error described in the issue\n3. We implemented the `__len__` method to return the length of the pipeline's steps list\n4. We verified our fix with a reproduction script and comprehensive tests\n5. We confirmed that existing scikit-learn tests still pass with our changes\n\nOur implementation is minimal and focused on the specific issue, adding only the `__len__` method without changing any existing functionality.\n\nAt this point we can conclude that the issue has been successfully resolved. The Pipeline class now supports `len(pipe)`, which allows expressions like `pipe[:len(pipe)]` to work correctly.\n\nLet's make sure we've run the specific test for pipeline indexing:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\npython -m pytest -xvs sklearn/tests/test_pipeline.py | grep -i \"indexing\"\n\n\nLet's also check if there are any tests related to `__len__` or length in the pipeline tests:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\ngrep -r \"__len__\\|len(\" --include=\"*.py\" sklearn/tests/test_pipeline.py\n\n\nWith all our checks and tests, we're confident that our implementation of the `__len__` method in the Pipeline class correctly resolves the issue described. The method returns the number of steps in the pipeline, which is the expected behavior for `len(pipe)`.\n\nI believe our solution is complete and correct. The changes we made are minimal and focused on the specific issue, and we've thoroughly tested them to ensure they work as expected in various scenarios.\n\nLet's complete the task by running our comprehensive test one more time:\n\n\npython /workspace/comprehensive_pipeline_len_test.py\n\n\nI think we're good to go. Our implementation of the `__len__` method in the Pipeline class resolves the issue described, allowing expressions like `pipe[:len(pipe)]` to work correctly.\n\nLet's do a final check on the changes we made. The diff should show just the addition of the `__len__` method to the Pipeline class:\n\n\ncd /workspace/scikit-learn__scikit-learn__0.21\ndiff /workspace/pipeline.py.backup sklearn/pipeline.py\n\n\nWith all our checks and tests passing, we've successfully resolved the issue described in the problem statement.\n\nAlright, let's summarize our solution once more:\n\nThe issue was that Pipeline in scikit-learn didn't implement the `__len__` method, which was causing `pipe[:len(pipe)]` to raise an error, even though Pipeline supported indexing.\n\nTo resolve this, we added a `__len__` method to the Pipeline class that returns the number of steps in the pipeline. This allows expressions like `len(pipe)` and `pipe[:len(pipe)]` to work correctly, providing a more consistent and intuitive API for users.\n\nWe thoroughly tested our solution with:\n1. A script to reproduce the original error and confirm our fix\n2. Comprehensive tests for various scenarios, including edge cases\n3. Running the existing scikit-learn test suite to ensure no regressions\n\nThese tests confirm that our solution correctly resolves the issue while maintaining compatibility with the rest of the codebase.\n\nAnd that concludes our implementation.", - "signature": "ErUBCkYIARgCIkDnKEqW20HOxB9jS85ReLznKuKSV88wnAJf/M+y9ge7U+GOWAHd/trgmvpgHpJmOBkeUWUftR5g+81TNmAsDd2EEgxQHv0uMQBMvelCR+saDFvnSSfaNOj3NTVYpyIwG1ddKj5AVnXe382HERyVchmOTdn1ebogEFpXlRE66Mckq57xLSygusq6MwAV/6RvKh28HEuwcXNQn+BPIxQDVCwGGUh5g4fpZ6IgGk3mRg==", - }, - { - "type": "text", - "text": "I'll help you implement the `__len__` method for Pipeline in scikit-learn to fix the issue described. Let's start by exploring the repository and understanding the problem better.\n\n## Step 1: Explore the repository structure\n\nLet's first look at the repository structure:", - }, - ], - }, - {"role": "user", "content": [{"type": "text", "text": "Who do you know?"}]}, - ], - "max_tokens": 32768, - "thinking": {"type": "enabled", "budget_tokens": 30720}, - } - - response = litellm.completion(**params) - - assert response is not None - - -@pytest.mark.parametrize( - "model", - [ - "anthropic/claude-sonnet-4-5-20250929", - # "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - ], -) -def test_anthropic_redacted_thinking_in_assistant_message(model): - litellm._turn_on_debug() - params = { - "model": model, - "messages": [ - { - "role": "assistant", - "content": [ - { - "type": "redacted_thinking", - "data": "EqkBCkYIARgCKkAflgFkky5bvpaXt2GnDYgbA8QOCr+BF53t+UmiRA22Z7Ply9z2xfTGYSqvjlhIEsV6WDPdVoXndztvhKCzE2PUEgxwXpRD1hBLUSajVWoaDEftxmhqdg0mRwPUGCIwcht1EH91+gznPoaMNquU4sGeaOLFaeyNeG4dJXsYT/Jc4OG3453LN5ra4uVxC/GgKhGMQ1A9aO2Ac0O5M+bOdp1RFw==Eo0CCkYIARgCKkCcHATldbjR0vfU1DlNaQr3J2GKem6OjFybQyshp4C9XnysT/6y1CNcI+VGsbX99GfKLGqcsGYr81WlM+d7NscJEgxzkyZuwL3QnnxFiUUaDIA3nZpQa15D5XD72yIwyIGpJwhdavzXvE1bQLZj43aNtznG6Uwsxx4ZlLv83SUqH7GqzMxvm3stLj3cYmKMKnUqqhpeluvoxODUY/fhhF6Bjsj9C1MIRL+9urDH2EtAmZ+BrvLoXjRlbEH9+DtzLE57I1ShMDbUqLJXxXTcjhPkmu3JscBYf0waXfUgrQl2Pnv5dAxM2S3ZASk8di7ak0XcRknVBhhaR2ykdDbVyxzFzyZo8Fc=EtcBCkYIARgCKkCl6nQeKqHIBgdZ1EByLfEwnlZxsZWoDwablEKqRAIrKvB10ccs6RZqrTMZgcMLaW3QpWwnI4fC/WiOe811B94JEgyvTK4+E/zB+a42bYcaDOPesimKdlIPLT7VQiIwplWjvDcbe16vZSJ0OezjHCHEvML4QJPyvGE3NRHcLzC9UiGYriFys5zgv0O7qKr5Kj/56IL1BbaFqSANA7vjGoW+GSlv294L4LzqNWCD0ANzDnEjlXlVeibNM74v+KKXRVwn/IInHPog4hJA0/3GQyA=EtwBCkYIARgCKkBda4XEzq+PTfE7niGdYVzvAXRTb+3ujsDVGhVNtFnPx6K/I6ORfxOWmwEuk7iXygehQA18p0CVYLsCU4AHFvtjEgzYH2JNCxa8F07pGioaDOA635mdHKbyiecBJSIwshUavES7HZBnA4l3k8l92LAhuJQV1C5tUgKkk0pHRT+/OzDfXvxsZSx7AmR7J3QXKkQwHL6K9yZEWdeh/B22ft/GxyRViO7nZrT95PAAux31u++rYQyeFJ+rv0Yrs/KoBnlNUg9YFOpDMo1bMWV9n4CGwq92bw==EtEBCkYIARgCKkCZdn2NBzxiOEJt/E8VOs6YLbYjRaCkvhEdz5apcEZlBQJpulvgv1JvamrMZD0FCJZVTwxd/65M9Ady/LbtYTh7EgwtL7W9DXSFjxPErCIaDGk0e/bXY8yJdjk3CSIwYS0TtiaFK8tJrREBFA9IOp+q+tnE8Wl338CbbskRvF5topYmtofuBIG4GQkHvbQjKjn2BmwrEic/CdSEVbvEix7AWEsw92DabVmseTQhUbbuYRa4Ou6jXMW2pMJFUBjMr95gF6BlVFr4iEA=EsUBCkYIARgCKkAsEmKjMN9TVYLyBdo1+0uopommcjQx8Fu65+mje5Ft05KOnyKAzuUyORtk5r73glan8L+WlygaOOrZ1hi81219EgwpdTA6qbcaggIWeTIaDDrJ0eTbsqku4VSY8CIw3mJfRyv7ISHih4mpAVioGuuduXbaie5eKn5a+WgQiOmm22uZ4Gv72uluCSGGriHnKi28bHMomrytYLvKNvhL51yf5/Tgm/lIgQ9gyTJLqVzVjGn6ng1sN8vUti/tuGw=EsoBCkYIARgCKkB+jJBrxqqpzyGt5RXDKTBVxTnE8IrYRysAL2U/H171INDMCxrDHxfts3M0wuQirXN/2fZXwmQJIZRzzumA+I2sEgw0ySDeyTfHgTiafo8aDKOTl485koQiPwXipyIwG9n/zWUZ+tgfFELW2rV5/yo6Pq/r9bJdrd2b25qCATwX2gd54gsjWhSvLDkD7pLJKjL6ZuiW4N6hVo6JIR4UL8LxcsP9tET0ElIgQZ/h8HOIi18fQKsEdtseWCFnuXse21KIeg==EtwBCkYIARgCKkDWMlgTA+iKsScbpNtZab6dgMKRZYpQSoJ274+n0TqvLAqHL8GxLm1sMVom81LcVWCZZeIVQFbkmbJxyBovvLoUEgxy6YGb0EeJW10P8XEaDKowL3qI/z000pgR2SIwZIczlDKkqw75UYcEOC6Cx9yc0CdYjJnmQOa4Ezni20SANA8YnBMIYJqW4osO/KalKkTLmgvJRQE1Hk8Bn3af9fIYt+vITYEY4Wr7/UVNBtSXBOMP0YoSgNyzjX/pu2N3oy2Blv/YAgtHIJ3Xwd43clN5F2wU+Q==EtQBCkYIARgCKkD3vxW2GsLyEGtmBpI6NdNyh4i/ea7E9rp5puSHdk/dSCpW5G1wI3nrFIS2bUqZsvsDu3YgcDixG8eeDnzacC/qEgzilh/V8vaE1X9lRlIaDAa17eq6kSgaRrsAfSIwFAXgLu5BUKldMeQdcomRqgmY9hDzkDlRnBrbO9GxXsrmpGTU9iqVZQ7z9OVW522bKjyB/GeuNlv4V8a8uricx1InN8q94coWGCRPvAJVAvhP/YMCcNlvrgoN8C2RGc13e88uDq01r6gpkWTlVDY=EssBCkYIARgCKkAOhKBpvfqIElQ1mlG7NiCiolHnqagXryuwNsODnttLBeVMGBsZ8DgpSGWonVE/22MQgciWLY7WaaeoDcpL3X/pEgx4xuL/KqOgxrBnau4aDH3pQ/Sqr1aHa68YiiIwR6+w9QOWFfut8ZG8z+QkAO/kZVePcELKabHp7ikY+DOjvOt4FfnaChwQFTSGzZhaKjPK4MwQukuZIT1PFGFIh20Hi6wMQlHvsChIF88nUV2EAz4Sgb/vWPiQBbWP3gT3hJBehQY=EtMBCkYIARgCKkCT0yD5m4Rvs3KBNkAC2g7aprLTzKRqF+vdHAeYte9KngJZhThexj65o+q9HOGhIIAsboRhz70xkAybdQdsrg8OEgzQm1M980FeZMCi1XsaDJSFOpIuOhUOkPIs+iIw62jO5yY9ZETmrYtEb+pYN5Cyf467YVOOv7FBo44gIFgUvFklU5+y09k3MGzrBNViKjvkopPoFbpYI9ilB3dN6pAzrzhDzOum+Rsx1N25+UYvdT+yYBilrIPW1XmLmzT+ZMs4eV5caG35ZsNsjQ==EtwBCkYIARgCKkCOShz0/2ZO3u0WH8PBN63fAwKo4TcNFM3axUJL9dK9JJDLtC0XwP9Ee4vqPZyLBao4RyAefbYmY3TJ1As/AbuvEgxbYiyN4UcjaJU9mwkaDP9L3FACdMRQ+UFOSSIwQ0btU6cKIRsSNzvBsP8Fa4Ab7vOnlo4YSAv2lD7ZdDKVcQaWQZHYsQb/QQDfIGKGKkRXhNoET9KyQkb/x8lVpUR1d2u/sHTdgKEjkUdQop88SUFHvkGcJrMUTvnuvUdO4MdHwKnN0IINbDHTEUjUXSQPkpfTTA==EtwBCkYIARgCKkCIwQCFJUrhd1aT8hGMNcPIl+CaSZWsqerPDUGzZnS2tt2+tAs+TAPcKVHC07BdEXj6aKSbrOb8b7OQ/KFbrWJ4Egz980omEnE4djm8t5UaDDXrDJWgFSuZ+LWFmSIw/RzMo5ncKnqvf0TZ1krxMi4/DpAZb0Lgmc1XxGT2JPA4At9EEHNVPrWLXwGM3vUYKkQltG8EJFOWL1In5541dca1pnRDyBg4JVRQ5CuvA/pUCI2e9ARiODI7D+ydZorcnWQ7j2Qc1DguMQVHMbPLyGbQx9vqgQ==EtsBCkYIARgCKkDiH+ww5G0OgaW7zSQD7ZKYdViZfi+KO+TkA/k4rlTKsIwpUILZZ/53ppu93xaEazsD92GXKKSG3B/jBCqjQRg7EgzR3K/BJFTt359xPOgaDEHyoGVloiLS71ufAiIwO77B26VivdVgd2Dmv3DOtUAFs/jDwLM9EmNCBeoivwJPD2hYEKNm6TUWTinGfO2jKkNbrYgpA5esB0y1iXA0qGwRAmnD8ykZc0DT40vvd9EDvb5gHCd7RyjEU9BKnXBPWpGdTi4U+LZKYQ9LEE6sJ8vBm8w3EtUBCkYIARgCKkBbxQIjnTzzKf8Qhfcu+so91+MMbpJNyga27D9tZBtTexYLMJtzDWux4urfCc5TjjX0MvK62lKkhcPLuJE7KiI8EgzFF+TlNgPNp6RoyQgaDBAUDEAsqBMj7z4kciIwUWEZMGkG8ZnjltVpuffHxw5Rqyc+Smh1MnqnWxo0JlCOC43W5JH5KoJ/4RDxX7IjKj2fs5F6eiRMEi+L4KyjDBIvoPoE/wrdC+Fo6c8lMJiYw0MJ/lXgJQv6p0GRe251X+pcfN+2lx067/GLP6qjEtsBCkYIARgCKkCItf9nN0FKJsetom0ZoZvccwboNM2erGP7tIAYsOzsA9lmh7rFI2mFbOOC2WZ1v+QkvxppQ2wO+N35t29LC7RPEgzyJgiM1GHTVN+VPPwaDOXyzSg9BQ85oi58DCIwu/JxKJwVECkbru1d05yhwMYDsJrSJW1BO2ZBrg8Tb48S+dpD6hEPd1itq8cSM3ChKkNv83rGY8Gjg2DiTWDsIqUCD0pb2drrwnjkherr5/EQWdhHC7MijF8zyvqU4tBZrxP+64GcII7P87ja8B4YxGUIw9J7Et0BCkYIARgCKkCInOjYRgGSjcV/WHJ6HjB983rvz/nrOZ9xZMdrTYdHURtXN4zMAjZYQ8ZBk31n4aFGv5PAtDfbjqcytZUaCKicEgwXQrjgS0FHWq/2PwAaDKjYgoXuPPq+RNJUvCIwh1VmSiLGu+3pl7RcCBxnH/ue38EUDZAIRYiDI59h8CVdZpDSqaH8yJvFlR5Jxc8xKkXcEPduWcuONY+vatnIo5AQeSh9HM4oM4DoDma1OvVfdPUpbvaTP3ZhEv4iOMjvwzHBBkvc8b9jV2oTb8Xe50COLFJvURk=EtcBCkYIARgCKkDM4CyfgVBHhusU4C0tg/RwXiAbNtjOoYfcufGUnFlQKcpuJnekvb61EAerBrELguIrvNJIbyqy0Kcd/r64hu1UEgyITWjG3/cVsm/o0JkaDKm1/y0HF1YpqoiFoCIwqImOpk6SngP99aXE4p5c7y9rOvVo3lmKidTUdi1lmtoEZ9sXdY49nLsGeCuCjPJKKj976uFmgrZWIEZIL+HQGVjDOJ7mK8NzAxjX3m0AELsWN5FgbGOHus/S4o2EKi43/MLaRervgaFdrxK9BKGE6LY=EtMBCkYIARgCKkDvEoH/lv1fRxN+JaknzdY53WmQrEGJ7yupv22X2TdxN2+GmY8l1KYONWboOxalfoSbSlp3+zVJXdvTCa60CYnnEgyUslgNTFL5iGt+aq0aDESsIoNRuPYqDc5fbCIw9gHGejHXKw9GMR0sw1RnIF2FBI5Zo5/4EK2AFZ8BU5yAYgJw0wTc16ZVEFEraKS+KjtqVPmiodedFzc+f4kr+U8dy+xQtcsmTe9KcvAYmskvZ6Kl6iCitm/PZdjl/7COePcTVu32QnxZuG4Mpw==EtEBCkYIARgCKkB/SdSv2Jo8DJ4pOOK4mYXhSsPrnf6/ESHL7voj6FbdYPsgg2f3XQByQV93Menel5tgcx0jvNfY7Z9nx4Rz3iTvEgxN/mWUwb6Lb/1BfkAaDBONEsjWD1fKeK8H/iIwy+yJUFPTde2wxI/j6em5uS8HWGsfX9pUB4u/K4QHAd85bn63rrXSxbe2DHIG620UKjk+C6q3aXztOAGAyvhjiN9lnNAFPv93GTnwj+14n07c/xPdHBQyXXi742UBjFdQkmwp3m6RWf5psYU=EuQBCkYIARgCKkBxavD9zRmeX22ltvtCNzZzXTpsAHmNwSuejX7ibJueaDQaSOykBjNJavdMn6yQ8mAxCpNrNmhtBhGxHBGZE668EgzFNqHVE2WctK5ZiN0aDGNFTI5T3/0vDCtFXiIwRDXV5+9nWYGzuih8cG8h4dCs+n90rcL/Tz78QKsfpZeLNpr4aZSU8KHO2OmcmFoOKkxdgzKPy/gOfcCELsudlawbVyobU4CIhOYacIPhi+0XvgjXpqP0JIANaOdawb2zWrKhBKNA4VCHzbFkDm9cV1WrGIw0cEJ3oRU7idRgEsEBCkYIARgCKkDJUpJz2Ct4ZZJlWkAGg1Lc/rVqCd/V5rq01yehv9GkTIaq9H2jgjVKnUV1e4o9F1cUxmMk6fn4XK01sp/szP2GEgyvuemo2Di0USGKingaDCAMXK1kWRk6KofoyyIwxr/Jdwz2RrUytRWMGjrs4MkcQ2rhrVL/00Ktebga9cwrqeDOq+7nN8L64V+XEwsJKimHdmpCQPqYz8rIX25+v2XqcBDXzoBW8+eqdJKRhKcYooLbBXK3DUgRVQ==", - }, - { - "type": "text", - "text": "I'm not able to respond to special commands or trigger phrases like the one you've shared. Those types of strings don't activate any special modes or features in my system. Is there something specific I can help you with today? I'm happy to assist with questions, have a conversation, provide information, or help with various tasks within my normal capabilities.", - }, - ], - }, - {"role": "user", "content": [{"type": "text", "text": "Who do you know?"}]}, - ], - "max_tokens": 32768, - "thinking": {"type": "enabled", "budget_tokens": 30720}, - } - - response = litellm.completion(**params) - - assert response is not None - - -def test_just_system_message(): - litellm._turn_on_debug() - litellm.modify_params = True - params = { - "model": "anthropic/claude-sonnet-4-5-20250929", - "messages": [{"role": "system", "content": "You are a helpful assistant."}], - } - - response = litellm.completion(**params) - - assert response is not None - - @pytest.mark.parametrize( "model", ["anthropic/claude-3-sonnet-20240229", "anthropic/claude-3-opus-20240229"], @@ -1772,32 +1664,6 @@ def test_anthropic_strict_not_present(): assert "strict" not in tool["input_schema"] -def test_anthropic_structured_output_chat_completion_api(): - response = litellm.completion( - model="claude-sonnet-4-5-20250929", - messages=[{"role": "user", "content": "What is the capital of France?"}], - response_format={ - "type": "json_schema", - "json_schema": { - "name": "final_output", - "strict": True, - "schema": { - "description": 'Progress report for the thinking process\n\nThis model represents a snapshot of the agent\'s current progress during\nthe thinking process, providing a brief description of the current activity.\n\nAttributes:\n agent_doing: Brief description of what the agent is currently doing.\n Should be kept under 10 words. Example: "Learning about home automation"', - "properties": { - "agent_doing": {"title": "Agent Doing", "type": "string"} - }, - "required": ["agent_doing"], - "title": "ThinkingStep", - "type": "object", - "additionalProperties": False, - }, - }, - }, - ) - assert response is not None - print(f"response: {response}") - - def _make_transform_request(optional_params: dict, litellm_params: dict) -> dict: from litellm.llms.anthropic.chat.transformation import AnthropicConfig diff --git a/tests/llm_translation/test_azure_ai.py b/tests/llm_translation/test_azure_ai.py index 5be6ade80ab..f00409f280b 100644 --- a/tests/llm_translation/test_azure_ai.py +++ b/tests/llm_translation/test_azure_ai.py @@ -270,7 +270,7 @@ async def test_azure_ai_request_format(): @pytest.mark.asyncio -@pytest.mark.parametrize("model", ["azure/gpt5_series/gpt-5-mini", "azure/gpt-5-mini"]) +@pytest.mark.parametrize("model", ["azure/gpt5_series/gpt-5-mini"]) async def test_azure_gpt5_reasoning(model): litellm._turn_on_debug() response = await litellm.acompletion( diff --git a/tests/llm_translation/test_azure_o_series.py b/tests/llm_translation/test_azure_o_series.py index 7a223739844..2ee9bdb2be2 100644 --- a/tests/llm_translation/test_azure_o_series.py +++ b/tests/llm_translation/test_azure_o_series.py @@ -11,6 +11,10 @@ from base_llm_unit_tests import BaseLLMChatTest, BaseOSeriesModelsTest class TestAzureOpenAIO3Mini(BaseOSeriesModelsTest, BaseLLMChatTest): + test_content_list_handling = None + test_empty_tools = None + test_function_calling_with_tool_response = None + def get_base_completion_call_args(self): # Clear the LLM client cache to prevent test pollution from cached clients litellm.in_memory_llm_clients_cache.flush_cache() diff --git a/tests/llm_translation/test_azure_openai.py b/tests/llm_translation/test_azure_openai.py index e6528e77749..df1892638b0 100644 --- a/tests/llm_translation/test_azure_openai.py +++ b/tests/llm_translation/test_azure_openai.py @@ -729,18 +729,3 @@ def test_azure_with_content_safety_error(): ] == "high" ) - - -def test_azure_openai_with_prompt_cache_key(): - """ - E2E test for Azure OpenAI with prompt cache key param on /chat/completions API. - """ - litellm._turn_on_debug() - response = litellm.completion( - model="azure/gpt-4.1-mini", - api_key=os.getenv("AZURE_AI_API_KEY"), - api_base=os.getenv("AZURE_AI_API_BASE"), - api_version="2024-12-01-preview", - messages=[{"role": "user", "content": "What is the weather in San Francisco?"}], - prompt_cache_key="test_streaming_azure_openai", - ) diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 00744da2185..4161e08235b 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -425,55 +425,6 @@ def test_completion_bedrock_claude_aws_bedrock_client(bedrock_session_token_cred # test_completion_bedrock_claude_sts_client_auth() -@pytest.mark.parametrize( - "image_url", - [ - "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAL0AAAC9CAMAAADRCYwCAAAAh1BMVEX///8AAAD8/Pz5+fkEBAT39/cJCQn09PRNTU3y8vIMDAwzMzPe3t7v7+8QEBCOjo7FxcXR0dHn5+elpaWGhoYYGBivr686OjocHBy0tLQtLS1TU1PY2Ni6urpaWlpERER3d3ecnJxoaGiUlJRiYmIlJSU4ODhBQUFycnKAgIDBwcFnZ2chISE7EjuwAAAI/UlEQVR4nO1caXfiOgz1bhJIyAJhX1JoSzv8/9/3LNlpYd4rhX6o4/N8Z2lKM2cURZau5JsQEhERERERERERERERERERERHx/wBjhDPC3OGN8+Cc5JeMuheaETSdO8vZFyCScHtmz2CsktoeMn7rLM1u3h0PMAEhyYX7v/Q9wQvoGdB0hlbzm45lEq/wd6y6G9aezvBk9AXwp1r3LHJIRsh6s2maxaJpmvqgvkC7WFS3loUnaFJtKRVUCEoV/RpCnHRvAsesVQ1hw+vd7Mpo+424tLs72NplkvQgcdrsvXkW/zJWqH/fA0FT84M/xnQJt4to3+ZLuanbM6X5lfXKHosO9COgREqpCR5i86pf2zPS7j9tTj+9nO7bQz3+xGEyGW9zqgQ1tyQ/VsxEDvce/4dcUPNb5OD9yXvR4Z2QisuP0xiGWPnemgugU5q/troHhGEjIF5sTOyW648aC0TssuaaCEsYEIkGzjWXOp3A0vVsf6kgRyqaDk+T7DIVWrb58b2tT5xpUucKwodOD/5LbrZC1ws6YSaBZJ/8xlh+XZSYXaMJ2ezNqjB3IPXuehPcx2U6b4t1dS/xNdFzguUt8ie7arnPeyCZroxLHzGgGdqVcspwafizPWEXBee+9G1OaufGdvNng/9C+gwgZ3PH3r87G6zXTZ5D5De2G2DeFoANXfbACkT+fxBQ22YFsTTJF9hjFVO6VbqxZXko4WJ8s52P4PnuxO5KRzu0/hlix1ySt8iXjgaQ+4IHPA9nVzNkdduM9LFT/Aacj4FtKrHA7iAw602Vnht6R8Vq1IOS+wNMKLYqayAYfRuufQPGeGb7sZogQQoLZrGPgZ6KoYn70Iw30O92BNEDpvwouCFn6wH2uS+EhRb3WF/HObZk3HuxfRQM3Y/Of/VH0n4MKNHZDiZvO9+m/ABALfkOcuar/7nOo7B95ACGVAFaz4jMiJwJhdaHBkySmzlGTu82gr6FSTik2kJvLnY9nOd/D90qcH268m3I/cgI1xg1maE5CuZYaWLH+UHANCIck0yt7Mx5zBm5vVHXHwChsZ35kKqUpmo5Svq5/fzfAI5g2vDtFPYo1HiEA85QrDeGm9g//LG7K0scO3sdpj2CBDgCa+0OFs0bkvVgnnM/QBDwllOMm+cN7vMSHlB7Uu4haHKaTwgGkv8tlK+hP8fzmFuK/RQTpaLPWvbd58yWIo66HHM0OsPoPhVqmtaEVL7N+wYcTLTbb0DLdgp23Eyy2VYJ2N7bkLFAAibtoLPe5sLt6Oa2bvU+zyeMa8wrixO0gRTn9tO9NCSThTLGqcqtsDvphlfmx/cPBZVvw24jg1LE2lPuEo35Mhi58U0I/Ga8n5w+NS8i34MAQLos5B1u0xL1ZvCVYVRw/Fs2q53KLaXJMWwOZZ/4MPYV19bAHmgGDKB6f01xoeJKFbl63q9J34KdaVNPJWztQyRkzA3KNs1AdAEDowMxh10emXTCx75CkurtbY/ZpdNDGdsn2UcHKHsQ8Ai3WZi48IfkvtjOhsLpuIRSKZTX9FA4o+0d6o/zOWqQzVJMynL9NsxhSJOaourq6nBVQBueMSyubsX2xHrmuABZN2Ns9jr5nwLFlLF/2R6atjW/67Yd11YQ1Z+kA9Zk9dPTM/o6dVo6HHVgC0JR8oUfmI93T9u3gvTG94bAH02Y5xeqRcjuwnKCK6Q2+ajl8KXJ3GSh22P3Zfx6S+n008ROhJn+JRIUVu6o7OXl8w1SeyhuqNDwNI7SjbK08QrqPxS95jy4G7nCXVq6G3HNu0LtK5J0e226CfC005WKK9sVvfxI0eUbcnzutfhWe3rpZHM0nZ/ny/N8tanKYlQ6VEW5Xuym8yV1zZX58vwGhZp/5tFfhybZabdbrQYOs8F+xEhmPsb0/nki6kIyVvzZzUASiOrTfF+Sj9bXC7DoJxeiV8tjQL6loSd0yCx7YyB6rPdLx31U2qCG3F/oXIuDuqd6LFO+4DNIJuxFZqSsU0ea88avovFnWKRYFYRQDfCfcGaBCLn4M4A1ntJ5E57vicwqq2enaZEF5nokCYu9TbKqCC5yCDfL+GhLxT4w4xEJs+anqgou8DOY2q8FMryjb2MehC1dRJ9s4g9NXeTwPkWON4RH+FhIe0AWR/S9ekvQ+t70XHeimGF78LzuU7d7PwrswdIG2VpgF8C53qVQsTDtBJc4CdnkQPbnZY9mbPdDFra3PCXBBQ5QBn2aQqtyhvlyYM4Hb2/mdhsxCUen04GZVvIJZw5PAamMOmjzq8Q+dzAKLXDQ3RUZItWsg4t7W2DP+JDrJDymoMH7E5zQtuEpG03GTIjGCW3LQqOYEsXgFc78x76NeRwY6SNM+IfQoh6myJKRBIcLYxZcwscJ/gI2isTBty2Po9IkYzP0/SS4hGlxRjFAG5z1Jt1LckiB57yWvo35EaolbvA+6fBa24xodL2YjsPpTnj3JgJOqhcgOeLVsYYwoK0wjY+m1D3rGc40CukkaHnkEjarlXrF1B9M6ECQ6Ow0V7R7N4G3LfOHAXtymoyXOb4QhaYHJ/gNBJUkxclpSs7DNcgWWDDmM7Ke5MJpGuioe7w5EOvfTunUKRzOh7G2ylL+6ynHrD54oQO3//cN3yVO+5qMVsPZq0CZIOx4TlcJ8+Vz7V5waL+7WekzUpRFMTnnTlSCq3X5usi8qmIleW/rit1+oQZn1WGSU/sKBYEqMNh1mBOc6PhK8yCfKHdUNQk8o/G19ZPTs5MYfai+DLs5vmee37zEyyH48WW3XA6Xw6+Az8lMhci7N/KleToo7PtTKm+RA887Kqc6E9dyqL/QPTugzMHLbLZtJKqKLFfzVWRNJ63c+95uWT/F7R0U5dDVvuS409AJXhJvD0EwWaWdW8UN11u/7+umaYjT8mJtzZwP/MD4r57fihiHlC5fylHfaqnJdro+Dr7DajvO+vi2EwyD70s8nCH71nzIO1l5Zl+v1DMCb5ebvCMkGHvobXy/hPumGLyX0218/3RyD1GRLOuf9u/OGQyDmto32yMiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIv7GP8YjWPR/czH2AAAAAElFTkSuQmCC", - "https://avatars.githubusercontent.com/u/29436595?v=", - ], -) -def test_bedrock_claude_3(image_url): - try: - litellm.set_verbose = True - data = { - "max_tokens": 100, - "stream": False, - "temperature": 0.3, - "messages": [ - {"role": "user", "content": "Hi"}, - {"role": "assistant", "content": "Hi"}, - { - "role": "user", - "content": [ - {"text": "describe this image", "type": "text"}, - { - "image_url": { - "detail": "high", - "url": image_url, - }, - "type": "image_url", - }, - ], - }, - ], - } - response: ModelResponse = completion( - model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - num_retries=3, - **data, - ) # type: ignore - # Add any assertions here to check the response - assert len(response.choices) > 0 - assert len(response.choices[0].message.content) > 0 - - except litellm.InternalServerError: - pass - except RateLimitError: - pass - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - @pytest.mark.parametrize( "stop", [""], @@ -911,49 +862,6 @@ def test_completion_bedrock_external_client_region(monkeypatch): pytest.fail(f"Error occurred: {e}") -def test_bedrock_tool_calling(): - """ - # related issue: https://github.com/BerriAI/litellm/issues/5007 - # Bedrock tool names must satisfy regular expression pattern: [a-zA-Z][a-zA-Z0-9_]* ensure this is true - """ - litellm.set_verbose = True - response = litellm.completion( - model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0", - fallbacks=["bedrock/meta.llama3-1-8b-instruct-v1:0"], - messages=[ - { - "role": "user", - "content": "What's the weather like in Boston today in Fahrenheit?", - } - ], - tools=[ - { - "type": "function", - "function": { - "name": "-DoSomethingVeryCool-forLitellm_Testin999229291-0293993", - "description": "use this to get the current weather", - "parameters": {"type": "object", "properties": {}}, - }, - } - ], - ) - - print("bedrock response") - print(response) - - # Assert that the tools in response have the same function name as the input - _choice_1 = response.choices[0] - if _choice_1.message.tool_calls is not None: - print(_choice_1.message.tool_calls) - for tool_call in _choice_1.message.tool_calls: - _tool_Call_name = tool_call.function.name - if _tool_Call_name is not None and "DoSomethingVeryCool" in _tool_Call_name: - assert ( - _tool_Call_name - == "-DoSomethingVeryCool-forLitellm_Testin999229291-0293993" - ) - - def test_bedrock_tools_pt_valid_names(): """ # related issue: https://github.com/BerriAI/litellm/issues/5007 @@ -2031,6 +1939,14 @@ def test_bedrock_supports_tool_call(model, expected_supports_tool_call): class TestBedrockConverseChatCrossRegion(BaseLLMChatTest): + test_content_list_handling = None + test_developer_role_translation = None + test_function_calling_with_tool_response = None + test_image_url = None + test_json_response_format_stream = None + test_tool_call_with_empty_enum_property = None + test_tool_call_with_property_type_array = None + def get_base_completion_call_args(self) -> dict: os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") @@ -2070,6 +1986,9 @@ class TestBedrockConverseChatCrossRegion(BaseLLMChatTest): class TestBedrockConverseAnthropicUnitTests(BaseAnthropicChatTest): + test_completion_thinking_with_max_tokens = None + test_completion_thinking_without_max_tokens = None + def get_base_completion_call_args(self) -> dict: return { "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", @@ -2083,6 +2002,11 @@ class TestBedrockConverseAnthropicUnitTests(BaseAnthropicChatTest): class TestBedrockConverseChatNormal(BaseLLMChatTest): + test_content_list_handling = None + test_empty_tools = None + test_function_calling_with_tool_response = None + test_image_url = None + def get_base_completion_call_args(self) -> dict: os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") @@ -2098,6 +2022,10 @@ class TestBedrockConverseChatNormal(BaseLLMChatTest): class TestBedrockConverseNovaTestSuite(BaseLLMChatTest): + test_content_list_handling = None + test_function_calling_with_tool_response = None + test_image_url = None + def get_base_completion_call_args(self) -> dict: os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") diff --git a/tests/llm_translation/test_bedrock_gpt_oss.py b/tests/llm_translation/test_bedrock_gpt_oss.py index b264c16601f..777b374ee66 100644 --- a/tests/llm_translation/test_bedrock_gpt_oss.py +++ b/tests/llm_translation/test_bedrock_gpt_oss.py @@ -9,6 +9,8 @@ from litellm.llms.custom_httpx.http_handler import HTTPHandler class TestBedrockGPTOSS(BaseLLMChatTest): + test_json_response_format = None + def get_base_completion_call_args(self) -> dict: return { "model": "bedrock/converse/openai.gpt-oss-20b-1:0", diff --git a/tests/llm_translation/test_bedrock_invoke_tests.py b/tests/llm_translation/test_bedrock_invoke_tests.py index cf53899ecf6..46386b207cb 100644 --- a/tests/llm_translation/test_bedrock_invoke_tests.py +++ b/tests/llm_translation/test_bedrock_invoke_tests.py @@ -6,6 +6,16 @@ import litellm from litellm.types.llms.bedrock import BedrockInvokeNovaRequest +_LITELLM_LOGO_IMAGE_URL = ( + "https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0/" + "ui/litellm-dashboard/public/assets/logos/litellm_logo.jpg" +) +_AWSMP_LOGO_IMAGE_URL = ( + "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/" + "c233c9ade2ccb5491072ae232c814942.png" +) + + @pytest.mark.flaky(retries=3, delay=5) class TestBedrockInvokeClaudeJson(BaseLLMChatTest): def get_base_completion_call_args(self) -> dict: @@ -18,8 +28,27 @@ class TestBedrockInvokeClaudeJson(BaseLLMChatTest): """Test that tool calls with no arguments is translated correctly. Relevant issue: https://github.com/BerriAI/litellm/issues/6833""" pass + @pytest.mark.parametrize( + "image_url, detail", + [ + (_LITELLM_LOGO_IMAGE_URL, None), + (_LITELLM_LOGO_IMAGE_URL, "low"), + (_LITELLM_LOGO_IMAGE_URL, "high"), + (_AWSMP_LOGO_IMAGE_URL, "low"), + (_AWSMP_LOGO_IMAGE_URL, "high"), + ], + ) + @pytest.mark.flaky(retries=4, delay=2) + def test_image_url(self, image_url, detail): + super().test_image_url(detail=detail, image_url=image_url) + test_content_list_handling = None + test_image_url_string = None + test_pdf_handling = None + class TestBedrockInvokeNovaJson(BaseLLMChatTest): + test_json_response_format = None + def get_base_completion_call_args(self) -> dict: return { "model": "bedrock/invoke/us.amazon.nova-micro-v1:0", diff --git a/tests/llm_translation/test_bedrock_llama.py b/tests/llm_translation/test_bedrock_llama.py index 6c1a7073c13..b02b482b955 100644 --- a/tests/llm_translation/test_bedrock_llama.py +++ b/tests/llm_translation/test_bedrock_llama.py @@ -5,6 +5,10 @@ import litellm class TestBedrockTestSuite(BaseLLMChatTest): + test_content_list_handling = None + test_empty_tools = None + test_function_calling_with_tool_response = None + def test_tool_call_no_arguments(self, tool_call_no_arguments): pass diff --git a/tests/llm_translation/test_bedrock_moonshot.py b/tests/llm_translation/test_bedrock_moonshot.py index 3bf047c51a5..5323a87c366 100644 --- a/tests/llm_translation/test_bedrock_moonshot.py +++ b/tests/llm_translation/test_bedrock_moonshot.py @@ -30,6 +30,8 @@ class TestBedrockMoonshotInvoke(BaseLLMChatTest): Inherits all standard LLM tests from BaseLLMChatTest. """ + test_json_response_format_stream = None + def get_base_completion_call_args(self) -> dict: litellm._turn_on_debug() return { diff --git a/tests/llm_translation/test_bedrock_nova_json.py b/tests/llm_translation/test_bedrock_nova_json.py index 754ef4e3525..f9531c99b52 100644 --- a/tests/llm_translation/test_bedrock_nova_json.py +++ b/tests/llm_translation/test_bedrock_nova_json.py @@ -5,6 +5,14 @@ import litellm class TestBedrockNovaJson(BaseLLMChatTest): + test_content_list_handling = None + test_developer_role_translation = None + test_empty_tools = None + test_function_calling_with_tool_response = None + test_json_response_format_stream = None + test_tool_call_with_empty_enum_property = None + test_tool_call_with_property_type_array = None + def get_base_completion_call_args(self) -> dict: litellm._turn_on_debug() return { diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index 1a34e404d7f..7b0b741563d 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -74,6 +74,16 @@ GEMINI_3_IMAGE_SIZE_MAPPINGS = [ class TestGoogleAIStudioGemini(BaseLLMChatTest): + test_async_pdf_handling_with_file_id = None + test_content_list_handling = None + test_developer_role_translation = None + test_function_calling_with_tool_response = None + test_image_url = None + test_json_response_nested_json_schema = None + test_json_response_nested_pydantic_obj = None + test_json_response_pydantic_obj = None + test_web_search = None + def get_base_completion_call_args(self) -> dict: return {"model": "gemini/gemini-2.5-flash"} diff --git a/tests/llm_translation/test_groq.py b/tests/llm_translation/test_groq.py index fbecbeab08b..ce2d5461d60 100644 --- a/tests/llm_translation/test_groq.py +++ b/tests/llm_translation/test_groq.py @@ -18,6 +18,10 @@ from litellm.llms.groq.chat.transformation import ( class TestGroq(BaseLLMChatTest): + test_content_list_handling = None + test_empty_tools = None + test_web_search = None + def get_base_completion_call_args(self) -> dict: return { "model": "groq/openai/gpt-oss-120b", diff --git a/tests/llm_translation/test_openai.py b/tests/llm_translation/test_openai.py index f3b4ba0e8a6..d748a56e90c 100644 --- a/tests/llm_translation/test_openai.py +++ b/tests/llm_translation/test_openai.py @@ -274,6 +274,7 @@ async def test_vision_with_custom_model(): class TestOpenAIChatCompletion(BaseLLMChatTest): test_basic_tool_calling = None + test_function_calling_with_tool_response = None def get_base_completion_call_args(self) -> dict: return {"model": "gpt-4o-mini"} @@ -687,17 +688,6 @@ def test_openai_tool_calling(): response = litellm.completion(**completion_params) -@pytest.mark.asyncio -async def test_openai_gpt5_reasoning(): - response = await litellm.acompletion( - model="openai/gpt-5-mini", - messages=[{"role": "user", "content": "What is the capital of France?"}], - reasoning_effort="minimal", - ) - print("response: ", response) - assert response.choices[0].message.content is not None - - @pytest.mark.asyncio async def test_openai_safety_identifier_parameter(): """Test that safety_identifier parameter is correctly passed to the OpenAI API.""" diff --git a/tests/llm_translation/test_openai_o1.py b/tests/llm_translation/test_openai_o1.py index fd25e04d67d..e3c81e3920e 100644 --- a/tests/llm_translation/test_openai_o1.py +++ b/tests/llm_translation/test_openai_o1.py @@ -142,6 +142,10 @@ def test_litellm_responses(): class TestOpenAIO1(BaseOSeriesModelsTest, BaseLLMChatTest): + test_empty_tools = None + test_tool_call_with_empty_enum_property = None + test_tool_call_with_property_type_array = None + def get_base_completion_call_args(self): return { "model": "o1", @@ -162,6 +166,9 @@ class TestOpenAIO1(BaseOSeriesModelsTest, BaseLLMChatTest): class TestOpenAIO3(BaseOSeriesModelsTest, BaseLLMChatTest): + test_basic_tool_calling = None + test_function_calling_with_tool_response = None + def get_base_completion_call_args(self): return { "model": "o3-mini", @@ -188,27 +195,3 @@ def test_o3_reasoning_effort(): reasoning_effort="high", ) assert resp.choices[0].message.content is not None - - -@pytest.mark.parametrize("model", ["o1", "o3-mini"]) -def test_streaming_response(model): - """Test that streaming response is returned correctly""" - from litellm import completion - - response = completion( - model=model, - messages=[ - {"role": "system", "content": "Be a good bot!"}, - {"role": "user", "content": "Hello!"}, - ], - stream=True, - ) - - assert response is not None - - chunks = [] - for chunk in response: - chunks.append(chunk) - - resp = litellm.stream_chunk_builder(chunks=chunks) - print(resp) diff --git a/tests/llm_translation/test_together_ai.py b/tests/llm_translation/test_together_ai.py index 2c17bee9a7c..1cf4834ebf7 100644 --- a/tests/llm_translation/test_together_ai.py +++ b/tests/llm_translation/test_together_ai.py @@ -16,6 +16,14 @@ import pytest class TestTogetherAI(BaseLLMChatTest): test_basic_tool_calling = None + test_empty_tools = None + test_function_calling_with_tool_response = None + test_json_response_format = None + test_json_response_nested_json_schema = None + test_json_response_nested_pydantic_obj = None + test_json_response_pydantic_obj = None + test_tool_call_with_empty_enum_property = None + test_tool_call_with_property_type_array = None def get_base_completion_call_args(self) -> dict: litellm.set_verbose = True diff --git a/tests/llm_translation/test_xai.py b/tests/llm_translation/test_xai.py index d6d42ed215e..4f3346b5477 100644 --- a/tests/llm_translation/test_xai.py +++ b/tests/llm_translation/test_xai.py @@ -8,7 +8,6 @@ from unittest.mock import AsyncMock import httpx import pytest -import litellm from litellm import Choices, Message, ModelResponse, EmbeddingResponse, Usage from litellm import completion from unittest.mock import patch @@ -179,31 +178,7 @@ class TestXAIChat(BaseLLMChatTest): """Test that tool calls with no arguments is translated correctly. Relevant issue: https://github.com/BerriAI/litellm/issues/6833""" pass - def test_web_search(self): - """Web search is only supported for Grok 4 family models""" - from litellm.utils import supports_web_search - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - litellm._turn_on_debug() - - # Use grok-4-1-fast which supports web search - model = "xai/grok-4-1-fast" - - if not supports_web_search(model, None): - pytest.skip("Model does not support web search") - - response = completion( - model=model, - messages=[ - {"role": "user", "content": "What's the weather like in Boston today?"} - ], - web_search_options={}, - max_tokens=100, - ) - - assert response is not None + test_web_search = None def test_xai_streaming_with_include_usage(): diff --git a/tests/local_testing/test_acooldowns_router.py b/tests/local_testing/test_acooldowns_router.py index 18c58a5cfac..61e947b1322 100644 --- a/tests/local_testing/test_acooldowns_router.py +++ b/tests/local_testing/test_acooldowns_router.py @@ -4,8 +4,6 @@ import asyncio import os import time -import traceback - import pytest import concurrent @@ -19,113 +17,9 @@ from litellm import Router load_dotenv() -def _make_model_list(): - return [ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "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", - "litellm_params": { - "model": "gpt-3.5-turbo", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - "tpm": 1000000, - "rpm": 9000, - }, - ] - - -def _make_kwargs(): - return { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": "Hey, how's it going?"}], - } - - -@pytest.mark.flaky(retries=3, delay=1) -def test_multiple_deployments_sync(): - import concurrent - import time - - litellm.set_verbose = False - results = [] - kwargs = _make_kwargs() - router = Router( - model_list=_make_model_list(), - redis_host=os.getenv("REDIS_HOST"), - redis_password=os.getenv("REDIS_PASSWORD"), - redis_port=int(os.getenv("REDIS_PORT")), # type: ignore - routing_strategy="simple-shuffle", - set_verbose=True, - num_retries=1, - ) # type: ignore - try: - for _ in range(3): - response = router.completion(**kwargs) - results.append(response) - print(results) - router.reset() - except Exception as e: - print(f"FAILED TEST!") - pytest.fail(f"An error occurred - {traceback.format_exc()}") - - # test_multiple_deployments_sync() -def test_multiple_deployments_parallel(): - litellm.set_verbose = False # Corrected the syntax for setting verbose to False - results = [] - futures = {} - kwargs = _make_kwargs() - start_time = time.time() - router = Router( - model_list=_make_model_list(), - redis_host=os.getenv("REDIS_HOST"), - redis_password=os.getenv("REDIS_PASSWORD"), - redis_port=int(os.getenv("REDIS_PORT")), # type: ignore - routing_strategy="simple-shuffle", - set_verbose=True, - num_retries=1, - ) # type: ignore - # Assuming you have an executor instance defined somewhere in your code - with concurrent.futures.ThreadPoolExecutor() as executor: - for _ in range(5): - future = executor.submit(router.completion, **kwargs) - futures[future] = future - - # Retrieve the results from the futures - while futures: - done, not_done = concurrent.futures.wait( - futures.values(), - timeout=10, - return_when=concurrent.futures.FIRST_COMPLETED, - ) - for future in done: - try: - result = future.result() - results.append(result) - del futures[future] # Remove the done future - except Exception as e: - print(f"Exception: {e}; traceback: {traceback.format_exc()}") - del futures[future] # Remove the done future with exception - - print(f"Remaining futures: {len(futures)}") - router.reset() - end_time = time.time() - print(results) - print(f"ELAPSED TIME: {end_time - start_time}") - - # Assuming litellm, router, and executor are defined somewhere in your code diff --git a/tests/local_testing/test_amazing_vertex_completion.py b/tests/local_testing/test_amazing_vertex_completion.py index e34cee90f65..7f4044fc87e 100644 --- a/tests/local_testing/test_amazing_vertex_completion.py +++ b/tests/local_testing/test_amazing_vertex_completion.py @@ -137,30 +137,6 @@ def load_vertex_ai_credentials(): os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name) -@pytest.mark.asyncio -async def test_get_response(): - load_vertex_ai_credentials() - prompt = '\ndef count_nums(arr):\n """\n Write a function count_nums which takes an array of integers and returns\n the number of elements which has a sum of digits > 0.\n If a number is negative, then its first signed digit will be negative:\n e.g. -123 has signed digits -1, 2, and 3.\n >>> count_nums([]) == 0\n >>> count_nums([-1, 11, -11]) == 1\n >>> count_nums([1, 1, 2]) == 3\n """\n' - try: - response = await acompletion( - model="gemini-2.5-flash-lite", - messages=[ - { - "role": "system", - "content": "Complete the given code with no more explanation. Remember that there is a 4-space indent before the first line of your generated code.", - }, - {"role": "user", "content": prompt}, - ], - ) - return response - except litellm.RateLimitError: - pass - except litellm.UnprocessableEntityError as e: - pass - except Exception as e: - pytest.fail(f"An error occurred - {str(e)}") - - # test_vertex_ai_anthropic_streaming() @@ -341,35 +317,6 @@ def test_avertex_ai_stream(): # test_vertex_ai_stream() -@pytest.mark.flaky(retries=3, delay=1) -@pytest.mark.asyncio -async def test_async_vertexai_response_basic(): - load_vertex_ai_credentials() - try: - user_message = "Hello, how are you?" - messages = [{"content": user_message, "role": "user"}] - response = await acompletion( - model="gemini-3.5-flash", - messages=messages, - temperature=0.7, - timeout=5, - vertex_location="global", - ) - print(f"response: {response}") - except litellm.NotFoundError as e: - pass - except litellm.RateLimitError as e: - pass - except litellm.Timeout as e: - pass - except litellm.APIError as e: - pass - except litellm.InternalServerError as e: - pass - except Exception as e: - pytest.fail(f"An exception occurred: {e}") - - @pytest.mark.flaky(retries=3, delay=1) @pytest.mark.asyncio async def test_async_vertexai_streaming_response(): @@ -434,49 +381,6 @@ async def test_async_vertexai_streaming_response(): pytest.fail(f"An exception occurred: {e}") -@pytest.mark.parametrize("load_pdf", [False]) # True, -@pytest.mark.flaky(retries=3, delay=1) -def test_completion_function_plus_pdf(load_pdf): - litellm.set_verbose = True - load_vertex_ai_credentials() - try: - import base64 - - import requests - - # URL of the file - url = "https://storage.googleapis.com/cloud-samples-data/generative-ai/pdf/2403.05530.pdf" - - # Download the file - if load_pdf: - response = requests.get(url) - file_data = response.content - - encoded_file = base64.b64encode(file_data).decode("utf-8") - url = f"data:application/pdf;base64,{encoded_file}" - - image_content = [ - {"type": "text", "text": "What's this file about?"}, - { - "type": "image_url", - "image_url": {"url": url}, - }, - ] - image_message = {"role": "user", "content": image_content} - - response = completion( - model="vertex_ai_beta/gemini-2.5-flash-lite", - messages=[image_message], - stream=False, - ) - - print(response) - except litellm.InternalServerError as e: - pass - except Exception as e: - pytest.fail("Got={}".format(str(e))) - - def encode_image(image_path): import base64 @@ -1470,90 +1374,6 @@ async def test_gemini_pro_httpx_custom_api_base(model): # @pytest.mark.skip(reason="exhausted vertex quota. need to refactor to mock the call") -@pytest.mark.parametrize("sync_mode", [True]) -@pytest.mark.parametrize("provider", ["vertex_ai"]) -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=1) -async def test_gemini_pro_function_calling(provider, sync_mode): - try: - load_vertex_ai_credentials() - litellm.set_verbose = True - - messages = [ - { - "role": "system", - "content": "Your name is Litellm Bot, you are a helpful assistant", - }, - # User asks for their name and weather in San Francisco - { - "role": "user", - "content": "Hello, what is your name and can you tell me the weather?", - }, - # Assistant replies with a tool call - { - "role": "assistant", - "content": "", - "tool_calls": [ - { - "id": "call_123", - "type": "function", - "index": 0, - "function": { - "name": "get_weather", - "arguments": '{"location":"San Francisco, CA"}', - }, - } - ], - }, - # The result of the tool call is added to the history - { - "role": "tool", - "tool_call_id": "call_123", - "content": "27 degrees celsius and clear in San Francisco, CA", - }, - # Now the assistant can reply with the result of the tool call. - ] - - tools = [ - { - "type": "function", - "function": { - "name": "get_weather", - "description": "Get the current weather in a given location", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "The city and state, e.g. San Francisco, CA", - } - }, - "required": ["location"], - }, - }, - } - ] - - data = { - "model": "{}/gemini-2.5-flash-lite".format(provider), - "messages": messages, - "tools": tools, - } - if sync_mode: - response = litellm.completion(**data) - else: - response = await litellm.acompletion(**data) - - print(f"response: {response}") - except litellm.RateLimitError as e: - pass - except Exception as e: - if "429 Quota exceeded" in str(e): - pass - else: - pytest.fail("An unexpected exception occurred - {}".format(str(e))) - - # gemini_pro_function_calling() @@ -3522,46 +3342,6 @@ def test_vertex_ai_llama_tool_calling(): assert response._hidden_params["response_cost"] > 0 -def test_vertex_schema_test(): - load_vertex_ai_credentials() - litellm._turn_on_debug() - - def tool_call(text: str | None) -> str: - return text or "No text provided" - - tool = { - "type": "function", - "function": { - "name": "git_create_branch", - "description": "Creates a new branch from an optional base branch", - "parameters": { - "type": "object", - "properties": { - "repo_path": {"title": "Repo Path", "type": "string"}, - "branch_name": {"title": "Branch Name", "type": "string"}, - "base_branch": { - "anyOf": [{"type": "string"}, {"type": "null"}], - "default": None, - "title": "Base Branch", - }, - }, - "required": ["repo_path", "branch_name"], - "title": "GitCreateBranch", - }, - }, - } - - response = litellm.completion( - model="vertex_ai/gemini-3.5-flash", - messages=[{"role": "user", "content": "call the tool"}], - tools=[tool], - tool_choice="required", - vertex_location="global", - ) - - print(response) - - def test_gemini_nullable_object_tool_schema_httpx(): """ Ensure nullable object tool params preserve nested properties in Vertex schema conversion. diff --git a/tests/local_testing/test_arize_ai.py b/tests/local_testing/test_arize_ai.py index 138858cee03..d427e686dfa 100644 --- a/tests/local_testing/test_arize_ai.py +++ b/tests/local_testing/test_arize_ai.py @@ -35,26 +35,6 @@ async def test_async_otel_callback(): await asyncio.sleep(2) -@pytest.mark.asyncio() -async def test_async_dynamic_arize_config(): - litellm.set_verbose = True - - verbose_proxy_logger.setLevel(logging.DEBUG) - verbose_logger.setLevel(logging.DEBUG) - litellm.success_callback = ["arize"] - - await litellm.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi test from arize dynamic config"}], - temperature=0.1, - user="OTEL_USER", - arize_api_key=os.getenv("ARIZE_SPACE_API_KEY"), - arize_space_key=os.getenv("ARIZE_SPACE_KEY"), - ) - - await asyncio.sleep(2) - - @pytest.fixture def mock_env_vars(monkeypatch): monkeypatch.setenv("ARIZE_SPACE_KEY", "test_space_key") diff --git a/tests/local_testing/test_async_fn.py b/tests/local_testing/test_async_fn.py index e2b3a62bd28..a7b105bfc68 100644 --- a/tests/local_testing/test_async_fn.py +++ b/tests/local_testing/test_async_fn.py @@ -215,43 +215,6 @@ async def test_hf_completion_tgi(): # test_get_cloudflare_response_streaming() -def test_get_response_streaming(): - import asyncio - - async def test_async_call(): - user_message = "write a short poem in one sentence" - messages = [{"content": user_message, "role": "user"}] - try: - litellm.set_verbose = True - response = await acompletion( - model="gpt-3.5-turbo", messages=messages, stream=True, timeout=5 - ) - print(type(response)) - - import inspect - - is_async_generator = inspect.isasyncgen(response) - print(is_async_generator) - - output = "" - i = 0 - async for chunk in response: - token = chunk["choices"][0]["delta"].get("content", "") - if token == None: - continue # openai v1.0.0 returns content=None - output += token - assert output is not None, "output cannot be None." - assert isinstance(output, str), "output needs to be of type str" - assert len(output) > 0, "Length of output needs to be greater than 0." - print(f"output: {output}") - except litellm.Timeout as e: - pass - except Exception as e: - pytest.fail(f"An exception occurred: {e}") - - asyncio.run(test_async_call()) - - # test_get_response_streaming() diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py index 2d8983c2fc8..5ff7d79e3f8 100644 --- a/tests/local_testing/test_completion.py +++ b/tests/local_testing/test_completion.py @@ -192,242 +192,6 @@ def test_completion_empower(): pytest.fail(f"Error occurred: {e}") -def test_completion_claude_3_empty_response(): - litellm.set_verbose = True - - messages = [ - { - "role": "system", - "content": [{"type": "text", "text": "You are 2twNLGfqk4GMOn3ffp4p."}], - }, - {"role": "user", "content": "Hi gm!", "name": "ishaan"}, - {"role": "assistant", "content": "Good morning! How are you doing today?"}, - { - "role": "user", - "content": "I was hoping we could chat a bit", - }, - ] - try: - response = litellm.completion( - model="claude-sonnet-4-5-20250929", messages=messages - ) - print(response) - except litellm.InternalServerError as e: - pytest.skip(f"InternalServerError - {str(e)}") - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - -def test_completion_claude_3(): - litellm.set_verbose = True - messages = [ - { - "role": "user", - "content": "\nWhat is the query for `console.log` => `console.error`\n", - }, - { - "role": "assistant", - "content": "\nThis is the GritQL query for the given before/after examples:\n\n`console.log` => `console.error`\n\n", - }, - { - "role": "user", - "content": "\nWhat is the query for `console.info` => `consdole.heaven`\n", - }, - ] - try: - # test without max tokens - response = completion( - model="anthropic/claude-sonnet-4-5-20250929", - messages=messages, - ) - # Add any assertions, here to check response args - print(response) - except litellm.InternalServerError as e: - pytest.skip(f"InternalServerError - {str(e)}") - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - -@pytest.mark.parametrize( - "model", - ["anthropic/claude-sonnet-4-5-20250929", "us.anthropic.claude-sonnet-4-5-20250929-v1:0"], -) -def test_completion_claude_3_function_call(model): - litellm.set_verbose = True - tools = [ - { - "type": "function", - "function": { - "name": "get_current_weather", - "description": "Get the current weather in a given location", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "The city and state, e.g. San Francisco, CA", - }, - "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]}, - }, - "required": ["location"], - }, - }, - } - ] - messages = [ - { - "role": "user", - "content": "What's the weather like in Boston today in Fahrenheit?", - } - ] - try: - # test without max tokens - response = completion( - model=model, - messages=messages, - tools=tools, - tool_choice={ - "type": "function", - "function": {"name": "get_current_weather"}, - }, - drop_params=True, - ) - - # Add any assertions here to check response args - print(response) - assert isinstance(response.choices[0].message.tool_calls[0].function.name, str) - assert isinstance( - response.choices[0].message.tool_calls[0].function.arguments, str - ) - - messages.append( - response.choices[0].message.model_dump() - ) # Add assistant tool invokes - tool_result = ( - '{"location": "Boston", "temperature": "72", "unit": "fahrenheit"}' - ) - # Add user submitted tool results in the OpenAI format - messages.append( - { - "tool_call_id": response.choices[0].message.tool_calls[0].id, - "role": "tool", - "name": response.choices[0].message.tool_calls[0].function.name, - "content": tool_result, - } - ) - # In the second response, Claude should deduce answer from tool results - second_response = completion( - model=model, - messages=messages, - tools=tools, - tool_choice="auto", - drop_params=True, - ) - print(second_response) - except litellm.InternalServerError: - pass - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - -@pytest.mark.parametrize("sync_mode", [True]) -@pytest.mark.parametrize( - "model, api_key, api_base", - [ - ("gpt-3.5-turbo", None, None), - ("claude-sonnet-4-5-20250929", None, None), - ("us.anthropic.claude-sonnet-4-5-20250929-v1:0", None, None), - # ( - # "azure_ai/command-r-plus", - # os.getenv("AZURE_COHERE_API_KEY"), - # os.getenv("AZURE_COHERE_API_BASE"), - # ), - ], -) -@pytest.mark.asyncio -async def test_model_function_invoke(model, sync_mode, api_key, api_base): - try: - litellm.set_verbose = True - - messages = [ - { - "role": "system", - "content": "Your name is Litellm Bot, you are a helpful assistant", - }, - # User asks for their name and weather in San Francisco - { - "role": "user", - "content": "Hello, what is your name and can you tell me the weather?", - }, - # Assistant replies with a tool call - { - "role": "assistant", - "content": "", - "tool_calls": [ - { - "id": "call_123", - "type": "function", - "index": 0, - "function": { - "name": "get_weather", - "arguments": '{"location": "San Francisco, CA"}', - }, - } - ], - }, - # The result of the tool call is added to the history - { - "role": "tool", - "tool_call_id": "call_123", - "content": "27 degrees celsius and clear in San Francisco, CA", - }, - # Now the assistant can reply with the result of the tool call. - ] - - tools = [ - { - "type": "function", - "function": { - "name": "get_weather", - "description": "Get the current weather in a given location", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "The city and state, e.g. San Francisco, CA", - } - }, - "required": ["location"], - }, - }, - } - ] - - data = { - "model": model, - "messages": messages, - "tools": tools, - "api_key": api_key, - "api_base": api_base, - } - if sync_mode: - response = litellm.completion(**data) - else: - response = await litellm.acompletion(**data) - - print(f"response: {response}") - except litellm.InternalServerError: - pass - except litellm.RateLimitError as e: - pass - except Exception as e: - if "429 Quota exceeded" in str(e): - pass - else: - pytest.fail("An unexpected exception occurred - {}".format(str(e))) - - @pytest.mark.asyncio async def test_anthropic_no_content_error(): """ @@ -540,48 +304,6 @@ def test_parse_xml_params(): assert response["unit"] == "fahrenheit" -def test_completion_claude_3_multi_turn_conversations(): - litellm.set_verbose = True - litellm.modify_params = True - messages = [ - {"role": "assistant", "content": "?"}, # test first user message auto injection - {"role": "user", "content": "Hi!"}, - { - "role": "user", - "content": [{"type": "text", "text": "What is the weather like today?"}], - }, - {"role": "assistant", "content": "Hi! I am Claude. "}, - {"role": "assistant", "content": "Today is a sunny "}, - ] - try: - response = completion( - model="anthropic/claude-sonnet-4-5-20250929", - messages=messages, - ) - print(response) - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - -def test_completion_claude_3_stream(): - litellm.set_verbose = False - messages = [{"role": "user", "content": "Hello, world"}] - try: - # test without max tokens - response = completion( - model="anthropic/claude-sonnet-4-5-20250929", - messages=messages, - max_tokens=10, - stream=True, - ) - # Add any assertions, here to check response args - print(response) - for chunk in response: - print(chunk) - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - def encode_image(image_path): import base64 @@ -2253,25 +1975,6 @@ async def test_re_use_azure_async_client(): pytest.fail("got Exception", e) -def test_re_use_openaiClient(): - try: - print("gpt-3.5 with client test\n\n") - litellm.set_verbose = True - import openai - - client = openai.OpenAI( - api_key=os.environ["OPENAI_API_KEY"], - ) - ## Test OpenAI call - for _ in range(2): - response = litellm.completion( - model="gpt-3.5-turbo", messages=messages, client=client - ) - print(f"response: {response}") - except Exception as e: - pytest.fail("got Exception", e) - - @pytest.mark.skip( reason="this is bad test. It doesn't actually fail if the token is not set in the header. " ) @@ -3347,60 +3050,7 @@ def test_completion_gemini(model): # test_completion_gemini() -@pytest.mark.asyncio -async def test_acompletion_gemini(): - litellm.set_verbose = True - model_name = "gemini/gemini-2.5-flash-lite" - messages = [{"role": "user", "content": "Hey, how's it going?"}] - try: - response = await litellm.acompletion(model=model_name, messages=messages) - # Add any assertions here to check the response - print(f"response: {response}") - except litellm.Timeout as e: - pass - except litellm.APIError as e: - pass - except Exception as e: - if "InternalServerError" in str(e): - pass - else: - pytest.fail(f"Error occurred: {e}") - - # Deepseek tests -def test_completion_deepseek(): - litellm.set_verbose = True - model_name = "deepseek/deepseek-chat" - tools = [ - { - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather of an location, the user shoud supply a location first", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "The city and state, e.g. San Francisco, CA", - } - }, - "required": ["location"], - }, - }, - }, - ] - messages = [{"role": "user", "content": "How's the weather in Hangzhou?"}] - try: - response = completion(model=model_name, messages=messages, tools=tools) - # Add any assertions here to check the response - print(response) - except litellm.APIError as e: - pass - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - @pytest.mark.skip(reason="Account deleted by IBM.") def test_completion_watsonx_error(): litellm.set_verbose = True @@ -4107,37 +3757,3 @@ def test_completion_gpt_4o_empty_str(): messages=[{"role": "user", "content": ""}], ) assert resp.choices[0].message.content is not None - - -def test_edit_note(): - litellm.callbacks = ["langfuse_otel"] - response = completion( - model="gpt-4o", - messages=[ - { - "role": "system", - "content": "Your only job is to call the edit_note tool with the content specified in the user's message.", - }, - { - "role": "user", - "content": "Edit the note with the content: 'This is a test note.'", - }, - ], - tools=[ - { - "type": "function", - "function": { - "name": "edit_note", - "description": "Edit the note with the content specified in the user's message.", - "parameters": { - "type": "object", - "properties": { - "content": {"type": "string"}, - }, - }, - }, - }, - ], - ) - - return response diff --git a/tests/local_testing/test_function_call_parsing.py b/tests/local_testing/test_function_call_parsing.py index ebb13e0018d..6c1d1c7c5af 100644 --- a/tests/local_testing/test_function_call_parsing.py +++ b/tests/local_testing/test_function_call_parsing.py @@ -136,7 +136,7 @@ def trade(model_name: str) -> List[Trade]: # type: ignore @pytest.mark.parametrize( - "model", ["claude-haiku-4-5-20251001", "us.anthropic.claude-haiku-4-5-20251001-v1:0"] + "model", ["us.anthropic.claude-haiku-4-5-20251001-v1:0"] ) @pytest.mark.flaky(retries=6, delay=10) def test_function_call_parsing(model): diff --git a/tests/local_testing/test_function_calling.py b/tests/local_testing/test_function_calling.py index 164cdda1e50..2914f29182c 100644 --- a/tests/local_testing/test_function_calling.py +++ b/tests/local_testing/test_function_calling.py @@ -8,7 +8,7 @@ import io import pytest from unittest.mock import patch, MagicMock, AsyncMock import litellm -from litellm import RateLimitError, Timeout, completion, completion_cost, embedding +from litellm import RateLimitError, Timeout, completion_cost, embedding litellm.num_retries = 0 litellm.cache = None @@ -324,153 +324,6 @@ def test_groq_parallel_function_call(): pytest.fail(f"Error occurred: {e}") -@pytest.mark.parametrize( - "model", - [ - "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - ], -) -def test_passing_tool_result_as_list(model): - litellm.set_verbose = True - litellm._turn_on_debug() - messages = [ - { - "content": [ - { - "type": "text", - "text": "You are a helpful assistant that have the ability to interact with a computer to solve tasks.", - } - ], - "role": "system", - }, - { - "content": [ - { - "type": "text", - "text": "Write a git commit message for the current staging area and commit the changes.", - } - ], - "role": "user", - }, - { - "content": [ - { - "type": "text", - "text": "I'll help you commit the changes. Let me first check the git status to see what changes are staged.", - } - ], - "role": "assistant", - "tool_calls": [ - { - "index": 1, - "function": { - "arguments": '{"command": "git status", "thought": "Checking git status to see staged changes"}', - "name": "execute_bash", - }, - "id": "toolu_01V1paXrun4CVetdAGiQaZG5", - "type": "function", - } - ], - }, - { - "content": [ - { - "type": "text", - "text": 'OBSERVATION:\nOn branch master\r\n\r\nNo commits yet\r\n\r\nChanges to be committed:\r\n (use "git rm --cached ..." to unstage)\r\n\tnew file: hello.py\r\n\r\n\r\n[Python Interpreter: /openhands/poetry/openhands-ai-5O4_aCHf-py3.12/bin/python]\nroot@openhands-workspace:/workspace # \n[Command finished with exit code 0]', - } - ], - "role": "tool", - "tool_call_id": "toolu_01V1paXrun4CVetdAGiQaZG5", - "name": "execute_bash", - }, - ] - tools = [ - { - "type": "function", - "function": { - "name": "execute_bash", - "description": 'Execute a bash command in the terminal.\n* Long running commands: For commands that may run indefinitely, it should be run in the background and the output should be redirected to a file, e.g. command = `python3 app.py > server.log 2>&1 &`.\n* Interactive: If a bash command returns exit code `-1`, this means the process is not yet finished. The assistant must then send a second call to terminal with an empty `command` (which will retrieve any additional logs), or it can send additional text (set `command` to the text) to STDIN of the running process, or it can send command=`ctrl+c` to interrupt the process.\n* Timeout: If a command execution result says "Command timed out. Sending SIGINT to the process", the assistant should retry running the command in the background.\n', - "parameters": { - "type": "object", - "properties": { - "thought": { - "type": "string", - "description": "Reasoning about the action to take.", - }, - "command": { - "type": "string", - "description": "The bash command to execute. Can be empty to view additional logs when previous exit code is `-1`. Can be `ctrl+c` to interrupt the currently running process.", - }, - }, - "required": ["command"], - }, - }, - }, - { - "type": "function", - "function": { - "name": "finish", - "description": "Finish the interaction.\n* Do this if the task is complete.\n* Do this if the assistant cannot proceed further with the task.\n", - }, - }, - { - "type": "function", - "function": { - "name": "str_replace_editor", - "description": "Custom editing tool for viewing, creating and editing files\n* State is persistent across command calls and discussions with the user\n* If `path` is a file, `view` displays the result of applying `cat -n`. If `path` is a directory, `view` lists non-hidden files and directories up to 2 levels deep\n* The `create` command cannot be used if the specified `path` already exists as a file\n* If a `command` generates a long output, it will be truncated and marked with ``\n* The `undo_edit` command will revert the last edit made to the file at `path`\n\nNotes for using the `str_replace` command:\n* The `old_str` parameter should match EXACTLY one or more consecutive lines from the original file. Be mindful of whitespaces!\n* If the `old_str` parameter is not unique in the file, the replacement will not be performed. Make sure to include enough context in `old_str` to make it unique\n* The `new_str` parameter should contain the edited lines that should replace the `old_str`\n", - "parameters": { - "type": "object", - "properties": { - "command": { - "description": "The commands to run. Allowed options are: `view`, `create`, `str_replace`, `insert`, `undo_edit`.", - "enum": [ - "view", - "create", - "str_replace", - "insert", - "undo_edit", - ], - "type": "string", - }, - "path": { - "description": "Absolute path to file or directory, e.g. `/repo/file.py` or `/repo`.", - "type": "string", - }, - "file_text": { - "description": "Required parameter of `create` command, with the content of the file to be created.", - "type": "string", - }, - "old_str": { - "description": "Required parameter of `str_replace` command containing the string in `path` to replace.", - "type": "string", - }, - "new_str": { - "description": "Optional parameter of `str_replace` command containing the new string (if not given, no string will be added). Required parameter of `insert` command containing the string to insert.", - "type": "string", - }, - "insert_line": { - "description": "Required parameter of `insert` command. The `new_str` will be inserted AFTER the line `insert_line` of `path`.", - "type": "integer", - }, - "view_range": { - "description": "Optional parameter of `view` command when `path` points to a file. If none is given, the full file is shown. If provided, the file will be shown in the indicated line number range, e.g. [11, 12] will show lines 11 and 12. Indexing at 1 to start. Setting `[start_line, -1]` shows all lines from `start_line` to the end of the file.", - "items": {"type": "integer"}, - "type": "array", - }, - }, - "required": ["command", "path"], - }, - }, - }, - ] - for _ in range(2): - resp = completion(model=model, messages=messages, tools=tools) - print(resp) - - if model == "claude-sonnet-4-5-20250929": - assert resp.usage.prompt_tokens_details.cached_tokens > 0 - - @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio @pytest.mark.flaky(retries=6, delay=1) diff --git a/tests/local_testing/test_lowest_cost_routing.py b/tests/local_testing/test_lowest_cost_routing.py index 631271ca710..a0214ed10f7 100644 --- a/tests/local_testing/test_lowest_cost_routing.py +++ b/tests/local_testing/test_lowest_cost_routing.py @@ -10,7 +10,6 @@ load_dotenv() import copy import pytest -from litellm import Router from litellm.router_strategy.lowest_cost import LowestCostLoggingHandler from litellm.caching.caching import DualCache @@ -96,37 +95,6 @@ async def test_get_available_deployments_custom_price(): assert selected_model["model_info"]["id"] == "chatgpt-v-1" -@pytest.mark.asyncio -async def test_lowest_cost_routing(): - """ - Test if router, returns model with the lowest cost - """ - model_list = [ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4"}, - "model_info": {"id": "openai-gpt-4"}, - }, - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": {"id": "gpt-3.5-turbo"}, - }, - ] - - # init router - router = Router(model_list=model_list, routing_strategy="cost-based-routing") - response = await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "Hey, how's it going?"}], - ) - print(response) - print( - response._hidden_params["model_id"] - ) # expect groq-llama, since groq/llama has lowest cost - assert "gpt-3.5-turbo" == response._hidden_params["model_id"] - - async def _deploy(lowest_cost_logger, deployment_id, tokens_used, duration): kwargs = { "litellm_params": { diff --git a/tests/local_testing/test_router.py b/tests/local_testing/test_router.py index 4c62c28530d..4965fa631a9 100644 --- a/tests/local_testing/test_router.py +++ b/tests/local_testing/test_router.py @@ -64,71 +64,6 @@ def test_router_multi_org_list(): assert len(router.get_model_list()) == 3 -@pytest.mark.asyncio() -async def test_router_provider_wildcard_routing(): - """ - Pass list of orgs in 1 model definition, - expect a unique deployment for each to be created - """ - litellm.set_verbose = True - router = litellm.Router( - model_list=[ - { - "model_name": "openai/*", - "litellm_params": { - "model": "openai/*", - "api_key": os.environ["OPENAI_API_KEY"], - "api_base": "https://api.openai.com/v1", - }, - }, - { - "model_name": "anthropic/*", - "litellm_params": { - "model": "anthropic/*", - "api_key": os.environ["ANTHROPIC_API_KEY"], - }, - }, - { - "model_name": "groq/*", - "litellm_params": { - "model": "groq/*", - "api_key": os.environ["GROQ_API_KEY"], - }, - }, - ] - ) - - print("router model list = ", router.get_model_list()) - - response1 = await router.acompletion( - model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}", - messages=[{"role": "user", "content": "hello"}], - ) - - print("response 1 = ", response1) - - response2 = await router.acompletion( - model="openai/gpt-3.5-turbo", - messages=[{"role": "user", "content": "hello"}], - ) - - print("response 2 = ", response2) - - response3 = await router.acompletion( - model="groq/openai/gpt-oss-120b", - messages=[{"role": "user", "content": "hello"}], - ) - - print("response 3 = ", response3) - - response4 = await router.acompletion( - model=os.environ.get( - "CI_CD_DEFAULT_ANTHROPIC_MODEL", "claude-haiku-4-5-20251001" - ), - messages=[{"role": "user", "content": "hello"}], - ) - - @pytest.mark.asyncio() async def test_router_provider_wildcard_routing_regex(): """ @@ -986,176 +921,16 @@ def test_function_calling_on_router(): ### IMAGE GENERATION -@pytest.mark.asyncio -async def test_aimg_gen_on_router(): - litellm.set_verbose = True - try: - model_list = [ - { - "model_name": "gpt-image-1", - "litellm_params": { - "model": "gpt-image-1", - }, - } - ] - router = Router(model_list=model_list, num_retries=3) - response = await router.aimage_generation( - model="gpt-image-1", prompt="A cute baby sea otter" - ) - print(response) - assert len(response.data) > 0 - router.reset() - except litellm.InternalServerError as e: - pass - except Exception as e: - if "Your task failed as a result of our safety system." in str(e): - pass - elif "Operation polling timed out" in str(e): - pass - elif "Connection error" in str(e): - pass - else: - traceback.print_exc() - pytest.fail(f"Error occurred: {e}") - - # asyncio.run(test_aimg_gen_on_router()) -def test_img_gen_on_router(): - litellm.set_verbose = True - try: - model_list = [ - { - "model_name": "gpt-image-1", - "litellm_params": { - "model": "gpt-image-1", - }, - } - ] - router = Router(model_list=model_list) - response = router.image_generation( - model="gpt-image-1", prompt="A cute baby sea otter" - ) - print(response) - assert len(response.data) > 0 - router.reset() - except litellm.RateLimitError as e: - pass - except Exception as e: - traceback.print_exc() - pytest.fail(f"Error occurred: {e}") - - # test_img_gen_on_router() ### -def test_aembedding_on_router(): - litellm.set_verbose = True - try: - model_list = [ - { - "model_name": "text-embedding-ada-002", - "litellm_params": { - "model": "text-embedding-ada-002", - }, - "tpm": 100000, - "rpm": 10000, - }, - ] - router = Router(model_list=model_list) - - async def embedding_call(): - ## Test 1: user facing function - response = await router.aembedding( - model="text-embedding-ada-002", - input=["good morning from litellm", "this is another item"], - ) - print(response) - - ## Test 2: underlying function - response = await router._aembedding( - model="text-embedding-ada-002", - input=["good morning from litellm 2"], - ) - print(response) - router.reset() - - asyncio.run(embedding_call()) - - print("\n Making sync Embedding call\n") - ## Test 1: user facing function - response = router.embedding( - model="text-embedding-ada-002", - input=["good morning from litellm 2"], - ) - print(response) - router.reset() - - ## Test 2: underlying function - response = router._embedding( - model="text-embedding-ada-002", - input=["good morning from litellm 2"], - ) - print(response) - router.reset() - except Exception as e: - if "Your task failed as a result of our safety system." in str(e): - pass - elif "Operation polling timed out" in str(e): - pass - elif "Connection error" in str(e): - pass - else: - traceback.print_exc() - pytest.fail(f"Error occurred: {e}") - - # test_aembedding_on_router() -def test_azure_embedding_on_router(): - """ - [PROD Use Case] - Makes an aembedding call + embedding call - """ - litellm.set_verbose = True - try: - model_list = [ - { - "model_name": "text-embedding-ada-002", - "litellm_params": { - "model": "azure/text-embedding-ada-002", - "api_key": os.environ["AZURE_AI_API_KEY"], - "api_base": os.environ["AZURE_AI_API_BASE"], - }, - "tpm": 100000, - "rpm": 10000, - }, - ] - router = Router(model_list=model_list) - - async def embedding_call(): - response = await router.aembedding( - model="text-embedding-ada-002", input=["good morning from litellm"] - ) - print(response) - - asyncio.run(embedding_call()) - - print("\n Making sync Azure Embedding call\n") - - response = router.embedding( - model="text-embedding-ada-002", - input=["test 2 from litellm. async embedding"], - ) - print(response) - router.reset() - except Exception as e: - traceback.print_exc() - pytest.fail(f"Error occurred: {e}") - - # test_azure_embedding_on_router() @@ -1163,30 +938,6 @@ def test_azure_embedding_on_router(): # test openai-compatible endpoint -@pytest.mark.asyncio -async def test_mistral_on_router(): - litellm._turn_on_debug() - model_list = [ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "mistral/mistral-small-latest", - }, - }, - ] - router = Router(model_list=model_list) - response = await router.acompletion( - model="gpt-3.5-turbo", - messages=[ - { - "role": "user", - "content": "hello from litellm test", - } - ], - ) - print(response) - - # asyncio.run(test_mistral_on_router()) diff --git a/tests/local_testing/test_streaming.py b/tests/local_testing/test_streaming.py index c59ed667242..6e102b89554 100644 --- a/tests/local_testing/test_streaming.py +++ b/tests/local_testing/test_streaming.py @@ -435,28 +435,6 @@ def test_completion_azure_stream(): pytest.fail(f"Error occurred: {e}") -def test_completion_azure_function_calling_stream(): - try: - litellm.set_verbose = False - user_message = "What is the current weather in Boston?" - messages = [{"content": user_message, "role": "user"}] - response = completion( - model="azure/gpt-4.1-mini", - messages=messages, - stream=True, - tools=tools_schema, - ) - # Add any assertions here to check the response - for chunk in response: - print(chunk) - if chunk["choices"][0]["finish_reason"] == "stop": - break - print(chunk["choices"][0]["finish_reason"]) - print(chunk["choices"][0]["delta"]["content"]) - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - @pytest.mark.skip("Flaky ollama test - needs to be fixed") def test_completion_ollama_hosted_stream(): try: diff --git a/tests/local_testing/test_timeout.py b/tests/local_testing/test_timeout.py index 784e2c73cd7..c0187014c71 100644 --- a/tests/local_testing/test_timeout.py +++ b/tests/local_testing/test_timeout.py @@ -15,35 +15,6 @@ import litellm from tests.fake_openai_endpoint import FAKE_OPENAI_API_BASE -@pytest.mark.parametrize( - "model, provider", - [ - ("gpt-3.5-turbo", "openai"), - ("azure/gpt-4.1-mini", "azure"), - ], -) -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_httpx_timeout(model, provider, sync_mode): - """ - Test if setting httpx.timeout works for completion calls - """ - timeout_val = httpx.Timeout(10.0, connect=60.0) - - messages = [{"role": "user", "content": "Hey, how's it going?"}] - - if sync_mode: - response = litellm.completion( - model=model, messages=messages, timeout=timeout_val - ) - else: - response = await litellm.acompletion( - model=model, messages=messages, timeout=timeout_val - ) - - print(f"response: {response}") - - def test_timeout(): # this Will Raise a timeout litellm.set_verbose = False diff --git a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py b/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py index 48e81836035..566af351a98 100644 --- a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py +++ b/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py @@ -77,20 +77,6 @@ def validate_stream_chunk(chunk): assert isinstance(chunk.created, int) -def test_streaming_response(): - client = get_test_client() - stream = client.responses.create( - model="gpt-5.5", input="just respond with the word 'ping'", stream=True - ) - - collected_chunks = [] - for chunk in stream: - print("stream chunk=", chunk) - collected_chunks.append(chunk) - - assert len(collected_chunks) > 0 - - def test_model_not_found_error(): client = get_test_client() with pytest.raises(NotFoundError): diff --git a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py index d354ddafd00..ffbbf261e89 100644 --- a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py +++ b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py @@ -1,7 +1,7 @@ import json import os from datetime import datetime -from typing import AsyncIterator, Dict, Any +from typing import Dict, Any import asyncio import unittest.mock from unittest.mock import AsyncMock, MagicMock @@ -69,6 +69,9 @@ def _validate_anthropic_response(response: Dict[str, Any]): class TestAnthropicDirectAPI(BaseAnthropicMessagesTest): """Tests for direct Anthropic API calls""" + test_non_streaming_base = None + test_streaming_base = None + @property def model_config(self) -> Dict[str, Any]: return { @@ -87,6 +90,8 @@ class TestAnthropicDirectAPI(BaseAnthropicMessagesTest): class TestAnthropicBedrockAPI(BaseAnthropicMessagesTest): """Tests for Anthropic via Bedrock""" + test_streaming_base = None + @property def model_config(self) -> Dict[str, Any]: return { @@ -104,6 +109,8 @@ class TestAnthropicBedrockAPI(BaseAnthropicMessagesTest): class TestAnthropicOpenAIAPI(BaseAnthropicMessagesTest): """Tests for OpenAI via Anthropic messages interface""" + test_streaming_base = None + @property def model_config(self) -> Dict[str, Any]: return { @@ -126,67 +133,6 @@ class TestAnthropicOpenAIAPI(BaseAnthropicMessagesTest): pass -@pytest.mark.asyncio -async def test_anthropic_messages_streaming_with_bad_request(): - """ - Test the anthropic_messages with streaming request - """ - error = None - try: - response = await litellm.anthropic.messages.acreate( - messages=[{"role": "user", "content": "hi"}], - api_key=os.getenv("ANTHROPIC_API_KEY"), - model="claude-haiku-4-5-20251001", - max_tokens=100, - stream=True, - ) - print(response) - if isinstance(response, AsyncIterator): - async for chunk in response: - print("chunk=", chunk) - except Exception as e: - error = e - - if error is not None: - assert getattr(error, "status_code", 400) == 400, f"got {vars(error)}" - - -@pytest.mark.asyncio -async def test_anthropic_messages_router_streaming_with_bad_request(): - """ - Test the anthropic_messages with streaming request - """ - error = None - try: - router = Router( - model_list=[ - { - "model_name": "claude-special-alias", - "litellm_params": { - "model": "claude-haiku-4-5-20251001", - "api_key": os.getenv("ANTHROPIC_API_KEY"), - }, - } - ] - ) - - response = await router.aanthropic_messages( - messages=[{"role": "user", "content": "hi"}], - model="claude-special-alias", - max_tokens=100, - stream=True, - ) - print(response) - if isinstance(response, AsyncIterator): - async for chunk in response: - print("chunk=", chunk) - except Exception as e: - error = e - - if error is not None: - assert getattr(error, "status_code", 400) == 400, f"got {vars(error)}" - - @pytest.mark.asyncio async def test_anthropic_messages_litellm_router_non_streaming(): """ diff --git a/tests/test_openai_endpoints.py b/tests/test_openai_endpoints.py index 9fc0ea3c378..16f8de65236 100644 --- a/tests/test_openai_endpoints.py +++ b/tests/test_openai_endpoints.py @@ -425,39 +425,6 @@ async def test_completion_streaming_usage_metrics(): assert last_chunk.usage.total_tokens > 0, "Total tokens should be greater than 0" -@pytest.mark.asyncio -async def test_chat_completion_anthropic_structured_output(): - """ - Ensure nested pydantic output is returned correctly - """ - from pydantic import BaseModel - - class CalendarEvent(BaseModel): - name: str - date: str - participants: list[str] - - class EventsList(BaseModel): - events: list[CalendarEvent] - - messages = [ - {"role": "user", "content": "List 5 important events in the XIX century"} - ] - - client = AsyncOpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000") - - res = await client.beta.chat.completions.parse( - model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - messages=messages, - response_format=EventsList, - timeout=60, - ) - message = res.choices[0].message - - if message.parsed: - print(message.parsed.events) - - @pytest.mark.asyncio async def test_proxy_all_models(): """ diff --git a/tests/unified_google_tests/base_google_test.py b/tests/unified_google_tests/base_google_test.py index b7134962a0c..d6de60f6ec2 100644 --- a/tests/unified_google_tests/base_google_test.py +++ b/tests/unified_google_tests/base_google_test.py @@ -10,7 +10,6 @@ import litellm from litellm.google_genai import ( generate_content, agenerate_content, - generate_content_stream, agenerate_content_stream, ) from google.genai.types import ContentDict, PartDict @@ -195,45 +194,6 @@ class BaseGoogleGenAITest: return response - @pytest.mark.parametrize("is_async", [False, True]) - @pytest.mark.asyncio - async def test_streaming_base(self, is_async: bool): - """Base test for streaming requests (parametrized for sync/async)""" - request_params = self.model_config - temp_file_path = load_vertex_ai_credentials(model=request_params["model"]) - if temp_file_path: - self._temp_files_to_cleanup.append(temp_file_path) - contents = ContentDict( - parts=[PartDict(text="Hello, can you tell me a short joke?")], - role="user", - ) - - print( - f"Testing {'async' if is_async else 'sync'} streaming with model config: {request_params}" - ) - print(f"Contents: {contents}") - - chunks = [] - - if is_async: - print("\n--- Testing async agenerate_content_stream ---") - response = await agenerate_content_stream( - contents=contents, **request_params - ) - async for chunk in response: - print(f"Async chunk: {chunk}") - chunks.append(chunk) - else: - print("\n--- Testing sync generate_content_stream ---") - response = generate_content_stream(contents=contents, **request_params) - for chunk in response: - print(f"Sync chunk: {chunk}") - chunks.append(chunk) - - self._validate_streaming_response(chunks) - - return chunks - @pytest.mark.asyncio async def test_async_non_streaming_with_logging(self): """Test async non-streaming Google GenAI generate content with logging""" diff --git a/tests/unified_google_tests/test_google_ai_studio.py b/tests/unified_google_tests/test_google_ai_studio.py index 2364a01cedb..3c4213bbb29 100644 --- a/tests/unified_google_tests/test_google_ai_studio.py +++ b/tests/unified_google_tests/test_google_ai_studio.py @@ -10,6 +10,8 @@ import json class TestGoogleGenAIStudio(BaseGoogleGenAITest, BaseGoogleGenAIProxySDKTest): """Test Google GenAI Studio""" + test_non_streaming_base = None + @property def model_config(self): return { diff --git a/tests/unified_google_tests/test_litellm_responses_bridge.py b/tests/unified_google_tests/test_litellm_responses_bridge.py index b2489dfe2a9..d32e0cccc73 100644 --- a/tests/unified_google_tests/test_litellm_responses_bridge.py +++ b/tests/unified_google_tests/test_litellm_responses_bridge.py @@ -15,6 +15,8 @@ from tests.unified_google_tests.base_interactions_test import ( class TestLiteLLMResponsesBridge(BaseInteractionsTest): """Test LiteLLM Responses bridge using the base test suite.""" + test_create_streaming = None + def get_model(self) -> str: """Return the model string for the bridge provider. From 2b85808011e2fd122888976edccce41c298b16a8 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Fri, 2 Oct 2026 10:28:04 -0700 Subject: [PATCH 009/139] feat(mcp)!: disable stdio MCP servers by default (#44066) * feat(mcp)!: disable stdio MCP servers by default stdio MCP servers now only run when the proxy is started with LITELLM_ENABLE_MCP_STDIO=true. While it is off, existing stdio servers stay registered but never start: tool listings skip them quietly, direct tool calls and health checks return a 403 naming the env var, and creating or updating a stdio server is rejected. The flag is read from the process environment only, so DB-stored environment_variables cannot turn it on. The UI reads mcp_stdio_enabled from /.well-known/litellm-ui-config to grey out the stdio transport, show a banner on stdio forms, and badge stdio server cards. BREAKING CHANGE: stdio MCP servers are off by default. Set LITELLM_ENABLE_MCP_STDIO=true in the proxy environment and restart to keep using them. * fix(mcp): ignore stdio flag from config file and read UI flag from the selected worker LITELLM_ENABLE_MCP_STDIO set under environment_variables in config.yaml is now skipped like the DB-stored value, so only the process environment can enable stdio. The dashboard reads mcp_stdio_enabled from the proxy it is managing, so a control plane shows each worker's own setting. * test(mcp): cover non-mapping payloads in the shared transport validator * fix(mcp): skip blocked stdio servers quietly in every listing and keep the UI unchanged until the flag loads Prompt, resource and resource-template listings now skip a blocked stdio server at debug level like tool listing does, instead of logging a warning per server on every call. The dashboard only treats stdio as disabled once the proxy explicitly reports mcp_stdio_enabled false, so a proxy with the flag on, or an older one without the field, renders exactly as before with no flicker while loading. * fix(mcp): route blocked stdio tool calls to the flag error and warn once per server A gateway tools/call naming a blocked stdio server's tool now returns the LITELLM_ENABLE_MCP_STDIO message instead of "Tool not found". The "will not start" warning moves out of build_mcp_server_from_table, which DB reload re-runs on every cycle for rows with a NULL updated_at and which drafts and test-connection also call. It now fires when a row first enters the registry or changes transport. * fix(ui): explain on the server detail page why a stdio server is inert The Overview and MCP Tools tabs showed "No tools available" with no reason while stdio is disabled. The detail page now shows the same warning banner as the edit form, and hands off to the form's banner once editing starts. * refactor(ui): name the stdio banner conditions on the server detail page Keeps local/no-long-condition-chain within its budget * fix(proxy): log the ignored DB-stored LITELLM_ENABLE_MCP_STDIO warning once The DB config sync re-reads environment_variables on every cycle, so a stored flag logged the warning on each sync per worker --- .circleci/scripts/run_integration.sh | 2 +- .../mcp_server/mcp_server_manager.py | 40 ++- .../_experimental/mcp_server/stdio_gate.py | 28 +++ litellm/proxy/_types.py | 61 +++-- .../ui_discovery_endpoints.py | 2 + litellm/proxy/proxy_server.py | 16 ++ .../ui_discovery_endpoints.py | 1 + tests/mcp_tests/test_proxy_mcp_e2e.py | 1 + .../mcp_server/test_mcp_server_manager.py | 228 +++++++++++++++++- .../mcp_server/test_rest_endpoints.py | 4 + .../test_ui_discovery_endpoints.py | 20 ++ .../test_mcp_connector_import.py | 12 +- .../test_mcp_management_endpoints.py | 3 +- .../proxy/proxy_server/test_proxy_config.py | 44 ++++ tests/unit/proxy/test__types.py | 63 +++++ .../CreateMCPServer.integration.test.tsx | 23 ++ .../_components/CreateMCPServer.tsx | 10 +- .../_components/MCPServerCard.test.tsx | 33 +++ .../mcp-servers/_components/MCPServerCard.tsx | 16 ++ .../_components/StdioAvailability.test.tsx | 110 +++++++++ .../_components/StdioAvailability.tsx | 37 +++ .../_components/mcp_server_edit.test.tsx | 50 ++++ .../_components/mcp_server_edit.tsx | 10 +- .../_components/mcp_server_view.test.tsx | 25 ++ .../_components/mcp_server_view.tsx | 10 +- .../_components/mcp_servers.test.tsx | 89 +++++++ .../mcp-servers/_components/mcp_servers.tsx | 6 + .../src/components/networking.tsx | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 + 29 files changed, 896 insertions(+), 54 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/stdio_gate.py create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/StdioAvailability.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/StdioAvailability.tsx diff --git a/.circleci/scripts/run_integration.sh b/.circleci/scripts/run_integration.sh index f3166e5db39..c03220224d2 100644 --- a/.circleci/scripts/run_integration.sh +++ b/.circleci/scripts/run_integration.sh @@ -169,7 +169,7 @@ start_proxy() { INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \ LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" LITELLM_UI_PATH="$LITELLM_UI_PATH" PROXY_BASE_URL="http://127.0.0.1:$port" \ LITELLM_LICENSE="${LITELLM_LICENSE:-}" \ - LITELLM_MODE=PRODUCTION STORE_MODEL_IN_DB=True "${cost_map_env[@]}" \ + LITELLM_MODE=PRODUCTION STORE_MODEL_IN_DB=True LITELLM_ENABLE_MCP_STDIO=true "${cost_map_env[@]}" \ AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 COVERAGE_FILE="$coverage_data" \ "${proxy_command[@]}" --config tests/integration/proxy_config.yaml \ --host 127.0.0.1 --port "$port" --num_workers 1 --telemetry False \ diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0682217e481..375c010a9a8 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -142,6 +142,12 @@ from litellm.proxy._experimental.mcp_server.result_conversion import ( from litellm.proxy._experimental.mcp_server.sampling_handler import ( MCP_SAMPLING_AVAILABLE, ) +from litellm.proxy._experimental.mcp_server.stdio_gate import ( + MCP_STDIO_DISABLED_MESSAGE, + is_mcp_stdio_blocked, + is_mcp_stdio_enabled, + warn_if_mcp_stdio_blocked, +) from litellm.proxy._experimental.mcp_server.tool_catalog_guard import ( CatalogAlert, apply_description_overrides, @@ -2464,6 +2470,7 @@ class MCPServerManager: alias=alias, server_name=server_name, ) + warn_if_mcp_stdio_blocked(server_name, server_config.get("transport")) auth_type = server_config.get("auth_type", None) manual_issuer = _blank_to_none(server_config.get("issuer")) @@ -3282,6 +3289,7 @@ class MCPServerManager: # `credentials` field is the only one still encrypted here). # Re-decrypting plaintext would zero the values, so build with # env_vars_are_encrypted=False. + self._warn_if_newly_blocked_stdio(mcp_server, None) new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False) self._assign_unique_short_prefix(new_server) self._invalidate_server_definition_caches(mcp_server.server_id) @@ -4280,6 +4288,8 @@ class MCPServerManager: # Handle stdio transport if transport == MCPTransport.stdio: + if not is_mcp_stdio_enabled(): + raise HTTPException(status_code=403, detail=MCP_STDIO_DISABLED_MESSAGE) resolved_env: Final = ( stdio_env if stdio_env is not None @@ -4439,6 +4449,9 @@ class MCPServerManager: global_mcp_tool_registry, ) + if self._skip_blocked_stdio_listing(server, "tool"): + return [] + verbose_logger.debug("Connecting to url: %s", server.url) verbose_logger.info("_get_tools_from_server for %s...", server.name) @@ -4638,6 +4651,19 @@ class MCPServerManager: ) return server.server_id, hashlib.sha256(material.encode()).hexdigest() + @staticmethod + def _warn_if_newly_blocked_stdio(row: LiteLLM_MCPServerTable, previous: MCPServer | None) -> None: + if previous is None or previous.transport != row.transport: + warn_if_mcp_stdio_blocked(row.alias or row.server_name, row.transport) + + def _skip_blocked_stdio_listing(self, server: MCPServer, listing: str) -> bool: + if not is_mcp_stdio_blocked(server.transport): + return False + verbose_logger.debug( + "Skipping %s listing for MCP server %s: %s", listing, server.name, MCP_STDIO_DISABLED_MESSAGE + ) + return True + async def get_prompts_from_server( self, server: MCPServer, @@ -4648,6 +4674,8 @@ class MCPServerManager: raw_headers: dict[str, str] | None = None, client_ip: str | None = None, ) -> list[Prompt]: + if self._skip_blocked_stdio_listing(server, "prompt"): + return [] try: headers: Final = ( dict( @@ -4694,6 +4722,8 @@ class MCPServerManager: raw_headers: dict[str, str] | None = None, client_ip: str | None = None, ) -> list[Resource]: + if self._skip_blocked_stdio_listing(server, "resource"): + return [] try: headers: Final = ( dict( @@ -4740,6 +4770,8 @@ class MCPServerManager: raw_headers: dict[str, str] | None = None, client_ip: str | None = None, ) -> list[ResourceTemplate]: + if self._skip_blocked_stdio_listing(server, "resource template"): + return [] try: headers: Final = ( dict( @@ -6324,6 +6356,8 @@ class MCPServerManager: mcp_server = fallback if mcp_server is None: raise ValueError(f"Tool {name} not found") + if is_mcp_stdio_blocked(mcp_server.transport): + raise HTTPException(status_code=403, detail=MCP_STDIO_DISABLED_MESSAGE) if resolved_by_server_name_only and not self.server_exposes_tool(mcp_server, name): raise ValueError(f"Tool {name} not found") @@ -6708,7 +6742,10 @@ class MCPServerManager: if matched is not None: matched_prefix, original_tool_name = matched matched_server: Final = prefix_to_server.get(matched_prefix) - if matched_server is not None and self.server_exposes_tool(matched_server, original_tool_name): + if matched_server is not None and ( + self.server_exposes_tool(matched_server, original_tool_name) + or is_mcp_stdio_blocked(matched_server.transport) + ): return matched_server return None @@ -6770,6 +6807,7 @@ class MCPServerManager: alias=getattr(server, "alias", None), server_name=getattr(server, "server_name", None), ) + self._warn_if_newly_blocked_stdio(server, existing_server) verbose_logger.debug("Building server from DB: %s (%s)", server.server_id, server.server_name) # raw_rows come straight from the DB, so their global env var # values (like credentials) are still encrypted here, unlike the diff --git a/litellm/proxy/_experimental/mcp_server/stdio_gate.py b/litellm/proxy/_experimental/mcp_server/stdio_gate.py new file mode 100644 index 00000000000..00fde229e2f --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/stdio_gate.py @@ -0,0 +1,28 @@ +import os +from typing import Final + +from litellm._logging import verbose_logger +from litellm.types.mcp import MCPTransport + +MCP_STDIO_ENABLED_ENV_VAR: Final = "LITELLM_ENABLE_MCP_STDIO" +MCP_STDIO_DISABLED_MESSAGE: Final = ( + f"stdio MCP servers are disabled on this proxy. " + f"Set {MCP_STDIO_ENABLED_ENV_VAR}=true on the proxy and restart to enable them" +) + + +def is_mcp_stdio_enabled() -> bool: + return os.getenv(MCP_STDIO_ENABLED_ENV_VAR, "").strip().lower() == "true" + + +def is_mcp_stdio_flag_key(env_var_name: str) -> bool: + return env_var_name.upper() == MCP_STDIO_ENABLED_ENV_VAR + + +def is_mcp_stdio_blocked(transport: str | None) -> bool: + return transport == MCPTransport.stdio and not is_mcp_stdio_enabled() + + +def warn_if_mcp_stdio_blocked(server_name: str | None, transport: str | None) -> None: + if is_mcp_stdio_blocked(transport): + verbose_logger.warning("MCP server '%s' will not start: %s", server_name, MCP_STDIO_DISABLED_MESSAGE) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index bb113715d2c..988590bfdef 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -28,6 +28,7 @@ from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( validate_langfuse_span_scope_value, validate_no_callback_env_reference, ) +from litellm.proxy._experimental.mcp_server.stdio_gate import MCP_STDIO_DISABLED_MESSAGE, is_mcp_stdio_enabled from litellm.types.agents import AgentCaller, AgentResponse from litellm.types.integrations.compression_interception import ( CompressionSavingsMetadata, @@ -1583,6 +1584,30 @@ def _reject_unsupported_per_server_oauth_discovery(values: object, require_auth_ raise _per_server_oauth_discovery_error() +def _validate_mcp_transport_fields(values: object) -> None: + if not isinstance(values, dict): + return + transport: Final = values.get("transport") + if transport in (MCPTransport.http, MCPTransport.sse): + if not values.get("url") and not values.get("spec_path"): + raise ValueError("url or spec_path is required for HTTP/SSE transport") + return + if transport != MCPTransport.stdio: + return + if not is_mcp_stdio_enabled(): + raise ValueError(MCP_STDIO_DISABLED_MESSAGE) + command: Final = values.get("command") + if not command: + raise ValueError("command is required for stdio transport") + if not values.get("args"): + raise ValueError("args is required for stdio transport") + if os.path.basename(str(command)) not in MCP_STDIO_ALLOWED_COMMANDS: + raise ValueError( + f"Command '{command}' is not in the allowed commands list " + f"for stdio transport. Allowed commands: {sorted(MCP_STDIO_ALLOWED_COMMANDS)}" + ) + + class NewMCPServerRequest(LiteLLMPydanticObjectBase): server_id: str | None = None server_name: str | None = None @@ -1649,23 +1674,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): @model_validator(mode="before") @classmethod def validate_transport_fields(cls, values): - if isinstance(values, dict): - transport: Final = values.get("transport") - if transport == MCPTransport.stdio: - if not values.get("command"): - raise ValueError("command is required for stdio transport") - if not values.get("args"): - raise ValueError("args is required for stdio transport") - # Validate command against allowlist to prevent arbitrary execution - base_command: Final = os.path.basename(values["command"]) - if base_command not in MCP_STDIO_ALLOWED_COMMANDS: - raise ValueError( - f"Command '{values['command']}' is not in the allowed commands list " - f"for stdio transport. Allowed commands: {sorted(MCP_STDIO_ALLOWED_COMMANDS)}" - ) - elif transport in [MCPTransport.http, MCPTransport.sse]: - if not values.get("url") and not values.get("spec_path"): - raise ValueError("url or spec_path is required for HTTP/SSE transport") + _validate_mcp_transport_fields(values) return values @model_validator(mode="before") @@ -1748,23 +1757,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): @model_validator(mode="before") @classmethod def validate_transport_fields(cls, values): - if isinstance(values, dict): - transport: Final = values.get("transport") - if transport == MCPTransport.stdio: - if not values.get("command"): - raise ValueError("command is required for stdio transport") - if not values.get("args"): - raise ValueError("args is required for stdio transport") - # Validate command against allowlist to prevent arbitrary execution - base_command: Final = os.path.basename(values["command"]) - if base_command not in MCP_STDIO_ALLOWED_COMMANDS: - raise ValueError( - f"Command '{values['command']}' is not in the allowed commands list " - f"for stdio transport. Allowed commands: {sorted(MCP_STDIO_ALLOWED_COMMANDS)}" - ) - elif transport in [MCPTransport.http, MCPTransport.sse]: - if not values.get("url") and not values.get("spec_path"): - raise ValueError("url or spec_path is required for HTTP/SSE transport") + _validate_mcp_transport_fields(values) return values @model_validator(mode="before") diff --git a/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py b/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py index 8b042d18cd0..4c85d6148c0 100644 --- a/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py +++ b/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py @@ -4,6 +4,7 @@ from typing import Final from fastapi import APIRouter +from litellm.proxy._experimental.mcp_server.stdio_gate import is_mcp_stdio_enabled from litellm.proxy.common_utils.html_forms.default_credentials_hint import should_hide_default_credentials_hint from litellm.types.proxy.discovery_endpoints.ui_discovery_endpoints import ( UiDiscoveryEndpoints, @@ -41,4 +42,5 @@ async def get_ui_config(): hide_default_credentials_hint=hide_default_credentials_hint, is_control_plane=is_control_plane, workers=proxy_config.worker_registry if is_control_plane else [], + mcp_stdio_enabled=is_mcp_stdio_enabled(), ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5d47cceee02..e406e6faec4 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -335,6 +335,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.openai_like.model_info import MODEL_INFO_REFRESH_SECONDS from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.proxy._experimental.mcp_server.byok_credential_cache import byok_credential_cache +from litellm.proxy._experimental.mcp_server.stdio_gate import MCP_STDIO_ENABLED_ENV_VAR, is_mcp_stdio_flag_key from litellm.proxy._lazy_features import attach_lazy_features, reserve_lazy_slot from litellm.proxy._types import * from litellm.proxy.analytics_endpoints.analytics_endpoints import ( @@ -5407,6 +5408,7 @@ class ProxyConfig: self._last_cyberark_config: dict[str, object] | None = None # mutable-ok: change-detection cache self._last_cleanup_schedule_attempt: tuple[object, ...] | None = None self._cleanup_reschedule_failed: bool = False + self._warned_db_mcp_stdio_flag_ignored: bool = False self._cyberark_boot_env: dict[str, str | None] | None = None # mutable-ok: deployment env snapshot, set once self.worker_registry: list[WorkerRegistryEntry] = [] self.config_sync_subscriber: ConfigSyncSubscriber | None = None @@ -6190,6 +6192,12 @@ class ProxyConfig: if key in self._BLOCKED_ENV_KEYS: verbose_proxy_logger.warning("Skipping blocked environment variable key: %s", key) continue + if isinstance(key, str) and is_mcp_stdio_flag_key(key): + verbose_proxy_logger.warning( + "Ignoring %s set in the config file. Set it in the proxy's environment instead", + MCP_STDIO_ENABLED_ENV_VAR, + ) + continue ######################################################### # handles this scenario: # ```yaml @@ -7591,6 +7599,14 @@ class ProxyConfig: """ decrypted_env_vars: Final = {} for k, v in environment_variables.items(): + if isinstance(k, str) and is_mcp_stdio_flag_key(k): + if not self._warned_db_mcp_stdio_flag_ignored: + verbose_proxy_logger.warning( + "Ignoring %s stored in the database. Set it in the proxy's environment instead", + MCP_STDIO_ENABLED_ENV_VAR, + ) + self._warned_db_mcp_stdio_flag_ignored = True + continue try: decrypted_value = decrypt_value_helper(value=v, key=k, return_original_value=return_original_value) if decrypted_value is not None: diff --git a/litellm/types/proxy/discovery_endpoints/ui_discovery_endpoints.py b/litellm/types/proxy/discovery_endpoints/ui_discovery_endpoints.py index d52877b7c1e..4e0faabc132 100644 --- a/litellm/types/proxy/discovery_endpoints/ui_discovery_endpoints.py +++ b/litellm/types/proxy/discovery_endpoints/ui_discovery_endpoints.py @@ -11,4 +11,5 @@ class UiDiscoveryEndpoints(BaseModel): sso_configured: bool hide_default_credentials_hint: bool = False is_control_plane: bool = False + mcp_stdio_enabled: bool = False workers: list[WorkerRegistryEntry] = [] diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index a730f6c10ee..92f26e4ab7e 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -55,6 +55,7 @@ def _clear_proxy_database_env() -> typing.Iterator[None]: # the config file. We must set it here so the lifespan doesn't reset it to None. mp.setenv("LITELLM_MASTER_KEY", "sk-1234") mp.setenv("LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY", "true") + mp.setenv("LITELLM_ENABLE_MCP_STDIO", "true") try: yield finally: diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index a900ad50dfb..46264de738a 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -314,8 +314,9 @@ class TestMCPServerManager: with patch.object(manager, "_get_general_settings", return_value={}): assert manager.get_mcp_server_by_id(server.server_id, client_ip="8.8.8.8") is None - async def test_create_mcp_client_stdio(self): + async def test_create_mcp_client_stdio(self, monkeypatch): """Test creating MCP client for stdio transport""" + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") manager = MCPServerManager() stdio_server = MCPServer( @@ -458,11 +459,12 @@ class TestMCPServerManager: assert exc_info.value.status_code == 500 assert "oauth2_id_jag" in str(exc_info.value.detail) - async def test_create_mcp_client_stdio_injects_npm_config_cache(self): + async def test_create_mcp_client_stdio_injects_npm_config_cache(self, monkeypatch): """Test that _create_mcp_client injects NPM_CONFIG_CACHE when not already set, and preserves user-provided NPM_CONFIG_CACHE when present.""" from litellm.constants import MCP_NPM_CACHE_DIR + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") manager = MCPServerManager() # Case 1: NPM_CONFIG_CACHE not set -> should be injected @@ -491,6 +493,173 @@ class TestMCPServerManager: client2 = await manager._create_mcp_client(server_with_cache) assert client2.stdio_config["env"]["NPM_CONFIG_CACHE"] == "/custom/cache" + async def test_create_mcp_client_refuses_to_start_a_stdio_server_while_stdio_is_not_enabled(self, monkeypatch): + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + manager = MCPServerManager() + server = MCPServer( + server_id="stdio-off", + name="stdio_off", + transport=MCPTransport.stdio, + command="python", + args=["server.py"], + ) + + with pytest.raises(HTTPException) as exc_info: + await manager._create_mcp_client(server) + + assert exc_info.value.status_code == 403 + assert "LITELLM_ENABLE_MCP_STDIO=true" in str(exc_info.value.detail) + + @pytest.mark.parametrize( + "listing", + [ + lambda manager, server: manager._get_tools_from_server(server), + lambda manager, server: manager.get_prompts_from_server(server, user_api_key_auth=None), + lambda manager, server: manager.get_resources_from_server(server, user_api_key_auth=None), + lambda manager, server: manager.get_resource_templates_from_server(server, user_api_key_auth=None), + ], + ids=["tools", "prompts", "resources", "resource_templates"], + ) + async def test_listing_skips_a_stdio_server_quietly_while_stdio_is_not_enabled(self, monkeypatch, caplog, listing): + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + manager = MCPServerManager() + server = MCPServer( + server_id="stdio-quiet", + name="stdio_quiet", + transport=MCPTransport.stdio, + command="python", + args=["server.py"], + ) + + with caplog.at_level(logging.DEBUG, logger="LiteLLM"): + items = await listing(manager, server) + + assert items == [] + assert any("stdio_quiet" in r.getMessage() for r in caplog.records if r.levelno == logging.DEBUG) + assert not [r for r in caplog.records if r.levelno >= logging.WARNING] + + async def test_calling_a_tool_on_a_stdio_server_names_the_flag_while_stdio_is_not_enabled(self, monkeypatch): + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + manager = MCPServerManager() + server = MCPServer( + server_id="stdio-call", + name="stdio_call", + alias="stdio_call", + transport=MCPTransport.stdio, + command="python", + args=["server.py"], + ) + manager.registry[server.server_id] = server + + with pytest.raises(HTTPException) as exc_info: + manager._resolve_mcp_server_for_tool_call(server_name="stdio_call", name="echo") + + assert exc_info.value.status_code == 403 + assert "LITELLM_ENABLE_MCP_STDIO=true" in str(exc_info.value.detail) + + async def test_calling_an_unknown_tool_on_an_enabled_stdio_server_is_still_not_found(self, monkeypatch): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") + manager = MCPServerManager() + server = MCPServer( + server_id="stdio-call", + name="stdio_call", + alias="stdio_call", + transport=MCPTransport.stdio, + command="python", + args=["server.py"], + ) + manager.registry[server.server_id] = server + + with pytest.raises(ValueError, match="Tool echo not found"): + manager._resolve_mcp_server_for_tool_call(server_name="stdio_call", name="echo") + + @pytest.mark.parametrize("flag, routed", [(None, True), ("true", False)]) + async def test_a_prefixed_tool_name_routes_to_its_blocked_stdio_server(self, monkeypatch, flag, routed): + if flag is None: + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + else: + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", flag) + manager = MCPServerManager() + server = MCPServer( + server_id="stdio-route", + name="stdio_route", + alias="stdio_route", + transport=MCPTransport.stdio, + command="python", + args=["server.py"], + ) + manager.registry[server.server_id] = server + + resolved = manager._get_mcp_server_from_tool_name("stdio_route-echo") + + assert (resolved is server) is routed + + async def test_health_check_reports_a_stdio_server_unhealthy_with_the_flag_to_set(self, monkeypatch): + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + manager = MCPServerManager() + server = MCPServer( + server_id="stdio-health", + name="stdio_health", + transport=MCPTransport.stdio, + command="python", + args=["server.py"], + ) + manager.registry[server.server_id] = server + + result = await manager.health_check_server(server.server_id) + + assert result.status == "unhealthy" + assert "LITELLM_ENABLE_MCP_STDIO=true" in (result.health_check_error or "") + + async def test_a_config_stdio_server_stays_registered_and_warns_while_stdio_is_not_enabled( + self, monkeypatch, config_only_mcp_manager_factory, caplog + ): + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + manager = config_only_mcp_manager_factory() + config = {"local_tools": {"transport": MCPTransport.stdio, "command": "python", "args": ["server.py"]}} + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.load_servers_from_config(config) + + assert [s.server_name for s in manager.config_mcp_servers.values()] == ["local_tools"] + warnings = [m for m in caplog.messages if "local_tools" in m] + assert len(warnings) == 1 + assert "LITELLM_ENABLE_MCP_STDIO=true" in warnings[0] + + async def test_a_config_stdio_server_loads_without_a_warning_once_stdio_is_enabled( + self, monkeypatch, config_only_mcp_manager_factory, caplog + ): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") + manager = config_only_mcp_manager_factory() + config = {"local_tools": {"transport": MCPTransport.stdio, "command": "python", "args": ["server.py"]}} + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.load_servers_from_config(config) + + assert [s.server_name for s in manager.config_mcp_servers.values()] == ["local_tools"] + assert not [m for m in caplog.messages if "LITELLM_ENABLE_MCP_STDIO" in m] + + async def test_a_db_stdio_server_stays_registered_and_warns_while_stdio_is_not_enabled(self, monkeypatch, caplog): + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + manager = MCPServerManager() + row = LiteLLM_MCPServerTable( + server_id="db-stdio", + alias="db_stdio", + transport=MCPTransport.stdio, + command="python", + args=["server.py"], + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.add_server(row) + await manager.update_server(row) + await manager.update_server(row) + + assert "db-stdio" in manager.registry + assert sum("db_stdio" in m and "LITELLM_ENABLE_MCP_STDIO=true" in m for m in caplog.messages) == 1 + def test_build_stdio_env_only_accepts_x_prefixed_placeholders(self): """Ensure only ${X-*} placeholders are substituted from headers.""" manager = MCPServerManager() @@ -9876,7 +10045,8 @@ class TestCreateMcpClientV2Graft: assert exc.value.status_code == 500 assert "credential" in str(exc.value.detail) - async def test_stdio_migrated_auth_type_still_defers_to_v1(self): + async def test_stdio_migrated_auth_type_still_defers_to_v1(self, monkeypatch): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") client = await MCPServerManager()._create_mcp_client( MCPServer( server_id="stdio-graft", @@ -13563,7 +13733,10 @@ async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_reques @pytest.mark.asyncio @pytest.mark.parametrize("transport", ["http", "stdio"]) -async def test_debug_reports_legacy_signing_and_non_http_transport(_mcp_request_ctx, transport: Literal["http", "stdio"]) -> None: +async def test_debug_reports_legacy_signing_and_non_http_transport( + _mcp_request_ctx, monkeypatch, transport: Literal["http", "stdio"] +) -> None: + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var from starlette.requests import Request @@ -14991,6 +15164,53 @@ class TestSharedIdentifierPrefixWarning: assert "'shared'" in shared_warnings[0] +@pytest.mark.asyncio +@pytest.mark.parametrize( + "flag,transports,expected_warnings", + [ + (None, ["stdio", "stdio", "stdio"], 1), + (None, ["http", "stdio", "stdio"], 1), + ("true", ["stdio", "stdio", "stdio"], 0), + ], +) +async def test_reload_warns_once_about_a_blocked_stdio_row_that_is_rebuilt_every_time( + monkeypatch, caplog, flag, transports, expected_warnings +): + if flag is None: + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + else: + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", flag) + manager = MCPServerManager() + repository = MagicMock() + + async def build_from_table(table, **_kwargs): + return MCPServer(server_id=table.server_id, name=table.server_name, transport=table.transport) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + return_value=repository, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch.object(manager, "build_mcp_server_from_table", new=build_from_table), + patch.object(manager, "_maybe_register_openapi_tools", new=AsyncMock()), + patch.object(manager, "_prime_oauth_metadata_discovery_for_servers"), + caplog.at_level(logging.WARNING, logger="LiteLLM"), + ): + for transport in transports: + row = LiteLLM_MCPServerTable( + server_id="srv-null-ts", server_name="null_ts", transport=transport, command="python", updated_at=None + ) + repository.table.find_many = AsyncMock(return_value=[MagicMock(model_dump=row.model_dump)]) + await manager.reload_servers_from_database() + + assert manager.registry["srv-null-ts"].transport == transports[-1] + assert sum("'null_ts' will not start" in m for m in caplog.messages) == expected_warnings + + @pytest.mark.asyncio @pytest.mark.parametrize("revision", ["auto", "2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]) async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_manager_factory, revision): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py index 759014b54c5..680b84469d9 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -3730,6 +3730,10 @@ class TestGetToolsForSingleServer: class TestStdioCommandAllowlist: """Tests for MCP stdio command allowlist validation.""" + @pytest.fixture(autouse=True) + def _stdio_enabled(self, monkeypatch): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") + def test_allowed_command_passes_validation(self): """npx, uvx, python, etc. should be accepted.""" req = NewMCPServerRequest( diff --git a/tests/unit/proxy/discovery_endpoints/test_ui_discovery_endpoints.py b/tests/unit/proxy/discovery_endpoints/test_ui_discovery_endpoints.py index 64a2eb69325..34a7f78659b 100644 --- a/tests/unit/proxy/discovery_endpoints/test_ui_discovery_endpoints.py +++ b/tests/unit/proxy/discovery_endpoints/test_ui_discovery_endpoints.py @@ -460,3 +460,23 @@ def test_ui_discovery_endpoints_is_control_plane_false_when_no_workers(): data = response.json() assert data["is_control_plane"] is False assert data["workers"] == [] + + +@pytest.mark.parametrize(("flag", "expected"), [(None, False), ("false", False), ("true", True)]) +def test_ui_config_tells_the_dashboard_whether_stdio_mcp_servers_are_enabled(monkeypatch, flag, expected): + if flag is None: + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + else: + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", flag) + app = FastAPI() + app.include_router(router) + + with ( + patch("litellm.proxy.utils.get_server_root_path", return_value="/"), + patch("litellm.proxy.utils.get_proxy_base_url", return_value=None), + patch("litellm.proxy.auth.auth_utils.has_user_setup_sso", return_value=False), + ): + response = TestClient(app).get("/.well-known/litellm-ui-config") + + assert response.status_code == 200 + assert response.json()["mcp_stdio_enabled"] is expected diff --git a/tests/unit/proxy/management_endpoints/test_mcp_connector_import.py b/tests/unit/proxy/management_endpoints/test_mcp_connector_import.py index 9b1a0fb4f98..d60cc1fbb15 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_connector_import.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_connector_import.py @@ -124,7 +124,8 @@ class TestConvertMcpServersMapping: assert isinstance(result, ConvertedConnector) assert result.request.transport == MCPTransport.sse - def test_stdio_connector(self): + def test_stdio_connector(self, monkeypatch): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") result = _single( { "mcpServers": { @@ -142,11 +143,18 @@ class TestConvertMcpServersMapping: assert result.request.args == ["-y", "@example/mcp-server"] assert result.request.env == {"API_KEY": "value"} - def test_disallowed_stdio_command_returns_error(self): + def test_disallowed_stdio_command_returns_error(self, monkeypatch): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") result = _single({"mcpServers": {"evil": {"command": "rm", "args": ["-rf", "/"]}}}) assert isinstance(result, ConnectorConversionError) assert "not in the allowed commands list" in result.error + def test_stdio_connector_is_reported_as_an_error_while_stdio_is_not_enabled(self, monkeypatch): + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + result = _single({"mcpServers": {"local": {"command": "npx", "args": ["-y", "@example/mcp-server"]}}}) + assert isinstance(result, ConnectorConversionError) + assert "LITELLM_ENABLE_MCP_STDIO=true" in result.error + def test_unsupported_type_returns_error(self): result = _single({"mcpServers": {"ws": {"type": "websocket", "url": "wss://x.example"}}}) assert isinstance(result, ConnectorConversionError) diff --git a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py index f5fc5ae24d4..88ee36a0fef 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -4881,7 +4881,8 @@ class TestMCPApprovalWorkflow: assert "team" in str(exc_info.value.detail).lower() @pytest.mark.asyncio - async def test_register_mcp_server_rejects_stdio_transport(self): + async def test_register_mcp_server_rejects_stdio_transport(self, monkeypatch): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") # stdio servers spawn a local subprocess on the proxy host. Accepting # them from the non-admin submission endpoint would let a team member # propose a config that an admin could rubber-stamp into local code diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index 51ecc210119..7fdf9277154 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -3527,6 +3527,50 @@ def test_ProxyConfig__decrypt_and_set_db_env_variables_sets_env(monkeypatch): } +@pytest.mark.parametrize("stored_key", ["LITELLM_ENABLE_MCP_STDIO", "litellm_enable_mcp_stdio"]) +def test_ProxyConfig__decrypt_and_set_db_env_variables_cannot_enable_mcp_stdio(monkeypatch, stored_key): + monkeypatch.setattr( + "litellm.proxy.proxy_server.decrypt_value_helper", + lambda value, key, return_original_value=False: value, + ) + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + monkeypatch.delenv(stored_key, raising=False) + monkeypatch.delenv("KEY_X", raising=False) + pc = ProxyConfig() + out = pc._decrypt_and_set_db_env_variables({stored_key: "true", "KEY_X": "x"}) + assert out == {"KEY_X": "x"} + assert os.environ.get("KEY_X") == "x" + assert os.environ.get(stored_key) is None + assert os.environ.get("LITELLM_ENABLE_MCP_STDIO") is None + + +def test_ProxyConfig__decrypt_and_set_db_env_variables_warns_once_about_the_ignored_mcp_stdio_flag( + monkeypatch, caplog +): + monkeypatch.setattr( + "litellm.proxy.proxy_server.decrypt_value_helper", + lambda value, key, return_original_value=False: value, + ) + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + pc = ProxyConfig() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + for _ in range(3): + pc._decrypt_and_set_db_env_variables({"LITELLM_ENABLE_MCP_STDIO": "true"}) + assert os.environ.get("LITELLM_ENABLE_MCP_STDIO") is None + assert sum("Ignoring LITELLM_ENABLE_MCP_STDIO stored in the database" in m for m in caplog.messages) == 1 + + +@pytest.mark.parametrize("config_key", ["LITELLM_ENABLE_MCP_STDIO", "litellm_enable_mcp_stdio"]) +def test_ProxyConfig__load_environment_variables_cannot_enable_mcp_stdio(monkeypatch, config_key): + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + monkeypatch.delenv(config_key, raising=False) + monkeypatch.delenv("KEY_X", raising=False) + ProxyConfig()._load_environment_variables({"environment_variables": {config_key: "true", "KEY_X": "x"}}) + assert os.environ.get("KEY_X") == "x" + assert os.environ.get(config_key) is None + assert os.environ.get("LITELLM_ENABLE_MCP_STDIO") is None + + def test_ProxyConfig__decrypt_and_set_db_env_variables_invalid_dict_raises(): pc = ProxyConfig() with pytest.raises(AttributeError): diff --git a/tests/unit/proxy/test__types.py b/tests/unit/proxy/test__types.py index adc3bc04bdf..50c3eb2908c 100644 --- a/tests/unit/proxy/test__types.py +++ b/tests/unit/proxy/test__types.py @@ -11,10 +11,12 @@ from litellm.proxy._types import ( LiteLLM_AuditLogs, LiteLLM_TeamMembership, LitellmUserRoles, + NewMCPServerRequest, NewUserRequest, OrganizationMemberUpdateRequest, ResetSpendRequest, UpdateKeyRequest, + UpdateMCPServerRequest, UpdateUserRequest, UserAPIKeyAuth, ) @@ -403,3 +405,64 @@ def test_mcp_metadata_rejects_unavailable_upstream_protocol(revision): for model in (NewMCPServerRequest, UpdateMCPServerRequest): with pytest.raises(ValidationError): model.model_validate(payload) + + +MCP_SERVER_REQUESTS = (NewMCPServerRequest, UpdateMCPServerRequest) +STDIO_SERVER_FIELDS = {"server_id": "stdio-1", "transport": "stdio", "command": "python", "args": ["server.py"]} + + +@pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS) +def test_a_stdio_mcp_server_is_refused_while_stdio_is_not_enabled(monkeypatch, request_model): + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + + with pytest.raises(ValidationError, match="LITELLM_ENABLE_MCP_STDIO=true"): + request_model(**STDIO_SERVER_FIELDS) + + +@pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS) +@pytest.mark.parametrize("flag", ["true", "TRUE", " True "]) +def test_a_stdio_mcp_server_is_accepted_once_stdio_is_enabled(monkeypatch, request_model, flag): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", flag) + + assert request_model(**STDIO_SERVER_FIELDS).command == "python" + + +@pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS) +@pytest.mark.parametrize("flag", ["false", "1", "yes", ""]) +def test_only_an_explicit_true_enables_stdio_mcp_servers(monkeypatch, request_model, flag): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", flag) + + with pytest.raises(ValidationError, match="LITELLM_ENABLE_MCP_STDIO=true"): + request_model(**STDIO_SERVER_FIELDS) + + +@pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS) +def test_a_stdio_command_outside_the_allowlist_is_refused_even_when_stdio_is_enabled(monkeypatch, request_model): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") + + with pytest.raises(ValidationError, match="not in the allowed commands list"): + request_model(**{**STDIO_SERVER_FIELDS, "command": "/bin/sh"}) + + +@pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS) +@pytest.mark.parametrize("missing", ["command", "args"]) +def test_an_enabled_stdio_mcp_server_still_needs_a_command_and_args(monkeypatch, request_model, missing): + monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") + + with pytest.raises(ValidationError, match=f"{missing} is required for stdio transport"): + request_model(**{k: v for k, v in STDIO_SERVER_FIELDS.items() if k != missing}) + + +@pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS) +def test_an_http_mcp_server_is_unaffected_by_the_stdio_flag(monkeypatch, request_model): + monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) + + assert request_model(server_id="http-1", transport="http", url="https://mcp.example.com").url == "https://mcp.example.com" + with pytest.raises(ValidationError, match="url or spec_path is required"): + request_model(server_id="http-1", transport="http") + + +@pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS) +def test_a_non_mapping_mcp_server_payload_gets_a_validation_error(request_model): + with pytest.raises(ValidationError, match="valid dictionary"): + request_model.model_validate("not-a-server") diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx index 6ab11d1c9ea..3cf2bcdf238 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx @@ -1944,6 +1944,29 @@ describe("CreateMCPServer", () => { expect(nameInput).toHaveValue("github_mcp"); }); }); + + const sqlitePrefill = { + name: "sqlite", + title: "SQLite", + description: "Local database", + category: "Databases", + transport: "stdio", + command: "uvx", + args: ["mcp-server-sqlite"], + }; + + it("explains that a catalog stdio server cannot be added while the proxy has stdio off", async () => { + render(); + + expect(await screen.findByText("stdio is disabled on this proxy")).toBeInTheDocument(); + }); + + it("shows no stdio banner for a catalog stdio server once stdio is enabled", async () => { + render(); + + await waitFor(() => expect(getServerNameInput()).toHaveValue("sqlite")); + expect(screen.queryByText("stdio is disabled on this proxy")).not.toBeInTheDocument(); + }); }); describe("with back to discovery button", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx index f656bd2fc60..2d6e83aff9b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx @@ -48,6 +48,7 @@ import MCPServerCostConfig from "./mcp_server_cost_config"; import MCPConnectionStatus from "./mcp_connection_status"; import MCPToolConfiguration from "./mcp_tool_configuration"; import StdioConfiguration from "./StdioConfiguration"; +import { StdioDisabledBanner, TransportSelectItems } from "./StdioAvailability"; import MCPPermissionManagement from "./MCPPermissionManagement"; import OpenAPIFormSection, { OpenAPIKeyTool } from "./OpenAPIFormSection"; import MCPLogoSelector from "./MCPLogoSelector"; @@ -82,6 +83,7 @@ interface CreateMCPServerProps { existingServers?: MCPServer[]; prefillData?: DiscoverableMCPServer | null; onBackToDiscovery?: () => void; + stdioEnabled?: boolean; } const payloadErrorMessage = (result: Exclude): string => { @@ -113,6 +115,7 @@ const CreateMCPServer: React.FC = ({ existingServers, prefillData, onBackToDiscovery, + stdioEnabled = true, }) => { const form = useForm({ mode: "onChange", defaultValues: CREATE_DEFAULTS }); const registry = useMountRegistry(); @@ -750,11 +753,7 @@ const CreateMCPServer: React.FC = ({ - {TRANSPORT_ITEMS.map((item) => ( - - {item.label} - - ))} + )} @@ -917,6 +916,7 @@ const CreateMCPServer: React.FC = ({ {transportType !== "stdio" && transportType !== "" && isAwsSigV4AuthType && } {/* Stdio Configuration - only show for stdio transport */} + {transportType === "stdio" && !stdioEnabled && }
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx index d298d9d8145..532424a3d5b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx @@ -138,3 +138,36 @@ describe("MCPServerCard network access", () => { expect(screen.queryByText(/^Hub:/)).not.toBeInTheDocument(); }); }); + +describe("MCPServerCard stdio availability", () => { + const stdioServer = { transport: "stdio", url: undefined, command: "python", args: ["server.py"], auth_type: "none" }; + + it("flags a stdio server with how to enable stdio when the proxy has it off", async () => { + const user = userEvent.setup(); + render( + , + ); + + await user.hover(screen.getByText("stdio disabled")); + + expect( + await screen.findByText( + "stdio MCP servers are disabled on this proxy. Set LITELLM_ENABLE_MCP_STDIO=true on the proxy and restart to enable them", + ), + ).toBeInTheDocument(); + }); + + it("does not flag a stdio server when the proxy has stdio on", () => { + render(); + + expect(screen.getByText("STDIO")).toBeInTheDocument(); + expect(screen.queryByText("stdio disabled")).not.toBeInTheDocument(); + }); + + it("does not flag a non-stdio server when the proxy has stdio off", () => { + render(); + + expect(screen.getByText("HTTP")).toBeInTheDocument(); + expect(screen.queryByText("stdio disabled")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx index bb153f94665..a9bb2334372 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx @@ -14,6 +14,7 @@ import { cn } from "@/lib/cva.config"; import { AUTH_TYPE, MCP_REACHABLE_DESCRIPTION, type MCPServer } from "@/components/mcp_tools/types"; import { Logo } from "@/components/molecules/logo/Logo"; import { getMaskedAndFullUrl, getMCPNetworkAccess } from "./utils"; +import { STDIO_DISABLED_MESSAGE } from "./StdioAvailability"; interface MCPServerCardProps { server: MCPServer; @@ -29,6 +30,7 @@ interface MCPServerCardProps { onByokConnect?: () => void; onOpenFillFields?: () => void; onDelete?: () => void; + stdioEnabled?: boolean; } const HEALTH_TONE: Record = { @@ -52,6 +54,7 @@ const MCPServerCard: FC = ({ onByokConnect, onOpenFillFields, onDelete, + stdioEnabled = true, }) => { const alias = server.alias || server.server_name || ""; const name = server.server_name || alias || server.server_id; @@ -220,6 +223,19 @@ const MCPServerCard: FC = ({ /> {displayTransport.toUpperCase()} {authType} + {transport === "stdio" && !stdioEnabled && ( + + + + stdio disabled + + } + /> + {STDIO_DISABLED_MESSAGE} + + )} {oauthFlowUnset && ( + + + + + + + , + ); +} + +const option = (name: RegExp) => screen.getByRole("option", { name }); + +describe("TransportSelectItems", () => { + it("greys out only the stdio option and explains how to enable it when stdio is off", async () => { + const user = userEvent.setup(); + openTransportSelect(false); + + expect(option(/Standard Input\/Output \(stdio\)/)).toHaveAttribute("data-disabled"); + expect(option(/Streamable HTTP/)).not.toHaveAttribute("data-disabled"); + expect(option(/Server-Sent Events/)).not.toHaveAttribute("data-disabled"); + expect(option(/OpenAPI Spec/)).not.toHaveAttribute("data-disabled"); + + await user.hover(screen.getByLabelText("question-circle")); + + expect(await screen.findByText(STDIO_DISABLED_MESSAGE)).toBeInTheDocument(); + }); + + it("offers stdio like any other transport when stdio is on", () => { + openTransportSelect(true); + + expect(option(/Standard Input\/Output \(stdio\)/)).not.toHaveAttribute("data-disabled"); + expect(screen.queryByLabelText("question-circle")).not.toBeInTheDocument(); + }); +}); + +describe("useMcpStdioEnabled", () => { + afterEach(() => { + switchToWorkerUrl(null); + vi.restoreAllMocks(); + }); + + const renderWithConfig = (config: object) => { + const fetchSpy = vi + .spyOn(globalThis, "fetch") + .mockImplementation(async () => new Response(JSON.stringify(config), { status: 200 })); + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + const wrapper = ({ children }: { children: React.ReactNode }) => ( + {children} + ); + const { result } = renderHook(() => useMcpStdioEnabled(), { wrapper }); + const settled = () => + waitFor(() => + expect( + client + .getQueryCache() + .getAll() + .map((query) => query.state.status), + ).toEqual(["success"]), + ); + return { result, fetchSpy, settled }; + }; + + it("reads the flag from the worker the dashboard is managing", async () => { + switchToWorkerUrl("http://worker-b.example:4000"); + const { result, fetchSpy, settled } = renderWithConfig({ mcp_stdio_enabled: false }); + + await settled(); + expect(result.current).toBe(false); + expect(fetchSpy.mock.calls.map(([request]) => (request as Request).url)).toEqual([ + "http://worker-b.example:4000/.well-known/litellm-ui-config", + ]); + }); + + it.each([ + [{ mcp_stdio_enabled: false }, false], + [{ mcp_stdio_enabled: true }, true], + [{}, true], + ])("treats %o as stdio enabled=%s", async (config, enabled) => { + const { result, settled } = renderWithConfig(config); + + await settled(); + expect(result.current).toBe(enabled); + }); + + it("keeps stdio available while the proxy has not answered yet", async () => { + const fetchSpy = vi.spyOn(globalThis, "fetch").mockImplementation(() => new Promise(() => {})); + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + const { result } = renderHook(() => useMcpStdioEnabled(), { + wrapper: ({ children }: { children: React.ReactNode }) => ( + {children} + ), + }); + + await waitFor(() => expect(fetchSpy).toHaveBeenCalled()); + expect(result.current).toBe(true); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/StdioAvailability.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/StdioAvailability.tsx new file mode 100644 index 00000000000..960f0dbbe5d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/StdioAvailability.tsx @@ -0,0 +1,37 @@ +import { type FC } from "react"; +import { TriangleAlert } from "lucide-react"; +import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert"; +import { SelectItem } from "@/components/ui/select"; +import { SimpleTooltip } from "@/components/ui/tooltip"; +import { TRANSPORT, TRANSPORT_ITEMS } from "@/components/mcp_tools/types"; +import { $api } from "@/lib/http/api"; + +export const STDIO_DISABLED_MESSAGE = + "stdio MCP servers are disabled on this proxy. Set LITELLM_ENABLE_MCP_STDIO=true on the proxy and restart to enable them"; + +export const useMcpStdioEnabled = (): boolean => + $api.useQuery("get", "/.well-known/litellm-ui-config").data?.mcp_stdio_enabled !== false; + +export const TransportSelectItems: FC<{ stdioEnabled: boolean }> = ({ stdioEnabled }) => ( + <> + {TRANSPORT_ITEMS.map((item) => { + const disabled = item.value === TRANSPORT.STDIO && !stdioEnabled; + return ( + + {item.label} + {disabled && } + + ); + })} + +); + +export const StdioDisabledBanner: FC = () => ( + + + stdio is disabled on this proxy + + {STDIO_DISABLED_MESSAGE}. Until then this server cannot start or be saved as stdio. + + +); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx index 5aec78ba926..f79dd178571 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx @@ -242,6 +242,56 @@ describe("MCPServerEdit (stdio)", () => { }); }); +describe("MCPServerEdit (stdio disabled on the proxy)", () => { + const stdioServer = { + server_id: "server-1", + server_name: "TestServer", + alias: "test", + transport: "stdio", + url: null, + auth_type: "none", + command: "npx", + args: ["-y", "@circleci/mcp-server-circleci"], + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + mcp_access_groups: [], + }; + const renderEdit = (mcpServer: object, stdioEnabled: boolean) => + render( + ["mcpServer"]} + accessToken={null} + onCancel={vi.fn()} + onSuccess={vi.fn()} + availableAccessGroups={[]} + stdioEnabled={stdioEnabled} + />, + ); + + it("explains why an existing stdio server cannot run or be saved as stdio", () => { + renderEdit(stdioServer, false); + + expect(screen.getByText("stdio is disabled on this proxy")).toBeInTheDocument(); + expect(screen.getByText(/Set LITELLM_ENABLE_MCP_STDIO=true on the proxy and restart/)).toBeInTheDocument(); + }); + + it("shows no banner for a stdio server once stdio is enabled", () => { + renderEdit(stdioServer, true); + + expect(screen.getByLabelText("Command")).toBeInTheDocument(); + expect(screen.queryByText("stdio is disabled on this proxy")).not.toBeInTheDocument(); + }); + + it("shows no banner for a non-stdio server while stdio is disabled", () => { + renderEdit({ ...stdioServer, transport: "http", url: "https://mcp.example.com/mcp" }, false); + + expect(screen.getByRole("tab", { name: "Server Configuration" })).toBeInTheDocument(); + expect(screen.queryByText("stdio is disabled on this proxy")).not.toBeInTheDocument(); + }); +}); + describe("MCPServerEdit (delegate auth)", () => { beforeEach(() => { vi.clearAllMocks(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx index d45909445e9..3894cb6cd0b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx @@ -44,6 +44,7 @@ import TruePassthroughWarning from "./TruePassthroughWarning"; import PassthroughAuthorizeSection from "./PassthroughAuthorizeSection"; import MCPToolConfiguration from "./mcp_tool_configuration"; import StdioConfiguration from "./StdioConfiguration"; +import { StdioDisabledBanner, TransportSelectItems } from "./StdioAvailability"; import TokenExchangeFormFields from "./TokenExchangeFormFields"; import IdJagFormFields from "./IdJagFormFields"; import OAuthFormFields from "./OAuthFormFields"; @@ -90,6 +91,7 @@ interface MCPServerEditProps { onSuccess: (server: MCPServer) => void; availableAccessGroups: string[]; existingServers?: MCPServer[]; + stdioEnabled?: boolean; } const AUTH_TYPES_REQUIRING_AUTH_VALUE = [AUTH_TYPE.API_KEY, AUTH_TYPE.BEARER_TOKEN, AUTH_TYPE.TOKEN, AUTH_TYPE.BASIC]; @@ -103,6 +105,7 @@ const MCPServerEdit: React.FC = ({ onSuccess, availableAccessGroups, existingServers, + stdioEnabled = true, }) => { const initialStaticHeaders = React.useMemo(() => { if (!mcpServer.static_headers) { @@ -822,6 +825,7 @@ const MCPServerEdit: React.FC = ({ void submitForm(); }} > + {isStdioTransport && !stdioEnabled && } = ({ - {TRANSPORT_ITEMS.map((item) => ( - - {item.label} - - ))} + )} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.test.tsx index 7e1eda9143f..784bd61382b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.test.tsx @@ -222,6 +222,31 @@ describe("MCPServerView", () => { expect(screen.queryByText("edit form")).not.toBeInTheDocument(); }); + it.each([ + { transport: "stdio", stdioEnabled: false, shown: true }, + { transport: "stdio", stdioEnabled: true, shown: false }, + { transport: "http", stdioEnabled: false, shown: false }, + ])( + "explains why a $transport server is inert when stdioEnabled=$stdioEnabled", + ({ transport, stdioEnabled, shown }) => { + renderView({ transport }, { stdioEnabled }); + + expect(screen.getByText("srv-1")).toBeInTheDocument(); + expect(screen.queryByText("stdio is disabled on this proxy") !== null).toBe(shown); + }, + ); + + it("leaves the stdio warning to the edit form once editing starts", async () => { + renderView({ transport: "stdio" }, { stdioEnabled: false }); + await userEvent.click(screen.getByRole("tab", { name: "Settings" })); + expect(screen.getByText("stdio is disabled on this proxy")).toBeInTheDocument(); + + await userEvent.click(screen.getByRole("button", { name: "Edit Settings" })); + + expect(screen.getByText("edit form")).toBeInTheDocument(); + expect(screen.queryByText("stdio is disabled on this proxy")).not.toBeInTheDocument(); + }); + it("opens on the tab named by initialTabIndex", async () => { renderView({}, { initialTabIndex: 1 }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx index 278a98a2fa6..02d961a818d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx @@ -7,7 +7,7 @@ import { Button } from "@/components/ui/button"; import { Card } from "@/components/ui/card"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; -import { MCPServer, handleTransport, handleAuth } from "@/components/mcp_tools/types"; +import { MCPServer, TRANSPORT, handleTransport, handleAuth } from "@/components/mcp_tools/types"; // TODO: Move Tools viewer from index file import { MCPToolsViewer } from "."; import MCPServerEdit, { EDIT_OAUTH_UI_STATE_KEY } from "./mcp_server_edit"; @@ -15,6 +15,7 @@ import { MCPServerUserCredentialsPanel } from "./MCPServerUserCredentialsPanel"; import { getSecureItem } from "@/utils/secureStorage"; import { isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles"; import MCPServerCostDisplay from "./mcp_server_cost_display"; +import { StdioDisabledBanner } from "./StdioAvailability"; import { getMaskedAndFullUrl, getMCPNetworkAccess } from "./utils"; import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils"; import { CheckIcon, CopyIcon } from "lucide-react"; @@ -31,6 +32,7 @@ interface MCPServerViewProps { availableAccessGroups: string[]; existingServers?: MCPServer[]; initialTabIndex?: number; + stdioEnabled?: boolean; } // True when this render is the return from the edit-settings OAuth redirect for this @@ -63,6 +65,7 @@ export const MCPServerView: React.FC = ({ availableAccessGroups, existingServers, initialTabIndex = 0, + stdioEnabled = true, }) => { // Open the editing Settings tab on first render when returning from the edit OAuth // redirect, so the "token fetched" feedback shows where the user left off (Settings=2). @@ -74,6 +77,8 @@ export const MCPServerView: React.FC = ({ const networkAccess = getMCPNetworkAccess(mcpServer); const [copiedStates, setCopiedStates] = useState>({}); const [selectedTabIndex, setSelectedTabIndex] = useState(returningFromEditOAuth ? 2 : initialTabIndex); + const editFormShowsStdioBanner = selectedTabIndex === 2 && editing && canEdit; + const showStdioBanner = mcpServer.transport === TRANSPORT.STDIO && !stdioEnabled && !editFormShowsStdioBanner; const canViewUserCredentials = userRole !== null && isProxyAdminTierRole(userRole); const canRevokeUserCredentials = userRole !== null && isProxyAdminRole(userRole) && !isViewOnly; @@ -144,6 +149,8 @@ export const MCPServerView: React.FC = ({ {mcpServer.description &&

{mcpServer.description}

}
+ {showStdioBanner && } + setSelectedTabIndex(Number(v))}> @@ -253,6 +260,7 @@ export const MCPServerView: React.FC = ({ onSuccess={handleSuccess} availableAccessGroups={availableAccessGroups} existingServers={existingServers} + stdioEnabled={stdioEnabled} /> ) : (
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx index 9a3ff0cc6cb..2f4fcfd5f69 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx @@ -20,8 +20,12 @@ vi.mock("@/components/networking", () => ({ listMCPUserEnvVarStatus: vi.fn().mockResolvedValue([]), fetchMCPGatewaySessions: vi.fn(), terminateMCPGatewaySessions: vi.fn(), + getUiConfig: vi.fn().mockResolvedValue({}), })); +const stubUiConfig = (config: object) => + vi.spyOn(globalThis, "fetch").mockImplementation(async () => new Response(JSON.stringify(config), { status: 200 })); + const createQueryClient = () => new QueryClient({ defaultOptions: { @@ -135,6 +139,7 @@ describe("MCPServers", () => { beforeEach(() => { vi.clearAllMocks(); + stubUiConfig({}); }); it("should render the MCPServers component with title", async () => { @@ -299,6 +304,90 @@ describe("MCPServers", () => { expect(screen.getByText("No servers match the current filters or search.")).toBeVisible(); }); + it.each([ + [false, true], + [true, false], + ])("marks stdio servers as disabled only when the proxy reports stdio off (enabled=%s)", async (enabled, flagged) => { + stubUiConfig({ mcp_stdio_enabled: enabled }); + vi.mocked(networking.fetchMCPServers).mockResolvedValue([ + { + server_id: "stdio-1", + server_name: "local_tools", + alias: "local_tools", + transport: "stdio", + command: "python", + args: ["server.py"], + created_by: "user", + updated_by: "user", + } as MCPServer, + ]); + vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue([]); + + render( + + + , + ); + + const grid = await screen.findByTestId("mcp-servers-grid"); + await waitFor(() => expect(globalThis.fetch).toHaveBeenCalled()); + await waitFor(() => expect(within(grid).queryByText("stdio disabled") !== null).toBe(flagged)); + expect(within(grid).getByText("STDIO")).toBeInTheDocument(); + }); + + const stdioServer = { + server_id: "stdio-1", + server_name: "local_tools", + alias: "local_tools", + transport: "stdio", + command: "python", + args: ["server.py"], + created_by: "user", + updated_by: "user", + } as MCPServer; + + it("greys out stdio in the create form when the proxy reports stdio off", async () => { + stubUiConfig({ mcp_stdio_enabled: false }); + vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); + const user = userEvent.setup(); + render( + + + , + ); + await waitFor(() => expect(globalThis.fetch).toHaveBeenCalled()); + + await user.click(await screen.findByRole("button", { name: "+ Submit MCP Server" })); + await user.click(await screen.findByRole("combobox", { name: /Transport Type/ })); + + expect(await screen.findByRole("option", { name: /Standard Input\/Output \(stdio\)/ })).toHaveAttribute( + "data-disabled", + ); + expect(screen.getByRole("option", { name: /Streamable HTTP/ })).not.toHaveAttribute("data-disabled"); + + await user.keyboard("{Escape}"); + await waitFor(() => expect(screen.queryByRole("listbox")).not.toBeInTheDocument()); + }); + + it("explains on the edit page why an existing stdio server cannot run when the proxy reports stdio off", async () => { + stubUiConfig({ mcp_stdio_enabled: false }); + vi.mocked(networking.fetchMCPServers).mockResolvedValue([stdioServer]); + vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue([]); + const user = userEvent.setup(); + render( + + + , + ); + const grid = await screen.findByTestId("mcp-servers-grid"); + await waitFor(() => expect(within(grid).getByText("stdio disabled")).toBeInTheDocument()); + + await user.click(within(grid).getAllByText("local_tools")[0]); + await user.click(await screen.findByRole("tab", { name: "Settings" })); + + expect(await screen.findByText("stdio is disabled on this proxy")).toBeInTheDocument(); + }); + it("should render mocked MCP servers data in the table", async () => { // Mock MCP servers data const mockServers = [ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index 56a18e0ca4c..13510cadf39 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -29,6 +29,7 @@ import CreateMCPServer from "./CreateMCPServer"; import ImportMCPServers from "./ImportMCPServers"; import MCPConnect from "./mcp_connect"; import MCPServerCard from "./MCPServerCard"; +import { useMcpStdioEnabled } from "./StdioAvailability"; import { MCPServerView } from "./mcp_server_view"; import type { DiscoverableMCPServer, @@ -225,6 +226,8 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i const [sortKey, setSortKey] = useState("created_desc"); const isInternalUser = userRole === "Internal User"; + const stdioEnabled = useMcpStdioEnabled(); + // Single bulk fetch of this user's per-server env-var status. Drives the // red "N user fields missing" footer on each card with no per-row request. const { data: envVarStatuses, refetch: refetchEnvVarStatus } = useQuery({ @@ -500,6 +503,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i availableAccessGroups={uniqueMcpAccessGroups} existingServers={mcpServers} prefillData={prefillData} + stdioEnabled={stdioEnabled} onBackToDiscovery={() => { setModalVisible(false); setPrefillData(null); @@ -614,6 +618,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i availableAccessGroups={uniqueMcpAccessGroups} existingServers={mcpServers} initialTabIndex={selectedServerId === toolsTabServerId ? 1 : 0} + stdioEnabled={stdioEnabled} /> ) : (
@@ -752,6 +757,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i onByokConnect={server.is_byok ? () => setByokModalServer(server) : undefined} onOpenFillFields={() => setEnvVarsModalServer(server)} onDelete={isAdminRole(userRole) ? () => handleDelete(server.server_id) : undefined} + stdioEnabled={stdioEnabled} /> ))}
diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 20037c9b12b..ba6e6609457 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -382,6 +382,7 @@ export interface LiteLLMWellKnownUiConfig { hide_default_credentials_hint?: boolean; is_control_plane?: boolean; workers?: WorkerInfo[]; + mcp_stdio_enabled?: boolean; } export interface CredentialsResponse { diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index b72f2503e5d..43098982a5e 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -47086,6 +47086,11 @@ export interface components { * @default false */ is_control_plane: boolean; + /** + * Mcp Stdio Enabled + * @default false + */ + mcp_stdio_enabled: boolean; /** Proxy Base Url */ proxy_base_url: string | null; /** Server Root Path */ From 2c9a9711712895a7e08ba27fca357291f4b86290 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 10:32:04 -0700 Subject: [PATCH 010/139] fix(proxy): exit when DATABASE_URL is set but the Prisma toolchain is missing (#44207) With DATABASE_URL set and no way to run the Prisma CLI (not on PATH, not importable), run_server used to print a plain notice and keep booting. The server then crashed later inside the DB exception handler with an unrelated ModuleNotFoundError traceback, and the migration-only entrypoint (--skip_server_startup) exited 0 without migrating. It now exits 1 with a red one-line message naming the missing toolchain and how to install it. The no-DATABASE_URL path is unchanged. Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/proxy_cli.py | 6 ++-- tests/unit/proxy/test_proxy_cli.py | 57 ++++++++++++++++++++++++++++++ 2 files changed, 61 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index cb65cae79a4..a9745565799 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -1447,9 +1447,11 @@ def run_server( sys.exit(1) else: print( - "Unable to connect to DB. DATABASE_URL found in environment, but the prisma CLI is neither on " - "PATH nor importable as a package." + "\033[1;31mLiteLLM Proxy: a database URL is set but the prisma CLI is neither on PATH nor importable " + "as a package, so the database cannot be set up. Install it with `pip install 'litellm[extra_proxy]'` " + "or run a shipped LiteLLM image.\033[0m" ) + sys.exit(1) pgbouncer_settings: Final = PgBouncerSettings() upstream_database_url: Final = os.getenv("DATABASE_URL") if pgbouncer_settings.enabled and upstream_database_url is not None: diff --git a/tests/unit/proxy/test_proxy_cli.py b/tests/unit/proxy/test_proxy_cli.py index 56980b141b6..47071827f2c 100644 --- a/tests/unit/proxy/test_proxy_cli.py +++ b/tests/unit/proxy/test_proxy_cli.py @@ -2265,6 +2265,63 @@ class TestRunServerDbSetup: assert "prisma CLI is neither on PATH" not in capsys.readouterr().out mock_setup_database.assert_called_once_with(use_migrate=True, use_v2_resolver=True) + @pytest.mark.parametrize( + ("database_url", "exits"), + (("postgresql://test:test@localhost:5432/test", True), (None, False)), + ids=("database-url-set", "no-database-url"), + ) + @patch("atexit.register") + def test_startup_exits_when_the_prisma_toolchain_is_missing_only_if_a_database_is_configured( + self, + mock_atexit_register, + database_url, + exits, + tmp_path, + capsys, + ): + """A DATABASE_URL with no way to run the Prisma CLI is fatal; no DATABASE_URL needs no Prisma at all.""" + from litellm_proxy_extras import prisma_toolchain + + from litellm.proxy.proxy_cli import run_server + + empty_bin = tmp_path / "emptybin" + empty_bin.mkdir() + real_find_spec = prisma_toolchain.importlib.util.find_spec + + def hide_prisma(name, package=None): + return None if name == "prisma" else real_find_spec(name, package) + + mock_proxy_module = MagicMock( + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), + ) + clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")} + clean_env["PATH"] = str(empty_bin) + if database_url is not None: + clean_env["DATABASE_URL"] = database_url + + with ( + patch.dict(os.environ, clean_env, clear=True), + patch.dict( + "sys.modules", + {"proxy_server": mock_proxy_module, "litellm.proxy.proxy_server": mock_proxy_module}, + ), + patch.object(prisma_toolchain.importlib.util, "find_spec", side_effect=hide_prisma), + pytest.raises(SystemExit) if exits else nullcontext() as exit_info, + ): + run_server.main(["--local", "--skip_server_startup"], standalone_mode=False) + + out = capsys.readouterr().out + if exits: + assert exit_info.value.code == 1 + assert "a database URL is set but the prisma CLI is neither on PATH nor importable" in out + assert "pip install 'litellm[extra_proxy]'" in out + else: + assert "prisma CLI" not in out + assert "Setup complete" in out + @patch("subprocess.run") @patch("atexit.register") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") From 3ae491a06ccbbaf83b7baf7a3882f1f0ce366229 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Fri, 2 Oct 2026 10:56:52 -0700 Subject: [PATCH 011/139] fix(proxy): preserve state through composed lifespans (#44214) --- backend/main.py | 16 ++- gateway/main.py | 16 ++- litellm/proxy/_lazy_features.py | 6 +- tests/unit/proxy/test_component_allowlists.py | 117 +++++++++++++++++- 4 files changed, 138 insertions(+), 17 deletions(-) diff --git a/backend/main.py b/backend/main.py index 292ece48e7d..e0cef90c979 100644 --- a/backend/main.py +++ b/backend/main.py @@ -8,9 +8,13 @@ Run with: uvicorn backend.main:app --host 0.0.0.0 --port 4001 """ +from collections.abc import AsyncGenerator, Mapping from contextlib import asynccontextmanager +from typing import Final -from fastapi.routing import Mount +from starlette.applications import Starlette +from starlette.routing import Mount +from starlette.types import Lifespan # See gateway/main.py for why we assemble DATABASE_URL(s) here before # importing proxy_server. @@ -43,14 +47,16 @@ def _is_backend_route(route) -> bool: # See gateway/main.py for why the trim runs inside the lifespan instead of at # module scope. -_proxy_lifespan = app.router.lifespan_context +_proxy_lifespan: Final = app.router.lifespan_context @asynccontextmanager -async def _backend_lifespan(app_): - async with _proxy_lifespan(app_): +async def _backend_lifespan( + app_: Starlette, lifespan: Lifespan[Starlette] = _proxy_lifespan +) -> AsyncGenerator[Mapping[str, object], None]: + async with lifespan(app_) as state: app_.router.routes = [r for r in app_.router.routes if _is_backend_route(r)] - yield + yield state if state is not None else {} app.router.lifespan_context = _backend_lifespan diff --git a/gateway/main.py b/gateway/main.py index 61b885b27e4..fb4ae830808 100644 --- a/gateway/main.py +++ b/gateway/main.py @@ -9,9 +9,13 @@ Run with: uvicorn gateway.main:app --host 0.0.0.0 --port 4000 """ +from collections.abc import AsyncGenerator, Mapping from contextlib import asynccontextmanager +from typing import Final -from fastapi.routing import Mount +from starlette.applications import Starlette +from starlette.routing import Mount +from starlette.types import Lifespan # Assemble DATABASE_URL (+ DATABASE_URL_READ_REPLICA) from the discrete # DATABASE_* env vars before proxy_server imports spin up Prisma. Handles @@ -54,14 +58,16 @@ def _is_gateway_route(route) -> bool: # register routes. A module-load filter would miss routes added during # startup; running inside the lifespan, after the inner __aenter__, catches # them while still completing before uvicorn opens the listener. -_proxy_lifespan = app.router.lifespan_context +_proxy_lifespan: Final = app.router.lifespan_context @asynccontextmanager -async def _gateway_lifespan(app_): - async with _proxy_lifespan(app_): +async def _gateway_lifespan( + app_: Starlette, lifespan: Lifespan[Starlette] = _proxy_lifespan +) -> AsyncGenerator[Mapping[str, object], None]: + async with lifespan(app_) as state: app_.router.routes = [r for r in app_.router.routes if _is_gateway_route(r)] - yield + yield state if state is not None else {} app.router.lifespan_context = _gateway_lifespan diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 0b470cf7bda..0119ce4430e 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -530,11 +530,11 @@ def _register_all_on_startup(inner: "Lifespan[FastAPI]", features: tuple[LazyFea (config pass-through endpoints), so the table is put back in lazy mode's order once it is up.""" @asynccontextmanager - async def lifespan(app: "FastAPI") -> AsyncGenerator[None]: + async def lifespan(app: "FastAPI") -> AsyncGenerator[Mapping[str, object]]: register_all_features(app, features) - async with inner(app): + async with inner(app) as state: _restore_registry_order(app, features) - yield + yield state if state is not None else {} return lifespan diff --git a/tests/unit/proxy/test_component_allowlists.py b/tests/unit/proxy/test_component_allowlists.py index 3641a2d9be9..1a210fcb445 100644 --- a/tests/unit/proxy/test_component_allowlists.py +++ b/tests/unit/proxy/test_component_allowlists.py @@ -26,7 +26,18 @@ RDS IAM token when ``IAM_TOKEN_DB_AUTH`` is set). import json import os import sys -from typing import Final +from collections.abc import AsyncGenerator, Mapping +from contextlib import asynccontextmanager +from functools import partial +from typing import Final, Literal + +import pytest +from starlette.applications import Starlette +from starlette.requests import Request +from starlette.responses import JSONResponse +from starlette.routing import Mount, Route +from starlette.testclient import TestClient +from starlette.types import Lifespan # Importing ``litellm.proxy.proxy_server`` runs its module-level setup, which # reads ``DATABASE_URL`` (Prisma) and ``LITELLM_MASTER_KEY``. Tier-zero CI @@ -43,7 +54,6 @@ _PRE_EXISTING_ENV = {key: os.environ.get(key) for key in _THROWAWAY_ENV} for _key, _value in _THROWAWAY_ENV.items(): os.environ.setdefault(_key, _value) -from fastapi.routing import Mount from prometheus_client import make_asgi_app # gateway/ and backend/ live at the repo root, not inside litellm/. @@ -53,6 +63,7 @@ if _REPO_ROOT not in sys.path: from backend.routes.allowlist import BACKEND_MOUNT_PATHS from gateway.routes.allowlist import GATEWAY_MOUNT_PATHS +from litellm.proxy._lazy_features import LazyFeature, attach_lazy_features from litellm.proxy.proxy_server import app from tests.test_litellm_rust.support.child_interpreter import run_child_interpreter @@ -74,7 +85,10 @@ _DB_ENV_KEYS = ( ) _PRE_DB_ENV = {_key: os.environ.pop(_key, None) for _key in _DB_ENV_KEYS} _PRE_COMPONENT_LIFESPAN = app.router.lifespan_context -from gateway.main import _is_gateway_route +from gateway.main import _gateway_lifespan, _is_gateway_route + +app.router.lifespan_context = _PRE_COMPONENT_LIFESPAN +from backend.main import _backend_lifespan app.router.lifespan_context = _PRE_COMPONENT_LIFESPAN for _key, _previous in _PRE_DB_ENV.items(): @@ -85,7 +99,7 @@ for _key, _previous in _PRE_DB_ENV.items(): _COVERAGE_PROBE: Final = """ import json, os, sys sys.path.insert(0, os.environ["LITELLM_COMPONENT_ALLOWLIST_REPO_ROOT"]) -from fastapi.routing import Mount +from starlette.routing import Mount from backend.routes.allowlist import BACKEND_EXACT_PATHS, BACKEND_PATH_PREFIXES from gateway.routes.allowlist import GATEWAY_EXACT_PATHS, GATEWAY_PATH_PREFIXES from litellm.proxy._lazy_features import loaded_lazy_modules @@ -112,6 +126,101 @@ json.dump({ """ +@pytest.mark.parametrize( + "component_lifespan", (None, _gateway_lifespan, _backend_lifespan), ids=("proxy", "gateway", "backend") +) +@pytest.mark.parametrize("eager", (False, True), ids=("lazy", "eager")) +@pytest.mark.parametrize("state_kind", ("enabled", "disabled", "stateless")) +def test_composed_lifespan_preserves_request_state_and_teardown( + monkeypatch: pytest.MonkeyPatch, + component_lifespan: Lifespan[Starlette] | None, + eager: bool, + state_kind: Literal["enabled", "disabled", "stateless"], +) -> None: + monkeypatch.setenv("LITELLM_DISABLE_LAZY_ROUTES", str(eager).lower()) + receiver: Final = object() + resource: Final = object() + state: Final[Mapping[str, object]] = { + "tracing_receiver": receiver if state_kind == "enabled" else None, + "other_resource": resource, + } + events: Final[list[str]] = [] # mutable-ok: observe startup, requests and teardown across the ASGI boundary + + async def trace_state(request: Request) -> JSONResponse: + events.append("request") + assert events[0] == "startup" and "shutdown" not in events + assert getattr(request.state, "other_resource", None) is (resource if state_kind != "stateless" else None) + assert getattr(request.state, "tracing_receiver", None) is (receiver if state_kind == "enabled" else None) + return JSONResponse({"keys": sorted(request.scope["state"])}) + + def register_trace_route(application: Starlette, module: object) -> None: + application.router.routes.append(Route("/v1/traces", trace_state)) + + @asynccontextmanager + async def stateful_lifespan(application: Starlette) -> AsyncGenerator[Mapping[str, object], None]: + events.append("startup") + application.router.routes.append(Route("/not-a-component-route", trace_state)) + try: + yield state + finally: + events.append("shutdown") + + @asynccontextmanager + async def stateless_lifespan(application: Starlette) -> AsyncGenerator[None, None]: + async with stateful_lifespan(application): + yield + + application: Final = type(app)(lifespan=stateless_lifespan if state_kind == "stateless" else stateful_lifespan) + feature: Final = LazyFeature("traces", __name__, ("/v1/traces",), register_fn=register_trace_route) + attach_lazy_features(application, (feature,)) + if component_lifespan is not None: + application.router.lifespan_context = partial(component_lifespan, lifespan=application.router.lifespan_context) + + with TestClient(application) as client: + response: Final = client.get("/v1/traces") + assert response.status_code == 200, response.text + assert response.json() == {"keys": [] if state_kind == "stateless" else sorted(state)} + filtered: Final = client.get("/not-a-component-route") + assert filtered.status_code == (200 if component_lifespan is None else 404), filtered.text + assert events == (["startup", "request", "request"] if component_lifespan is None else ["startup", "request"]) + assert events == ( + ["startup", "request", "request", "shutdown"] if component_lifespan is None else ["startup", "request", "shutdown"] + ) + + +@pytest.mark.parametrize( + "component_lifespan", (None, _gateway_lifespan, _backend_lifespan), ids=("proxy", "gateway", "backend") +) +@pytest.mark.parametrize("eager", (False, True), ids=("lazy", "eager")) +@pytest.mark.parametrize("phase", ("startup", "shutdown")) +def test_composed_lifespan_propagates_lifecycle_failures( + monkeypatch: pytest.MonkeyPatch, component_lifespan: Lifespan[Starlette] | None, eager: bool, phase: str +) -> None: + monkeypatch.setenv("LITELLM_DISABLE_LAZY_ROUTES", str(eager).lower()) + failure: Final = RuntimeError(f"{phase} failed") + events: Final[list[str]] = [] # mutable-ok: observe lifecycle events across the ASGI boundary + + @asynccontextmanager + async def inner_lifespan(application: Starlette) -> AsyncGenerator[Mapping[str, object], None]: + events.append("startup") + if phase == "startup": + raise failure + yield {} + events.append("shutdown") + raise failure + + application: Final = type(app)(lifespan=inner_lifespan) + attach_lazy_features(application, ()) + if component_lifespan is not None: + application.router.lifespan_context = partial(component_lifespan, lifespan=application.router.lifespan_context) + + with pytest.raises(RuntimeError) as caught: + with TestClient(application): + events.append("serving") + assert caught.value is failure + assert events == (["startup"] if phase == "startup" else ["startup", "serving", "shutdown"]) + + def test_gateway_plus_backend_covers_full_app(): """Every route on the proxy app must be served by gateway or backend. From 56bba4fbbe98a6db7b07e2b8c431b3b4250d3155 Mon Sep 17 00:00:00 2001 From: Fede Kamelhar <209537060+fede-kamel@users.noreply.github.com> Date: Fri, 2 Oct 2026 13:57:12 -0400 Subject: [PATCH 012/139] fix(oci): resolve the GenAI endpoint realm from the compartment OCID instead of hardcoding oraclecloud.com (#43180) * fix(oci): resolve the GenAI endpoint realm from the region instead of hardcoding oraclecloud.com Government (OC2/OC3/OC4) and other non-commercial realms live under a different second-level domain, so a request for us-luke-1 was sent to inference.generativeai.us-luke-1.oci.oraclecloud.com and failed DNS Delegate the lookup to the OCI SDK's region registry when it is installed, honour OCI_DEFAULT_REALM otherwise, and keep api_base as the explicit override * fix(oci): read per-region realm metadata without the SDK and cover the registry path Replace the global OCI_DEFAULT_REALM fallback, which would have redirected commercial regions too in a mixed deployment, with the SDK's own per-region sources: OCI_REGION_METADATA and ~/.oci/regions-config.json. Regions not described anywhere keep their commercial endpoint Exercise the SDK registry path with a fake oci.regions module so CI, which has no SDK, still covers it, and skip the real-SDK test on oci.regions so a namespace package named oci in the tests tree cannot masquerade as the SDK * fix(oci): validate regions-config.json entries individually and tolerate undecodable files One malformed entry no longer discards the valid ones, and the file is parsed from bytes so an undecodable file is logged and ignored instead of failing every OCI request built without the SDK * fix(oci): resolve the realm from the compartment OCID so Government regions work without the SDK The Docker image ships without the oci package and Government deployments rarely carry OCI_REGION_METADATA, so the reviewed fallback still sent us-luke-1 to oraclecloud.com. Every compartment OCID already names its realm (ocid1.compartment.oc2..), so map that key through the SDK's twenty realm domains first, then the metadata sources, then the SDK registry, then the commercial default. * fix(oci): consult the SDK registry before hand-parsed region metadata and harden the fallback Review follow-ups on the realm resolver. Read the compartment realm first, then the SDK registry when it is installed, and only then the hand-parsed metadata sources, so the same file is never parsed twice with different rules. Lowercase metadata values like the SDK does, accept single-label realm domains, and expand ~ with os.path so a container without a home directory cannot raise out of URL building. Type the compartment as str | None at the caller, drop populate_by_name, isolate the legacy region tests from the developer's ~/.oci, and keep the new tests on the immutable style. --- litellm/llms/oci/common_utils.py | 147 +++++++++- litellm/llms/oci/embed/transformation.py | 6 +- tests/unit/llms/oci/test_oci_common_utils.py | 271 +++++++++++++++++-- 3 files changed, 399 insertions(+), 25 deletions(-) diff --git a/litellm/llms/oci/common_utils.py b/litellm/llms/oci/common_utils.py index 3f703564b5a..679a9c21e43 100644 --- a/litellm/llms/oci/common_utils.py +++ b/litellm/llms/oci/common_utils.py @@ -1,16 +1,21 @@ import base64 import hashlib +import importlib import json import os import re +from collections.abc import Mapping from dataclasses import dataclass from email.utils import formatdate -from typing import Final, Protocol +from pathlib import Path +from types import MappingProxyType +from typing import Final, Protocol, runtime_checkable from urllib.parse import urlparse import httpx -from pydantic import JsonValue +from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter, ValidationError, field_validator +from litellm._logging import verbose_logger from litellm.llms.base_llm.chat.transformation import BaseLLMException try: @@ -154,7 +159,7 @@ _OCI_KEY_ENV: Final = "OCI_KEY" _OCI_COMPARTMENT_ID_ENV: Final = "OCI_COMPARTMENT_ID" -def resolve_oci_credentials(optional_params: dict) -> dict: +def resolve_oci_credentials(optional_params: Mapping[str, object]) -> dict: """ Merge OCI credentials from optional_params (explicit, always wins) and environment variables (fallback). @@ -174,11 +179,140 @@ def resolve_oci_credentials(optional_params: dict) -> dict: } -_OCI_REGION_RE: Final = re.compile(r"^[a-z][a-z0-9-]{0,30}[a-z0-9]$") +_OCI_REGION_PATTERN: Final = r"^[a-z][a-z0-9-]{0,30}[a-z0-9]$" +_OCI_REALM_DOMAIN_PATTERN: Final = r"^[a-z0-9]([a-z0-9-]*[a-z0-9])?(\.[a-z0-9]([a-z0-9-]*[a-z0-9])?)*$" +_OCI_REGION_RE: Final = re.compile(_OCI_REGION_PATTERN) _OCI_ACTION_PATH_RE: Final = re.compile(rf"/{OCI_API_VERSION}/actions/[^/?#]+/?$") +_OCI_COMMERCIAL_REALM_DOMAIN: Final = "oraclecloud.com" +_OCI_INFERENCE_ENDPOINT_TEMPLATE: Final = "https://inference.generativeai.{region}.oci.{secondLevelDomain}" +_OCI_REGION_METADATA_ENV: Final = "OCI_REGION_METADATA" +_OCI_REGIONS_CONFIG_FILE: Final = "~/.oci/regions-config.json" +_OCID_REALM_RE: Final = re.compile(r"^ocid1\.[a-z0-9]+\.([a-z0-9]+)\.", re.IGNORECASE) +_OCI_REALM_DOMAINS: Final = MappingProxyType( + { + "oc1": "oraclecloud.com", + "oc2": "oraclegovcloud.com", + "oc3": "oraclegovcloud.com", + "oc4": "oraclegovcloud.uk", + "oc8": "oraclecloud8.com", + "oc9": "oraclecloud9.com", + "oc10": "oraclecloud10.com", + "oc14": "oraclecloud14.com", + "oc15": "oraclecloud15.com", + "oc19": "oraclecloud.eu", + "oc20": "oraclecloud20.com", + "oc21": "oraclecloud21.com", + "oc23": "oraclecloud23.com", + "oc24": "oraclecloud24.com", + "oc26": "oraclecloud26.com", + "oc29": "oraclecloud29.com", + "oc35": "oraclecloud35.com", + "oc42": "oraclecloud42.com", + "oc51": "oraclecloud51.com", + "oc52": "oraclecloud52.com", + } +) -def get_oci_base_url(optional_params: dict, api_base: str | None = None) -> str: +class OCIRegionMetadata(BaseModel): + """One entry of the OCI SDK's region metadata schema, as found in + ``~/.oci/regions-config.json`` (a JSON array) or ``OCI_REGION_METADATA`` (one object). + Values are lowercased before validation, as the SDK does.""" + + model_config = ConfigDict(frozen=True, extra="ignore") + + region_identifier: str = Field(alias="regionIdentifier", pattern=_OCI_REGION_PATTERN) + realm_domain_component: str = Field(alias="realmDomainComponent", pattern=_OCI_REALM_DOMAIN_PATTERN) + + @field_validator("region_identifier", "realm_domain_component", mode="before") + @classmethod + def _lowercase(cls, value: object) -> object: + return value.lower() if isinstance(value, str) else value + + +_JSON_ARRAY: Final = TypeAdapter(tuple[JsonValue, ...]) + + +def _validated_region_metadata(raw: JsonValue, source: str) -> OCIRegionMetadata | None: + try: + return OCIRegionMetadata.model_validate(raw) + except ValidationError as e: + verbose_logger.warning("Ignoring OCI region metadata entry in %s: %s", source, e) + return None + + +def _region_metadata_from_file() -> tuple[OCIRegionMetadata, ...]: + path: Final = Path(os.path.expanduser(_OCI_REGIONS_CONFIG_FILE)) + if not path.is_file(): + return () + try: + raw_entries: Final = _JSON_ARRAY.validate_json(path.read_bytes()) + except (OSError, ValidationError) as e: + verbose_logger.warning("Ignoring OCI region metadata in %s: %s", path, e) + return () + candidates: Final = (_validated_region_metadata(raw, str(path)) for raw in raw_entries) + return tuple(entry for entry in candidates if entry is not None) + + +def _region_metadata_from_env() -> tuple[OCIRegionMetadata, ...]: + raw: Final = os.environ.get(_OCI_REGION_METADATA_ENV) + if not raw: + return () + try: + return (OCIRegionMetadata.model_validate_json(raw),) + except ValidationError as e: + verbose_logger.warning("Ignoring OCI region metadata in %s: %s", _OCI_REGION_METADATA_ENV, e) + return () + + +def _realm_domain_from_ocid(ocid: str | None) -> str | None: + match: Final = _OCID_REALM_RE.match(ocid) if ocid else None + return _OCI_REALM_DOMAINS.get(match.group(1).lower()) if match else None + + +def _realm_domain_from_metadata(region: str) -> str | None: + entries: Final = (*_region_metadata_from_file(), *_region_metadata_from_env()) + return next((entry.realm_domain_component for entry in entries if entry.region_identifier == region), None) + + +@runtime_checkable +class _OCIRegionRegistry(Protocol): + def endpoint_for(self, service: str, region: str, service_endpoint_template: str) -> str: ... + + +def _load_oci_region_registry() -> _OCIRegionRegistry | None: + try: + registry: Final = importlib.import_module("oci.regions") + except ImportError: + return None + return registry if isinstance(registry, _OCIRegionRegistry) else None + + +def resolve_oci_inference_endpoint(region: str, compartment_id: str | None = None) -> str: + """Return the GenAI inference endpoint for ``region`` in whichever OCI realm hosts it. + + The realm's second-level domain comes first from the realm key inside ``compartment_id`` + (``ocid1.compartment.oc2..`` is the Government realm), then from the OCI SDK's region + registry when the SDK is installed, then from the per-region metadata sources the SDK + reads, ``~/.oci/regions-config.json`` and ``OCI_REGION_METADATA``, and otherwise defaults + to the commercial realm. Realm domains per ``oci/regions_definitions.py`` in oci 2.187.0. + A region that is not described anywhere therefore keeps its commercial endpoint, so one + government deployment never redirects the others. + """ + realm_domain: Final = _realm_domain_from_ocid(compartment_id) + if realm_domain is not None: + return _OCI_INFERENCE_ENDPOINT_TEMPLATE.format(region=region, secondLevelDomain=realm_domain) + registry: Final = _load_oci_region_registry() + if registry is not None: + return registry.endpoint_for( + "generative_ai_inference", region=region, service_endpoint_template=_OCI_INFERENCE_ENDPOINT_TEMPLATE + ) + return _OCI_INFERENCE_ENDPOINT_TEMPLATE.format( + region=region, secondLevelDomain=_realm_domain_from_metadata(region) or _OCI_COMMERCIAL_REALM_DOMAIN + ) + + +def get_oci_base_url(optional_params: Mapping[str, object], api_base: str | None = None) -> str: """Return the OCI inference base URL, respecting any explicit api_base override. If ``api_base`` already ends with a fully-formed OCI action path @@ -196,7 +330,8 @@ def get_oci_base_url(optional_params: dict, api_base: str | None = None) -> str: f"Invalid OCI region {region!r}: must match ^[a-z][a-z0-9-]{{0,30}}[a-z0-9]$ (e.g. 'us-ashburn-1')." ), ) - return f"https://inference.generativeai.{region}.oci.oraclecloud.com" + compartment_id: Final = creds["oci_compartment_id"] + return resolve_oci_inference_endpoint(region, compartment_id if isinstance(compartment_id, str) else None) # --------------------------------------------------------------------------- diff --git a/litellm/llms/oci/embed/transformation.py b/litellm/llms/oci/embed/transformation.py index 2300c6ee403..dec43717387 100644 --- a/litellm/llms/oci/embed/transformation.py +++ b/litellm/llms/oci/embed/transformation.py @@ -77,7 +77,11 @@ class OCIEmbedConfig(BaseEmbeddingConfig): Required call-time params (via optional_params or env vars): - ``oci_compartment_id`` / ``OCI_COMPARTMENT_ID`` - - ``oci_region`` / ``OCI_REGION`` (default: ``us-ashburn-1``) + - ``oci_region`` / ``OCI_REGION`` (default: ``us-ashburn-1``). The realm comes from the realm + key in ``oci_compartment_id`` (``ocid1.compartment.oc2..`` is the Government realm), so + non-commercial realms need no extra setting. A realm unknown to litellm can be described in + ``OCI_REGION_METADATA`` or ``~/.oci/regions-config.json``, resolved through the OCI SDK when + it is installed, or given as ``api_base``. Optional call-time params: - ``oci_serving_mode``: ``"ON_DEMAND"`` (default) or ``"DEDICATED"`` diff --git a/tests/unit/llms/oci/test_oci_common_utils.py b/tests/unit/llms/oci/test_oci_common_utils.py index d306d7351dd..e66645c4dcd 100644 --- a/tests/unit/llms/oci/test_oci_common_utils.py +++ b/tests/unit/llms/oci/test_oci_common_utils.py @@ -5,10 +5,16 @@ Covers schema utilities, signing helpers, and credential resolution paths that require no real OCI credentials or network calls. """ -import pytest +import sys +import types +from types import MappingProxyType +from typing import Final from unittest.mock import MagicMock, patch +import pytest + from litellm.llms.oci.common_utils import ( + _OCI_REALM_DOMAINS, OCI_API_VERSION, OCIError, OCIRequestWrapper, @@ -40,7 +46,8 @@ def test_oci_api_version_constant(): def test_sha256_base64_known_value(): - import base64, hashlib + import base64 + import hashlib data = b"hello" expected = base64.b64encode(hashlib.sha256(data).digest()).decode() @@ -60,9 +67,7 @@ def test_sha256_base64_empty(): def test_build_signature_string_request_target(): headers = {"host": "example.com", "date": "Mon, 01 Jan 2024 00:00:00 GMT"} - result = build_signature_string( - "POST", "/20231130/actions/chat", headers, ["(request-target)", "host", "date"] - ) + result = build_signature_string("POST", "/20231130/actions/chat", headers, ["(request-target)", "host", "date"]) lines = result.split("\n") assert lines[0] == "(request-target): post /20231130/actions/chat" assert lines[1] == "host: example.com" @@ -161,12 +166,10 @@ def test_get_oci_base_url_explicit_api_base(): ], ) def test_get_oci_base_url_strips_trailing_action_path(api_base): - assert ( - get_oci_base_url({}, api_base=api_base) - == "https://inference.generativeai.us-chicago-1.oci.oraclecloud.com" - ) + assert get_oci_base_url({}, api_base=api_base) == "https://inference.generativeai.us-chicago-1.oci.oraclecloud.com" +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") def test_get_oci_base_url_from_region(): url = get_oci_base_url({"oci_region": "eu-frankfurt-1"}) assert url == "https://inference.generativeai.eu-frankfurt-1.oci.oraclecloud.com" @@ -192,6 +195,7 @@ def test_get_oci_base_url_rejects_unsafe_region(region): get_oci_base_url({"oci_region": region}) +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") def test_get_oci_base_url_empty_region_falls_back_to_default(monkeypatch): monkeypatch.delenv("OCI_REGION", raising=False) url = get_oci_base_url({"oci_region": ""}) @@ -209,11 +213,248 @@ def test_get_oci_base_url_empty_region_falls_back_to_default(monkeypatch): "ap", ], ) +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") def test_get_oci_base_url_accepts_valid_region(region): url = get_oci_base_url({"oci_region": region}) assert url == f"https://inference.generativeai.{region}.oci.oraclecloud.com" +_NON_COMMERCIAL_REALMS: Final = ( + ("oc2", "us-luke-1", "oraclegovcloud.com"), + ("oc3", "us-gov-ashburn-1", "oraclegovcloud.com"), + ("oc4", "uk-gov-london-1", "oraclegovcloud.uk"), + ("oc19", "eu-frankfurt-2", "oraclecloud.eu"), +) +_UNKNOWN_REGION: Final = "xx-nowhere-1" +_UNKNOWN_REALM_COMPARTMENT: Final = "ocid1.compartment.oc99..aaaaaaaaexample" +_UNKNOWN_REGION_METADATA: Final = '{"realmKey": "OCX", "realmDomainComponent": "example.test", "regionKey": "XNW", "regionIdentifier": "xx-nowhere-1"}' + + +def _compartment(realm): + return f"ocid1.compartment.{realm}..aaaaaaaaexample" + + +def _params(region: str, compartment_id: object = None) -> MappingProxyType[str, object]: + return MappingProxyType({"oci_region": region, "oci_compartment_id": compartment_id}) + + +@pytest.fixture +def without_oci_sdk(monkeypatch): + monkeypatch.setitem(sys.modules, "oci", None) + monkeypatch.setitem(sys.modules, "oci.regions", None) + + +@pytest.fixture +def isolated_region_metadata(monkeypatch, tmp_path): + monkeypatch.delenv("OCI_REGION_METADATA", raising=False) + monkeypatch.delenv("OCI_COMPARTMENT_ID", raising=False) + monkeypatch.setenv("HOME", str(tmp_path)) + return tmp_path + + +def test_realm_table_matches_installed_sdk(): + # Realm domains per the OCI Python SDK's oci.regions_definitions.REALMS (v2.187.0, checked 2026-09-27) + definitions: Final = pytest.importorskip("oci.regions_definitions") + assert ( + MappingProxyType({realm: definitions.REALMS.get(realm) for realm in _OCI_REALM_DOMAINS}) == _OCI_REALM_DOMAINS + ) + + +@pytest.mark.usefixtures("isolated_region_metadata") +@pytest.mark.parametrize(("realm", "region", "second_level_domain"), _NON_COMMERCIAL_REALMS) +def test_get_oci_base_url_resolves_realm_from_region_via_sdk(realm, region, second_level_domain): + pytest.importorskip("oci.regions") + # Realm domains per the OCI Python SDK's oci.regions_definitions (v2.187.0, checked 2026-09-27) + url: Final = get_oci_base_url(_params(region)) + assert url == f"https://inference.generativeai.{region}.oci.{second_level_domain}" + + +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") +@pytest.mark.parametrize(("realm", "region", "second_level_domain"), _NON_COMMERCIAL_REALMS) +def test_get_oci_base_url_resolves_realm_from_compartment_ocid(realm, region, second_level_domain): + url: Final = get_oci_base_url(_params(region, _compartment(realm))) + assert url == f"https://inference.generativeai.{region}.oci.{second_level_domain}" + + +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") +def test_get_oci_base_url_resolves_realm_from_compartment_env(monkeypatch): + monkeypatch.setenv("OCI_COMPARTMENT_ID", _compartment("oc2")) + url: Final = get_oci_base_url(_params("us-luke-1")) + assert url == "https://inference.generativeai.us-luke-1.oci.oraclegovcloud.com" + + +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") +def test_get_oci_base_url_reads_realm_key_case_insensitively(): + url: Final = get_oci_base_url(_params("us-luke-1", _compartment("OC2"))) + assert url == "https://inference.generativeai.us-luke-1.oci.oraclegovcloud.com" + + +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") +def test_get_oci_base_url_keeps_commercial_compartment_commercial(): + url: Final = get_oci_base_url(_params("us-chicago-1", _compartment("oc1"))) + assert url == "https://inference.generativeai.us-chicago-1.oci.oraclecloud.com" + + +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") +@pytest.mark.parametrize("compartment_id", (None, "not-an-ocid", _UNKNOWN_REALM_COMPARTMENT, 42)) +def test_get_oci_base_url_without_sdk_defaults_to_commercial_when_realm_unknown(compartment_id): + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION, compartment_id)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.oraclecloud.com" + + +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") +def test_get_oci_base_url_compartment_realm_wins_over_region_metadata(monkeypatch): + monkeypatch.setenv( + "OCI_REGION_METADATA", '{"regionIdentifier": "us-luke-1", "realmDomainComponent": "example.test"}' + ) + url: Final = get_oci_base_url(_params("us-luke-1", _compartment("oc2"))) + assert url == "https://inference.generativeai.us-luke-1.oci.oraclegovcloud.com" + + +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") +def test_get_oci_base_url_without_sdk_uses_region_metadata_env(monkeypatch): + monkeypatch.setenv("OCI_REGION_METADATA", _UNKNOWN_REGION_METADATA) + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION, _UNKNOWN_REALM_COMPARTMENT)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.example.test" + + +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") +def test_get_oci_base_url_without_sdk_region_metadata_leaves_other_regions_commercial(monkeypatch): + monkeypatch.setenv("OCI_REGION_METADATA", _UNKNOWN_REGION_METADATA) + url: Final = get_oci_base_url(_params("us-chicago-1")) + assert url == "https://inference.generativeai.us-chicago-1.oci.oraclecloud.com" + + +@pytest.mark.usefixtures("without_oci_sdk") +def test_get_oci_base_url_without_sdk_uses_regions_config_file(isolated_region_metadata): + oci_dir: Final = isolated_region_metadata / ".oci" + oci_dir.mkdir() + (oci_dir / "regions-config.json").write_text(f"[{_UNKNOWN_REGION_METADATA}]") + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.example.test" + + +@pytest.mark.usefixtures("without_oci_sdk") +def test_get_oci_base_url_without_sdk_keeps_valid_regions_config_entries_next_to_a_bad_one(isolated_region_metadata): + oci_dir: Final = isolated_region_metadata / ".oci" + oci_dir.mkdir() + (oci_dir / "regions-config.json").write_text( + f'[{{"regionIdentifier": "us-langley-1"}}, {_UNKNOWN_REGION_METADATA}]' + ) + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.example.test" + + +@pytest.mark.usefixtures("without_oci_sdk") +@pytest.mark.parametrize("content", (b"\xff\xfe\x00[", b'{"regionIdentifier": "xx-nowhere-1"}', b"not json")) +def test_get_oci_base_url_without_sdk_ignores_unusable_regions_config_file(isolated_region_metadata, content): + oci_dir: Final = isolated_region_metadata / ".oci" + oci_dir.mkdir() + (oci_dir / "regions-config.json").write_bytes(content) + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.oraclecloud.com" + + +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") +@pytest.mark.parametrize( + "metadata", + ( + '{"regionIdentifier": "xx-nowhere-1", "realmDomainComponent": "evil.com/#"}', + '{"regionIdentifier": "xx-nowhere-1", "realmDomainComponent": "-internal"}', + '{"regionIdentifier": "xx-nowhere-1"}', + "not json", + ), +) +def test_get_oci_base_url_without_sdk_ignores_invalid_region_metadata(monkeypatch, metadata): + monkeypatch.setenv("OCI_REGION_METADATA", metadata) + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.oraclecloud.com" + + +def _fake_oci_regions(endpoint_for=None): + module: Final = types.ModuleType("oci.regions") + if endpoint_for is not None: + module.endpoint_for = endpoint_for + return module + + +@pytest.mark.usefixtures("isolated_region_metadata") +def test_get_oci_base_url_uses_sdk_region_registry_when_realm_unknown(monkeypatch): + endpoint_for: Final = MagicMock( + side_effect=lambda service, region, service_endpoint_template: service_endpoint_template.format( + region=region, secondLevelDomain="example.test" + ) + ) + monkeypatch.setitem(sys.modules, "oci", types.ModuleType("oci")) + monkeypatch.setitem(sys.modules, "oci.regions", _fake_oci_regions(endpoint_for)) + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION, _UNKNOWN_REALM_COMPARTMENT)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.example.test" + endpoint_for.assert_called_once_with( + "generative_ai_inference", + region=_UNKNOWN_REGION, + service_endpoint_template="https://inference.generativeai.{region}.oci.{secondLevelDomain}", + ) + + +@pytest.mark.usefixtures("isolated_region_metadata") +def test_get_oci_base_url_skips_sdk_region_registry_when_compartment_realm_known(monkeypatch): + def endpoint_for(service, region, service_endpoint_template): + raise AssertionError("registry consulted") + + monkeypatch.setitem(sys.modules, "oci", types.ModuleType("oci")) + monkeypatch.setitem(sys.modules, "oci.regions", _fake_oci_regions(endpoint_for)) + url: Final = get_oci_base_url(_params("us-luke-1", _compartment("oc2"))) + assert url == "https://inference.generativeai.us-luke-1.oci.oraclegovcloud.com" + + +@pytest.mark.usefixtures("isolated_region_metadata") +def test_get_oci_base_url_prefers_sdk_region_registry_over_hand_parsed_metadata(monkeypatch): + def endpoint_for(service, region, service_endpoint_template): + return service_endpoint_template.format(region=region, secondLevelDomain="sdk.test") + + monkeypatch.setitem(sys.modules, "oci", types.ModuleType("oci")) + monkeypatch.setitem(sys.modules, "oci.regions", _fake_oci_regions(endpoint_for)) + monkeypatch.setenv("OCI_REGION_METADATA", _UNKNOWN_REGION_METADATA) + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.sdk.test" + + +@pytest.mark.usefixtures("isolated_region_metadata") +def test_get_oci_base_url_falls_back_to_metadata_when_sdk_registry_lacks_endpoint_for(monkeypatch): + monkeypatch.setitem(sys.modules, "oci", types.ModuleType("oci")) + monkeypatch.setitem(sys.modules, "oci.regions", _fake_oci_regions()) + monkeypatch.setenv("OCI_REGION_METADATA", _UNKNOWN_REGION_METADATA) + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.example.test" + + +@pytest.mark.usefixtures("without_oci_sdk", "isolated_region_metadata") +@pytest.mark.parametrize( + ("metadata", "second_level_domain"), + ( + ('{"regionIdentifier": "XX-NOWHERE-1", "realmDomainComponent": "Example.Test"}', "example.test"), + ('{"regionIdentifier": "xx-nowhere-1", "realmDomainComponent": "internal"}', "internal"), + ), +) +def test_get_oci_base_url_without_sdk_normalizes_region_metadata_like_the_sdk( + monkeypatch, metadata, second_level_domain +): + monkeypatch.setenv("OCI_REGION_METADATA", metadata) + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.{second_level_domain}" + + +@pytest.mark.usefixtures("without_oci_sdk") +def test_get_oci_base_url_without_sdk_tolerates_unresolvable_home(monkeypatch): + def no_passwd_entry(uid): + raise KeyError(uid) + + monkeypatch.delenv("HOME", raising=False) + monkeypatch.setattr("pwd.getpwuid", no_passwd_entry) + url: Final = get_oci_base_url(_params(_UNKNOWN_REGION)) + assert url == f"https://inference.generativeai.{_UNKNOWN_REGION}.oci.oraclecloud.com" + + # --------------------------------------------------------------------------- # validate_oci_environment # --------------------------------------------------------------------------- @@ -247,17 +488,13 @@ def test_sign_with_oci_signer_exception_wrapped(): bad_signer = MagicMock() bad_signer.do_request_sign.side_effect = RuntimeError("signing failed") with pytest.raises(OCIError, match="Failed to sign request"): - sign_with_oci_signer( - {}, {"oci_signer": bad_signer}, {"key": "val"}, "https://example.com" - ) + sign_with_oci_signer({}, {"oci_signer": bad_signer}, {"key": "val"}, "https://example.com") def test_sign_with_oci_signer_success(): signer = MagicMock() signer.do_request_sign.return_value = None - headers, body = sign_with_oci_signer( - {}, {"oci_signer": signer}, {"key": "val"}, "https://example.com" - ) + headers, body = sign_with_oci_signer({}, {"oci_signer": signer}, {"key": "val"}, "https://example.com") assert isinstance(body, bytes) signer.do_request_sign.assert_called_once() @@ -270,9 +507,7 @@ def test_sign_with_oci_signer_success(): def test_sign_oci_request_routes_to_signer(): signer = MagicMock() signer.do_request_sign.return_value = None - headers, body = sign_oci_request( - {}, {"oci_signer": signer}, {}, "https://example.com" - ) + headers, body = sign_oci_request({}, {"oci_signer": signer}, {}, "https://example.com") signer.do_request_sign.assert_called_once() From e5873adbc4ab27492ab55cd21ec09e30fd9ae9b6 Mon Sep 17 00:00:00 2001 From: Steffan W Date: Fri, 2 Oct 2026 19:57:37 +0200 Subject: [PATCH 013/139] fix(chatgpt): preserve requested service tier in Responses calls (#40108) * fix(chatgpt): preserve requested service tier in Responses calls * fix(chatgpt): avoid extra mutable collections in tier filtering * test(chatgpt): refresh service-tier cases after main rebase Keep the tier-preservation regression on the current subscription model and match the moved unit-test suite's formatting. The adapter still forwards preferences without claiming backend priority entitlement. * refactor(chatgpt): map service tier through a lookup table after the allowlist filter --------- Co-authored-by: ryan-crabbe-berri --- .../llms/chatgpt/responses/transformation.py | 8 +- .../test_chatgpt_responses_transformation.py | 78 ++++++++++++------- 2 files changed, 57 insertions(+), 29 deletions(-) diff --git a/litellm/llms/chatgpt/responses/transformation.py b/litellm/llms/chatgpt/responses/transformation.py index 9774b762396..78ee8a86a65 100644 --- a/litellm/llms/chatgpt/responses/transformation.py +++ b/litellm/llms/chatgpt/responses/transformation.py @@ -35,6 +35,8 @@ from ..common_utils import ( if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +_CHATGPT_SERVICE_TIERS: Final = {"default": "default", "priority": "priority", "fast": "priority"} + class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): def __init__(self) -> None: @@ -108,7 +110,11 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): "truncation", } - return {k: v for k, v in request.items() if k in allowed_keys} + filtered: Final = {k: v for k, v in request.items() if k in allowed_keys} + service_tier: Final = _CHATGPT_SERVICE_TIERS.get(request.get("service_tier")) + if service_tier is not None: + filtered["service_tier"] = service_tier + return filtered def transform_response_api_response( self, diff --git a/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py b/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py index 0b04dd0ed78..88868844bb1 100644 --- a/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py +++ b/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py @@ -6,6 +6,7 @@ Source: litellm/llms/chatgpt/responses/transformation.py import json from collections.abc import Generator +from typing import Final from unittest.mock import MagicMock, patch import httpx @@ -30,6 +31,46 @@ def local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> Generator[None, Non class TestChatGPTResponsesAPITransformation: + @pytest.mark.parametrize( + ("requested_tier", "expected_tier"), + [("default", "default"), ("priority", "priority"), ("fast", "priority")], + ) + @pytest.mark.parametrize("effort", ["low", "high"]) + def test_chatgpt_preserves_service_tier(self, requested_tier: str, expected_tier: str, effort: str) -> None: + config: Final = ChatGPTResponsesAPIConfig() + request: Final = config.transform_responses_api_request( + model="chatgpt/gpt-6.1-sol", + input=[{"role": "user", "content": "Reply with OK"}], + response_api_optional_request_params={ + "service_tier": requested_tier, + "reasoning": {"effort": effort}, + "max_output_tokens": 16, + "prompt_cache_options": {"ttl": "30m"}, + }, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert request["service_tier"] == expected_tier + assert request["reasoning"] == {"effort": effort} + assert request["stream"] is True + assert request["store"] is False + assert "max_output_tokens" not in request + assert "prompt_cache_options" not in request + + @pytest.mark.parametrize("requested_tier", [None, "auto", "flex", "unknown"]) + def test_chatgpt_does_not_introduce_unsupported_service_tier(self, requested_tier: str | None) -> None: + config: Final = ChatGPTResponsesAPIConfig() + request: Final = config.transform_responses_api_request( + model="chatgpt/gpt-6.1-sol", + input=[{"role": "user", "content": "Reply with OK"}], + response_api_optional_request_params={} if requested_tier is None else {"service_tier": requested_tier}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert "service_tier" not in request + @pytest.mark.parametrize( "model_name", [ @@ -55,7 +96,6 @@ class TestChatGPTResponsesAPITransformation: assert isinstance(config, ChatGPTResponsesAPIConfig) assert config.custom_llm_provider == LlmProviders.CHATGPT - @pytest.mark.parametrize( "model_name", [ @@ -92,14 +132,10 @@ class TestChatGPTResponsesAPITransformation: url = config.get_complete_url(api_base=None, litellm_params={}) assert url == "https://chatgpt.example.com/responses" - custom_url = config.get_complete_url( - api_base="https://custom.chatgpt.com", litellm_params={} - ) + custom_url = config.get_complete_url(api_base="https://custom.chatgpt.com", litellm_params={}) assert custom_url == "https://custom.chatgpt.com/responses" - url_with_slash = config.get_complete_url( - api_base="https://chatgpt.example.com/", litellm_params={} - ) + url_with_slash = config.get_complete_url(api_base="https://chatgpt.example.com/", litellm_params={}) assert url_with_slash == "https://chatgpt.example.com/responses" @patch("litellm.llms.chatgpt.responses.transformation.Authenticator") @@ -162,9 +198,7 @@ class TestChatGPTResponsesAPITransformation: "user": "user_123", "temperature": 0.2, "top_p": 0.9, - "context_management": [ - {"type": "compaction", "compact_threshold": 200000} - ], + "context_management": [{"type": "compaction", "compact_threshold": 200000}], "metadata": {"foo": "bar"}, "max_output_tokens": 123, "stream_options": {"include_usage": True}, @@ -203,9 +237,7 @@ class TestChatGPTResponsesAPITransformation: ("chatgpt/gpt-5.3-codex", "gpt-5.3-codex"), ], ) - def test_chatgpt_non_stream_sse_response_parsing( - self, model_name: str, response_model: str - ): + def test_chatgpt_non_stream_sse_response_parsing(self, model_name: str, response_model: str): config = ChatGPTResponsesAPIConfig() response_payload = { "id": "resp_test", @@ -228,9 +260,7 @@ class TestChatGPTResponsesAPITransformation: "", ] ) - raw_response = httpx.Response( - 200, headers={"content-type": "text/event-stream"}, text=sse_body - ) + raw_response = httpx.Response(200, headers={"content-type": "text/event-stream"}, text=sse_body) logging_obj = MagicMock() parsed = config.transform_response_api_response( @@ -248,9 +278,7 @@ class TestChatGPTResponsesAPITransformation: ("chatgpt/gpt-5.3-codex", "gpt-5.3-codex"), ], ) - def test_chatgpt_non_stream_sse_response_recovers_output_items( - self, model_name: str, response_model: str - ): + def test_chatgpt_non_stream_sse_response_recovers_output_items(self, model_name: str, response_model: str): config = ChatGPTResponsesAPIConfig() response_payload = { "id": "resp_test", @@ -273,9 +301,7 @@ class TestChatGPTResponsesAPITransformation: "", ] ) - raw_response = httpx.Response( - 200, headers={"content-type": "text/event-stream"}, text=sse_body - ) + raw_response = httpx.Response(200, headers={"content-type": "text/event-stream"}, text=sse_body) logging_obj = MagicMock() parsed = config.transform_response_api_response( @@ -315,9 +341,7 @@ class TestChatGPTResponsesAPITransformation: "", ] ) - raw_response = httpx.Response( - 200, headers={"content-type": "text/event-stream"}, text=sse_body - ) + raw_response = httpx.Response(200, headers={"content-type": "text/event-stream"}, text=sse_body) logging_obj = MagicMock() parsed = config.transform_response_api_response( @@ -350,9 +374,7 @@ class TestChatGPTResponsesAPITransformation: "", ] ) - raw_response = httpx.Response( - 502, headers={"content-type": "text/event-stream"}, text=sse_body - ) + raw_response = httpx.Response(502, headers={"content-type": "text/event-stream"}, text=sse_body) logging_obj = MagicMock() with pytest.raises(OpenAIError) as exc_info: From 4c648f181af364bea937de39425597e1e1bc4603 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 11:17:51 -0700 Subject: [PATCH 014/139] fix(ui): leave unset callback select params out of the save payload (#44213) Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../src/components/settings.test.tsx | 34 +++++++++++++++++++ .../src/components/settings.tsx | 11 ++++-- 2 files changed, 43 insertions(+), 2 deletions(-) diff --git a/ui/litellm-dashboard/src/components/settings.test.tsx b/ui/litellm-dashboard/src/components/settings.test.tsx index 9b67657dbca..280bffd8fe9 100644 --- a/ui/litellm-dashboard/src/components/settings.test.tsx +++ b/ui/litellm-dashboard/src/components/settings.test.tsx @@ -439,6 +439,40 @@ describe("Settings", () => { }); }); + it("should post the saved s3_v2 folder partitioning when unchanged", async () => { + mockS3Callback({ S3_LOG_PROMPTS_ONLY: null, S3_PARTITION_GRANULARITY: "hour" }, "s3_v2"); + const user = await openS3EditModal("s3_v2"); + + const dialog = screen.getByRole("dialog"); + const partitioning = await within(dialog).findByRole("combobox", { name: "Folder Partitioning" }); + expect(partitioning).toHaveTextContent("hour"); + + await user.click(await within(dialog).findByRole("button", { name: "Save Changes" })); + await waitFor(() => { + expect(vi.mocked(setCallbacksCall)).toHaveBeenCalledTimes(1); + }); + const [, payload] = vi.mocked(setCallbacksCall).mock.calls[0]; + expect(payload.environment_variables.s3_partition_granularity).toBe("hour"); + }); + + it("should leave an unset s3_v2 folder partitioning out of the save payload", async () => { + mockS3Callback({ S3_LOG_PROMPTS_ONLY: null, S3_PARTITION_GRANULARITY: null }, "s3_v2"); + const user = await openS3EditModal("s3_v2"); + + const dialog = screen.getByRole("dialog"); + const partitioning = await within(dialog).findByRole("combobox", { name: "Folder Partitioning" }); + expect(partitioning).toHaveTextContent("Select folder partitioning"); + expect(partitioning).not.toHaveTextContent(/day|hour/i); + + await user.click(await within(dialog).findByRole("button", { name: "Save Changes" })); + await waitFor(() => { + expect(vi.mocked(setCallbacksCall)).toHaveBeenCalledTimes(1); + }); + const [, payload] = vi.mocked(setCallbacksCall).mock.calls[0]; + expect(Object.keys(payload.environment_variables)).not.toContain("s3_partition_granularity"); + expect(payload.environment_variables.callback).toBe("s3_v2"); + }); + it("should not offer folder partitioning for the legacy s3 callback, which cannot honour it", async () => { mockS3Callback({ S3_LOG_PROMPTS_ONLY: null, S3_PARTITION_GRANULARITY: null }); const user = await openS3EditModal(); diff --git a/ui/litellm-dashboard/src/components/settings.tsx b/ui/litellm-dashboard/src/components/settings.tsx index e376d858df8..7b08c53df87 100644 --- a/ui/litellm-dashboard/src/components/settings.tsx +++ b/ui/litellm-dashboard/src/components/settings.tsx @@ -285,7 +285,7 @@ const getDynamicParamsForCallback = ( // Shared helper function to build callback payload const buildCallbackPayload = (formValues: Record, callbackName: string) => { return { - environment_variables: formValues, + environment_variables: Object.fromEntries(Object.entries(formValues).filter(([, value]) => value !== undefined)), litellm_settings: { success_callback: [callbackName], }, @@ -348,8 +348,15 @@ const Settings: React.FC = ({ accessToken, userRole, userID, ); const fieldNameFor = (variable: string) => params.find((param) => param.toUpperCase() === variable.toUpperCase()) ?? variable; + const callbackConfig = findCallbackConfig(callbackConfigs, selectedEditCallback.name); const normalized = Object.fromEntries( - Object.entries(selectedEditCallback.variables || {}).map(([k, v]) => [fieldNameFor(k), v ?? ""]), + Object.entries(selectedEditCallback.variables || {}).flatMap(([key, value]) => { + const fieldName = fieldNameFor(key); + if (value == null && callbackConfig?.dynamic_params?.[fieldName]?.type === "select") { + return []; + } + return [[fieldName, value ?? ""]]; + }), ); editForm.reset({ ...normalized, From 4758fce91a3cdbf54009705152a24828e7433681 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 11:21:23 -0700 Subject: [PATCH 015/139] fix(proxy-extras): hand libpq a root cert, not Prisma's sslcert, when the migration job builds indexes (#44203) The migration job runs DatabaseURLSettings.apply_to_env(), which rewrites DATABASE_SSLMODE=verify-full plus DATABASE_SSLROOTCERT into Prisma's TLS dialect: sslmode=require&sslcert=&sslaccept=strict. The request-log index build then hands that same URL to psycopg, and libpq reads sslcert as a client certificate, failing with "certificate present, but not private key file" on every verify-full deployment since #43948. _strip_prisma_query_params now undoes the Prisma dialect before psycopg sees the URL. sslaccept=strict (or any value Prisma treats as strict) becomes sslrootcert= plus sslmode=verify-full whatever sslmode said, since strict verifies chain and hostname and libpq only does that in verify-full; sslmode=disable stays off. Without strict, Prisma verifies nothing, so the CA is dropped and sslmode is kept as is. A URL that also carries sslkey is libpq's own client-certificate form and is left alone. Resolves LIT-9169 Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../litellm_proxy_extras/utils.py | 70 ++++++++++++----- .../test_litellm_proxy_extras_utils.py | 78 +++++++++++++++++++ 2 files changed, 128 insertions(+), 20 deletions(-) diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py index 3acc19d397d..2062ca93fb3 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/utils.py +++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py @@ -78,6 +78,23 @@ class _InvalidIndex: table_size: str MAX_MIGRATE_DEPLOY_ATTEMPTS = 4 +LIBPQ_URL_PARAMS: Final = frozenset( + { + "sslmode", + "sslcert", + "sslkey", + "sslrootcert", + "sslpassword", + "application_name", + "connect_timeout", + "client_encoding", + "options", + "service", + "gssencmode", + "krbsrvname", + "target_session_attrs", + } +) @dataclass(frozen=True) @@ -689,30 +706,43 @@ class ProxyExtrasDBManager: @staticmethod def _strip_prisma_query_params(url: str) -> str: - """Remove Prisma-specific query params (connection_limit, pool_timeout, - schema, etc.) from DATABASE_URL so psycopg can parse it.""" + """Rewrite a Prisma-dialect URL for libpq: drop the Prisma-only params + (connection_limit, pool_timeout, schema, pgbouncer, sslaccept, ...) and + translate Prisma's TLS params back, since libpq reads ``sslcert`` as a + client certificate where Prisma reads it as the CA.""" from urllib.parse import parse_qsl, quote, urlencode, urlparse, urlunparse - parsed = urlparse(url) + parsed: Final = urlparse(url) if not parsed.query: return url - libpq_params = { - "sslmode", - "sslcert", - "sslkey", - "sslrootcert", - "sslpassword", - "application_name", - "connect_timeout", - "client_encoding", - "options", - "service", - "gssencmode", - "krbsrvname", - "target_session_attrs", - } - kept = [(k, v) for k, v in parse_qsl(parsed.query) if k in libpq_params] - return urlunparse(parsed._replace(query=urlencode(kept, quote_via=quote))) + pairs: Final = tuple(parse_qsl(parsed.query)) + kept: Final = tuple((k, v) for k, v in pairs if k in LIBPQ_URL_PARAMS) + sslaccept: Final = next((v for k, v in pairs if k == "sslaccept"), None) + libpq_pairs: Final = ProxyExtrasDBManager._libpq_tls_params(kept, sslaccept) + return urlunparse(parsed._replace(query=urlencode(libpq_pairs, quote_via=quote))) + + @staticmethod + def _libpq_tls_params( + pairs: "tuple[tuple[str, str], ...]", sslaccept: "str | None" + ) -> "tuple[tuple[str, str], ...]": + """Undo ``translate_libpq_ssl_params``. Prisma's ``sslcert`` is the CA and + ``sslaccept=strict`` checks chain and hostname, which libpq only does in + ``sslmode=verify-full``, so strict becomes ``sslrootcert`` plus + ``verify-full`` whatever ``sslmode`` said (``disable`` stays off). Prisma + defaults an absent ``sslaccept`` to ``accept_invalid_certs`` and anything + else to strict. Without strict it checks nothing, so the CA is dropped and + ``sslmode`` is kept as is: libpq only verifies when a root cert is present. + A URL that also carries ``sslkey`` is libpq's own client-certificate form + and is kept.""" + keys: Final = frozenset(k for k, _ in pairs) + if "sslcert" not in keys or "sslkey" in keys: + return pairs + sslmode: Final = next((v for k, v in pairs if k == "sslmode"), None) + rest: Final = tuple((k, v) for k, v in pairs if k not in ("sslcert", "sslmode")) + if sslaccept in (None, "accept_invalid_certs") or sslmode == "disable": + return rest if sslmode is None else rest + (("sslmode", sslmode),) + root_cert: Final = tuple(("sslrootcert", v) for k, v in pairs if k == "sslcert" and "sslrootcert" not in keys) + return rest + root_cert + (("sslmode", "verify-full"),) @staticmethod def _warn_if_db_ahead_of_head(migrations_dir: str) -> None: diff --git a/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py index 5540cf54193..29c9ec56d91 100644 --- a/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py +++ b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py @@ -1029,6 +1029,84 @@ class TestJWTKeyMappingCascade: +class TestStripPrismaQueryParams: + """The psycopg URL the job connects with is derived from the Prisma-dialect + DATABASE_URL, whose TLS params mean something else to libpq.""" + + @staticmethod + def _query(url: str) -> dict[str, str]: + from urllib.parse import parse_qsl, urlparse + + return dict(parse_qsl(urlparse(url).query)) + + def test_prisma_ca_sslcert_becomes_sslrootcert_with_verify_full(self): + url = "postgresql://u:p@writer:5432/db?schema=public&sslmode=require&sslcert=/tmp/pinned.pem&sslaccept=strict" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/tmp/pinned.pem"} + assert cleaned.startswith("postgresql://u:p@writer:5432/db?") + + @pytest.mark.parametrize("sslmode", ["prefer", "require"]) + @pytest.mark.parametrize("sslaccept", ["strict", "unknown-mode-prisma-treats-as-strict"]) + def test_strict_verifies_chain_and_hostname_whatever_sslmode_prisma_was_given(self, sslmode, sslaccept): + url = f"postgresql://writer/db?sslmode={sslmode}&sslcert=/certs/ca.pem&sslaccept={sslaccept}" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/certs/ca.pem"} + + def test_strict_with_tls_disabled_stays_off(self): + url = "postgresql://writer/db?sslmode=disable&sslcert=/certs/ca.pem&sslaccept=strict" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == {"sslmode": "disable"} + + @pytest.mark.parametrize("sslaccept", ["&sslaccept=accept_invalid_certs", ""]) + def test_without_strict_the_ca_is_dropped_so_libpq_checks_nothing_like_prisma(self, sslaccept): + url = f"postgresql://writer/db?sslmode=require&sslcert=/certs/ca.pem{sslaccept}" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == {"sslmode": "require"} + + def test_a_ca_alone_without_strict_or_sslmode_leaves_libpq_its_defaults(self): + cleaned = ProxyExtrasDBManager._strip_prisma_query_params("postgresql://writer/db?sslcert=/certs/ca.pem") + + assert cleaned == "postgresql://writer/db" + + def test_a_libpq_client_certificate_pair_is_left_alone(self): + url = "postgresql://writer/db?sslmode=verify-full&sslrootcert=/ca.pem&sslcert=/client.crt&sslkey=/client.key" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == { + "sslmode": "verify-full", + "sslrootcert": "/ca.pem", + "sslcert": "/client.crt", + "sslkey": "/client.key", + } + + def test_an_explicit_sslrootcert_wins_over_the_prisma_sslcert(self): + url = "postgresql://writer/db?sslmode=require&sslrootcert=/ca.pem&sslcert=/pinned.pem&sslaccept=strict" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/ca.pem"} + + def test_prisma_only_params_are_dropped_and_plain_urls_pass_through(self): + url = "postgresql://u:p@pooler:6543/db?schema=tenant&pgbouncer=true&connection_limit=5&connect_timeout=3" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert cleaned == "postgresql://u:p@pooler:6543/db?connect_timeout=3" + assert ( + ProxyExtrasDBManager._strip_prisma_query_params("postgresql://u:p@writer/db") + == "postgresql://u:p@writer/db" + ) + + class TestBuildRequestLogIndexes: """The migration job hands the index build the direct database URL and the schema the migrations target, waits for it, and reports its result.""" From 0dc23406eb8b1f713e1a2677bb537bfe40c19a9f Mon Sep 17 00:00:00 2001 From: tin-berri Date: Fri, 2 Oct 2026 11:29:36 -0700 Subject: [PATCH 016/139] feat: add Laya gateway and OSS classifier providers (#43626) * feat: add Laya gateway and classifier backend * test: cover the pass-through model_group pin and repair the shard fakes MockRequest in tests/pass_through_unit_tests gains an httpx.URL and an ASGI scope, which get_request_route now reads inside _init_kwargs_for_pass_through_endpoint, and the POST-only /laya/v1/systemone route joins the protocol-constrained exemptions. A built-in pass-through pins metadata.model_group to the resolved model so a client cannot choose its own per-model budget key; test_pass_through_endpoints now proves that on a non-Laya route and drops a duplicated assertion. Co-Authored-By: Claude Fable 5.1 --------- Co-authored-by: Claude Fable 5.1 --- litellm/llms/laya/__init__.py | 1 + litellm/llms/laya/common_utils.py | 60 ++++ ...odel_prices_and_context_window_backup.json | 39 +++ litellm/proxy/_lazy_features.py | 1 + litellm/proxy/_lazy_openapi_snapshot.json | 24 ++ litellm/proxy/_types.py | 1 + litellm/proxy/auth/auth_utils.py | 9 + .../auto_router_endpoints.py | 6 +- .../model_management_endpoints.py | 70 +++-- .../auto_router_permissions.py | 44 +-- .../llm_passthrough_endpoints.py | 43 +++ .../typesafe_passthrough_logging_handler.py | 5 +- .../pass_through_endpoints.py | 54 +++- .../pass_through_endpoints/success_handler.py | 6 +- .../complexity_router/complexity_router.py | 42 ++- .../complexity_router/config.py | 162 ++++++++-- .../complexity_router/jev_classifier.py | 43 ++- .../router_utils/auto_router_model_naming.py | 25 +- model_prices_and_context_window.json | 39 +++ provider_endpoints_support.json | 15 + .../test_pass_through_unit_tests.py | 4 +- tests/unit/llms/laya/__init__.py | 1 + tests/unit/llms/laya/test_common_utils.py | 60 ++++ tests/unit/proxy/auth/test_auth_utils.py | 20 +- .../test_model_management_endpoints.py | 284 +++++++++++++++++- .../test_auto_router_permissions.py | 148 +++++++-- ...st_typesafe_passthrough_logging_handler.py | 80 +++++ .../test_llm_pass_through_endpoints.py | 148 +++++++++ .../test_pass_through_endpoints.py | 142 ++++++--- ...t_passthrough_guardrail_block_otel_span.py | 21 +- .../test_vertex_passthrough_load_balancing.py | 17 +- .../complexity_router/test_jev_classifier.py | 129 +++++++- .../test_auto_router_model_naming.py | 71 ++++- tests/unit/test_utils.py | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 127 +++++--- 35 files changed, 1672 insertions(+), 270 deletions(-) create mode 100644 litellm/llms/laya/__init__.py create mode 100644 litellm/llms/laya/common_utils.py create mode 100644 tests/unit/llms/laya/__init__.py create mode 100644 tests/unit/llms/laya/test_common_utils.py diff --git a/litellm/llms/laya/__init__.py b/litellm/llms/laya/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/litellm/llms/laya/__init__.py @@ -0,0 +1 @@ + diff --git a/litellm/llms/laya/common_utils.py b/litellm/llms/laya/common_utils.py new file mode 100644 index 00000000000..f400eef22d3 --- /dev/null +++ b/litellm/llms/laya/common_utils.py @@ -0,0 +1,60 @@ +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import Final, Literal, TypeAlias + +from pydantic import AnyHttpUrl, BaseModel, TypeAdapter, ValidationError + +from litellm.secret_managers.main import get_secret_str + +LayaCheckpoint: TypeAlias = Literal["english", "multilingual", "typed-decisions"] + + +def validate_laya_model(value: object) -> LayaCheckpoint: + try: + return TypeAdapter(LayaCheckpoint).validate_python(value) + except ValidationError as exc: + raise ValueError("Laya model must be 'english', 'multilingual', or 'typed-decisions'") from exc + + +def validate_laya_request(body: Mapping[str, object]) -> LayaCheckpoint: + if "custom_body" in body: + raise ValueError("custom_body is not supported for Laya requests") + if body.get("stream"): + raise ValueError("Streaming is not supported for Laya requests") + return validate_laya_model(body.get("model")) + + +@dataclass(frozen=True, slots=True) +class LayaConnection: + api_base: str + api_key: str | None = field(repr=False) + + +def validate_laya_api_base(value: str) -> str: + try: + url: Final = TypeAdapter(AnyHttpUrl).validate_python(value) + except ValidationError as exc: + raise ValueError("Laya api_base must be an HTTP or HTTPS server URL") from exc + if url.username or url.password or url.query or url.fragment: + raise ValueError("Laya api_base must not contain credentials, a query, or a fragment") + return str(url).rstrip("/") + + +def laya_connection(api_base: str | None = None, api_key: str | None = None) -> LayaConnection: + base: Final = api_base if api_base is not None else get_secret_str("LAYA_API_BASE") + if not base: + raise ValueError("Laya requires api_base or LAYA_API_BASE pointing to a self-hosted server") + key: Final = api_key if api_base is not None else api_key or get_secret_str("LAYA_API_KEY") + return LayaConnection(api_base=validate_laya_api_base(base), api_key=key) + + +class _LayaRouting(BaseModel): + model: str | None = None + + +def laya_response_model(response: Mapping[str, object], requested_model: str | None) -> str: + try: + routing: Final = TypeAdapter(_LayaRouting).validate_python(response.get("routing") or _LayaRouting()) + except ValidationError: + return requested_model or "unknown" + return routing.model or requested_model or "unknown" diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 6a548e7a82d..3985a232f62 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -72622,6 +72622,45 @@ "supports_audio_input": true, "supports_video_input": true }, + "laya/english": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, + "laya/multilingual": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, + "laya/typed-decisions": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, "typesafe/jev-1.13.0": { "input_cost_per_token": 4.2e-08, "litellm_provider": "typesafe", diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 0119ce4430e..837552e522f 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -229,6 +229,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( "/tinyfish/", "/transcribe", "/typesafe/", + "/laya/", "/openrouter/", "/vertex-ai/", "/vertex_ai/", diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 5de6e91a0f2..3346b0c9ff8 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -27761,6 +27761,30 @@ ] } }, + "/laya/v1/systemone": { + "post": { + "operationId": "laya_proxy_route_laya_v1_systemone_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Laya Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, "/milvus/{endpoint}": { "delete": { "description": "Enable using Milvus `/vectors` endpoint as a pass-through endpoint.", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 988590bfdef..ef4c545507b 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -507,6 +507,7 @@ class LiteLLMRoutes(enum.Enum): "/vllm", "/mistral", "/typesafe", + "/laya", "/openrouter", "/milvus", "/gigachat", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 813d72ed7ae..e5a1430b3d8 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1883,6 +1883,15 @@ def _extract_model_candidates_from_request( llm_router: Router | None = None, team_id: str | None = None, ) -> list[str]: + if route.rstrip("/") == "/laya/v1/systemone": + from litellm.llms.laya.common_utils import validate_laya_model + + try: + laya_request: Final = TypeAdapter(Mapping[str, object]).validate_python(request_data) + laya_model: Final = validate_laya_model(laya_request.get("model")) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + return _dedupe_model_candidates((f"laya/{laya_model}",)) if route == "/cost/predict-cache": prediction_models: Final = _cache_prediction_model_candidates(request_data, llm_router, team_id) # pyright: ignore[reportUnknownArgumentType] # the typed reader validates each deployment ID from this legacy payload return _dedupe_model_candidates(prediction_models) diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 8b73b8177f4..50145ed3923 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -316,7 +316,7 @@ async def _authorize_models_this_test_can_call( its calls through the proxy. Team and member budgets are already enforced on every route. """ models: Final = _models_this_test_can_call(config) - if not models and config.classifier_type != "jev": + if not models and config.classifier_type != "oss_classifier": return from litellm.proxy.proxy_server import proxy_logging_obj @@ -342,9 +342,9 @@ async def _authorize_models_this_test_can_call( code=status.HTTP_400_BAD_REQUEST, ) from e - if config.classifier_type == "jev" and user_api_key_dict.budget_throttle_pct is not None: + if config.classifier_type == "oss_classifier" and user_api_key_dict.budget_throttle_pct is not None: raise ProxyException( - message="Budget has been exceeded! JEV Test Routing requires available budget.", + message="Budget has been exceeded! OSS Classifier Test Routing requires available budget.", type=ProxyErrorTypes.budget_exceeded, param=None, code=status.HTTP_400_BAD_REQUEST, diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 0f0d156d649..50bb831d169 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -39,6 +39,7 @@ from litellm.litellm_core_utils.ptu_pricing import ( SEARCH_CONTEXT_SIZES, ptu_config_error, ) +from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload from litellm.proxy._types import ( BlockModelRequest, CommonProxyErrors, @@ -115,6 +116,7 @@ from litellm.router_strategy.complexity_router import ( normalize_classification_examples, normalize_classification_prompt, ) +from litellm.router_strategy.complexity_router.config import resolve_complexity_router_config_write from litellm.router_utils.auto_router_model_naming import ( GATED_AUTO_ROUTER_CAPABILITIES, STRATEGY_ROUTER_PARAM_FIELDS, @@ -187,6 +189,25 @@ class _ProxyModelRow(Protocol): def model_dump_json(self, *, exclude_none: bool = False) -> str: ... +def _model_write_response( + row: _ProxyModelRow, member_write: MemberAutoRouterWrite | None +) -> _ProxyModelRow | Mapping[str, object]: + if member_write is None: + return row + payload: Final = TypeAdapter(dict[str, object]).validate_json(row.model_dump_json()) + stored_params: Final = payload.get("litellm_params") + params: Final = ( + TypeAdapter(dict[str, object]).validate_json(stored_params) + if isinstance(stored_params, str) + else TypeAdapter(dict[str, object]).validate_python(stored_params) + ) + redacted: Final = redact_credentials_in_payload(params) + return { + **payload, + "litellm_params": json.dumps(redacted) if isinstance(stored_params, str) else redacted, + } + + class _ProxyModelTable(Protocol): def find_unique(self, *, where: Mapping[str, object]) -> Awaitable[BaseModel | None]: ... @@ -407,34 +428,13 @@ WHERE model_id <> $1 def _effective_complexity_router_config( incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None -) -> object: +) -> Mapping[str, object] | None: incoming: Final = None if incoming_params is None else incoming_params.complexity_router_config existing: Final = None if existing_params is None else existing_params.complexity_router_config - if incoming is None: - return existing - if existing is None or incoming.get("classifier_type") != "jev" or existing.get("classifier_type") != "jev": - return incoming - incoming_jev: Final[object] = incoming.get("jev_classifier_config") - existing_jev: Final[object] = existing.get("jev_classifier_config") - if not isinstance(incoming_jev, Mapping) or not isinstance(existing_jev, Mapping): - return incoming - supplied: Final = TypeAdapter(dict[str, object]).validate_python(incoming_jev) - stored: Final = TypeAdapter(dict[str, object]).validate_python(existing_jev) - same_base: Final = "api_base" not in supplied or supplied["api_base"] == stored.get("api_base") - transport: Final = MappingProxyType( - { - key: value - for key, value in stored.items() - if key in ("api_key", "api_base") and (key != "api_key" or same_base) - } - ) - return { - **incoming, - "jev_classifier_config": { - **transport, - **supplied, - }, - } + config_adapter: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None) + return resolve_complexity_router_config_write( + config_adapter.validate_python(incoming), config_adapter.validate_python(existing) + ).effective def _effective_model( @@ -1304,7 +1304,7 @@ async def patch_model( live_after=reload_outcome.live_after, ) - return updated_model + return _model_write_response(updated_model, member_write) except Exception as e: verbose_proxy_logger.exception("Error in patch_model: %s", e) @@ -1501,10 +1501,18 @@ async def _add_model_to_db( slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None, ) -> "_ProxyModelRow | LiteLLM_ProxyModelTable": # encrypt litellm params # - _litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True) + _litellm_params_dict: Final = TypeAdapter(dict[str, object]).validate_python( + model_params.litellm_params.model_dump(exclude_none=True) + ) + if "complexity_router_config" in _litellm_params_dict: + _litellm_params_dict["complexity_router_config"] = _effective_complexity_router_config( + model_params.litellm_params, None + ) _original_litellm_model_name: Final = model_params.litellm_params.model for k, v in _litellm_params_dict.items(): - encrypted_value = encrypt_value_helper(value=v, new_encryption_key=new_encryption_key) + encrypted_value = ( + encrypt_value_helper(value=v, new_encryption_key=new_encryption_key) if isinstance(v, str) else v + ) model_params.litellm_params[k] = encrypted_value _data: Final[dict] = { "model_id": model_params.model_info.id, @@ -2536,7 +2544,7 @@ async def add_new_model( live_after=reload_outcome.live_after, ) - return model_response + return _model_write_response(model_response, member_write) except Exception as e: verbose_proxy_logger.exception("litellm.proxy.proxy_server.add_new_model(): Exception occured - %s", e) @@ -2760,7 +2768,7 @@ async def update_model( live_after=reload_outcome.live_after, ) - return model_response + return None if model_response is None else _model_write_response(model_response, member_write) except Exception as e: verbose_proxy_logger.exception("litellm.proxy.proxy_server.update_model(): Exception occured - %s", e) if isinstance(e, HTTPException): diff --git a/litellm/proxy/management_helpers/auto_router_permissions.py b/litellm/proxy/management_helpers/auto_router_permissions.py index c8bb95eb3bb..e0d8fda5b1c 100644 --- a/litellm/proxy/management_helpers/auto_router_permissions.py +++ b/litellm/proxy/management_helpers/auto_router_permissions.py @@ -33,6 +33,10 @@ from litellm.repositories.prisma_protocols import DatabaseClient from litellm.repositories.project_repository import ProjectRepository from litellm.repositories.table_repositories import TeamMembershipRepository from litellm.router import Router +from litellm.router_strategy.complexity_router.config import ( + ComplexityRouterConfigWrite, + resolve_complexity_router_config_write, +) from litellm.router_utils.auto_router_model_naming import classify_strategy_router_model, strategy_router_dependencies from litellm.types.management_endpoints.auto_router_endpoints import RequestComplexityRouterConfig from litellm.types.router import Deployment, updateDeployment @@ -65,12 +69,12 @@ class _MemberRouterGenerationParams(BaseModel): stop: str | tuple[str, ...] | None = None -class _MemberJevClassifierConfig(BaseModel): - """The Jev classifier settings a team member may set. Credentials stay the proxy's own: a member-chosen - api_base would receive the proxy's TYPESAFE_API_KEY, and a member-chosen api_key would be sent from the proxy.""" +class _MemberOpenSourceClassifierConfig(BaseModel): + """Classifier settings a team member may set while the gateway owns the connection.""" model_config = ConfigDict(extra="forbid") + provider: Literal["jev", "laya"] = "jev" model: str api_key: None = None api_base: None = None @@ -123,14 +127,21 @@ def authorize_member_auto_router_team( def validate_member_auto_router_config(config: Mapping[str, object]) -> RequestComplexityRouterConfig: + return _validate_member_auto_router_config_write(resolve_complexity_router_config_write(config, None)) + + +def _validate_member_auto_router_config_write(write: ComplexityRouterConfigWrite) -> RequestComplexityRouterConfig: + if write.effective is None: + raise HTTPException(status_code=400, detail="A complexity_router_config is required.") try: - validated: Final = _MemberComplexityRouterConfig.model_validate(config) - for entries in validated.tier_model_configs.values(): - for entry in entries: - _MemberRouterGenerationParams.model_validate(entry.litellm_params) - if validated.jev_classifier_config is not None: - _MemberJevClassifierConfig.model_validate(validated.jev_classifier_config.model_dump()) - return validated + if write.submitted is not None: + validated: Final = _MemberComplexityRouterConfig.model_validate(write.submitted) + for entries in validated.tier_model_configs.values(): + for entry in entries: + _MemberRouterGenerationParams.model_validate(entry.litellm_params) + if validated.opensource_classifier_config is not None: + _MemberOpenSourceClassifierConfig.model_validate(validated.opensource_classifier_config.model_dump()) + return RequestComplexityRouterConfig.model_validate(write.effective) except ValidationError as exc: location: Final = ".".join(str(part) for part in exc.errors()[0]["loc"]) raise HTTPException(status_code=400, detail=f"Invalid member auto-router configuration at {location}.") from exc @@ -332,16 +343,15 @@ async def authorize_member_auto_router_write( if existing is not None and incoming.model_name not in (None, public_name, existing.model_name): raise HTTPException(status_code=403, detail="Team members cannot rename an auto router.") supplied_config: Final = _RouterConfigSource.model_validate(params.model_dump()).complexity_router_config - raw_config: Final = ( - supplied_config - if supplied_config is not None - else _RouterConfigSource.model_validate(existing.litellm_params.model_dump()).complexity_router_config + stored_config: Final = ( + _RouterConfigSource.model_validate(existing.litellm_params.model_dump()).complexity_router_config if existing is not None else None ) - if raw_config is None: - raise HTTPException(status_code=400, detail="A complexity_router_config is required.") - config: Final = validate_member_auto_router_config(raw_config) + resolved_config: Final = resolve_complexity_router_config_write(supplied_config, stored_config) + if resolved_config.supplied_connection_fields: + raise HTTPException(status_code=403, detail="Team members cannot change classifier connections.") + config: Final = _validate_member_auto_router_config_write(resolved_config) stored_default: Final = existing.litellm_params.complexity_router_default_model if existing is not None else None default_model: Final = ( params.complexity_router_default_model diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index c7b597b584f..ed2ea475c7a 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -24,6 +24,7 @@ from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast import httpx from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket from fastapi.responses import StreamingResponse +from pydantic import ConfigDict, TypeAdapter from starlette.websockets import WebSocketState from typing_extensions import ReadOnly, TypedDict @@ -57,6 +58,7 @@ from litellm.llms.deepgram.common_utils import ( deepgram_listen_websocket_target, ) from litellm.llms.fal_ai.cost_calculator import fal_ai_passthrough_cost, fal_ai_queue_base +from litellm.llms.laya.common_utils import laya_connection, validate_laya_request from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.passthrough.main import AsyncPassthroughStreamingResponse @@ -636,6 +638,47 @@ async def typesafe_proxy_route( return await endpoint_func(request, fastapi_response, user_api_key_dict) +@router.post( + "/laya/v1/systemone", + tags=["Laya Pass-through", "pass-through"], +) +async def laya_proxy_route( + request: Request, + fastapi_response: Response, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +) -> Response: + body: Final = TypeAdapter(dict[str, object]).validate_python(await _read_request_body(request)) + try: + _ = validate_laya_request(body) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + try: + connection: Final = laya_connection() + except ValueError as exc: + raise HTTPException( + status_code=503, detail="Laya server is not configured correctly; check LAYA_API_BASE" + ) from exc + base_url: Final = httpx.URL(connection.api_base) + updated_url: Final = base_url.copy_with( + path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, "/v1/systemone"), + ) + authorization: Final[Mapping[str, str]] = ( + MappingProxyType({"Authorization": f"Bearer {connection.api_key}"}) + if connection.api_key + else MappingProxyType({}) + ) + endpoint_func: Final = create_pass_through_route( + endpoint="v1/systemone", + target=str(updated_url), + custom_headers=MappingProxyType({**authorization, "Content-Type": "application/json"}), + custom_llm_provider="laya", + is_streaming_request=False, + ) + return TypeAdapter(Response, config=ConfigDict(arbitrary_types_allowed=True)).validate_python( + await endpoint_func(request, fastapi_response, user_api_key_dict) + ) + + @router.api_route( "/openrouter/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py index 3ad92acb48a..03ec559b83d 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py @@ -10,6 +10,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.litellm_logging import ( get_standard_logging_object_payload, # pyright: ignore[reportUnknownVariableType] # legacy helper has an untyped signature ) +from litellm.llms.laya.common_utils import laya_response_model from litellm.proxy._types import PassThroughEndpointLoggingTypedDict from litellm.types.utils import ModelResponse, StandardPassThroughResponseObject, Usage @@ -69,9 +70,11 @@ class TypeSafePassthroughLoggingHandler: **kwargs: object, ) -> PassThroughEndpointLoggingTypedDict: response: Final = _parse_typesafe_response(response_body) - response_model: Final = response.model request_model_value: Final = request_body.get("model") request_model: Final = request_model_value if isinstance(request_model_value, str) else None + response_model: Final = ( + laya_response_model(response_body, request_model) if custom_llm_provider == "laya" else response.model + ) logged_model: Final = response_model or request_model or "unknown" model_name: Final = f"{custom_llm_provider}/{logged_model}" usage: Final = response.usage or _TypeSafeUsage() diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 22ecdc06ed9..fa592629933 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -26,6 +26,7 @@ from fastapi import ( status, ) from fastapi.responses import StreamingResponse +from pydantic import TypeAdapter from starlette.datastructures import UploadFile as StarletteUploadFile from starlette.websockets import WebSocketState from websockets.asyncio.client import connect @@ -64,6 +65,7 @@ from litellm.llms.base_llm.managed_resources.utils import ( resolve_passthrough_managed_id_provider, ) from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.llms.laya.common_utils import validate_laya_request from litellm.passthrough import BasePassthroughUtils from litellm.proxy._types import ( ConfigFieldInfo, @@ -74,7 +76,11 @@ from litellm.proxy._types import ( ProxyException, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_utils import request_dispatched_to_pass_through_endpoint +from litellm.proxy.auth.auth_utils import ( + get_model_from_request, + get_request_route, + request_dispatched_to_pass_through_endpoint, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, @@ -100,6 +106,8 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above + _key_or_team_allows_client_pricing_override, # pyright: ignore[reportPrivateUsage] # reuse the proxy's pricing trust policy + _strip_client_pricing_overrides, # pyright: ignore[reportPrivateUsage] # sanitize before trusted hooks add guardrail costs ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path @@ -585,7 +593,18 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): """ Filter out litellm params from the request body """ + from litellm.proxy.proxy_server import llm_router + _parsed_body = _parsed_body or {} + managed_model: Final = get_model_from_request( + request_data=_parsed_body, + route=get_request_route(request), + request_headers=request.headers, + request_query_params=request.query_params, + llm_router=llm_router, + request=request, + team_id=user_api_key_dict.team_id, + ) litellm_keys_in_body: Final = MappingProxyType( {k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body} @@ -631,10 +650,19 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): # would attribute it to a budget the operator scoped to a LiteLLM model that # merely shares the name. if not request_dispatched_to_pass_through_endpoint(request): + _metadata["model_group"] = managed_model if isinstance(managed_model, str) else None _metadata["user_api_key_model_max_budget"] = user_api_key_dict.model_max_budget _metadata["user_api_key_team_model_max_budget"] = user_api_key_dict.team_model_max_budget _metadata["user_api_key_user_model_max_budget"] = user_api_key_dict.user_model_max_budget _metadata["user_api_key_end_user_model_max_budget"] = user_api_key_dict.end_user_model_max_budget + else: + for field in ( + "user_api_key_model_max_budget", + "user_api_key_team_model_max_budget", + "user_api_key_user_model_max_budget", + "user_api_key_end_user_model_max_budget", + ): + _metadata.pop(field, None) _metadata.update( LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict) ) @@ -1131,6 +1159,15 @@ async def pass_through_request( _parsed_body, ) + if not _key_or_team_allows_client_pricing_override(user_api_key_dict): + pricing_body: Final = TypeAdapter(dict[str, object]).validate_python(_parsed_body) + _strip_client_pricing_overrides(pricing_body) + _parsed_body = pricing_body + if custom_llm_provider == "laya": + laya_request: Final = TypeAdapter(Mapping[str, object]).validate_python(_parsed_body) + checkpoint: Final = validate_laya_request(laya_request) + _parsed_body["model"] = f"laya/{checkpoint}" + ### COLLECT GUARDRAILS FOR PASSTHROUGH ENDPOINT ### # Passthrough endpoints are opt-in only for guardrails # When enabled, collect guardrails from org/team/key levels + passthrough-specific @@ -1186,6 +1223,17 @@ async def pass_through_request( call_type="pass_through_endpoint", endpoint_type=endpoint_type, ) + if custom_llm_provider == "laya": + hook_body: Final = TypeAdapter(dict[str, object]).validate_python(_parsed_body) + hook_model: Final = hook_body.get("model") + laya_body: Final = MappingProxyType( + { + **hook_body, + "model": hook_model.removeprefix("laya/") if isinstance(hook_model, str) else hook_model, + } + ) + _ = validate_laya_request(laya_body) + _parsed_body = TypeAdapter(dict[str, object]).validate_python(laya_body) resolved_timeout: Final = resolve_pass_through_request_timeout(timeout) async_client_obj: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.PassThroughEndpoint, @@ -2389,7 +2437,9 @@ async def websocket_passthrough_request( # with the existing _init_kwargs_for_pass_through_endpoint function class DummyRequest: def __init__(self, url: str, method: str = "WEBSOCKET", headers: dict | None = None): - self.url = url + self.url = httpx.URL(url) + self.scope = websocket.scope + self.query_params = websocket.query_params self.method = method self.headers = headers or {} diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 6bba879b6c1..3c4733d0bf0 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -334,8 +334,10 @@ class PassThroughEndpointLogging: ) standard_logging_response_object = transcribe_handler_result["result"] # rebind-ok: elif-chain kwargs = transcribe_handler_result["kwargs"] # rebind-ok: elif-chain contract - elif self.is_typesafe_route(custom_llm_provider) or self.is_openrouter_decisions_route( - url_route, custom_llm_provider + elif ( + self.is_typesafe_route(custom_llm_provider) + or custom_llm_provider == "laya" + or self.is_openrouter_decisions_route(url_route, custom_llm_provider) ): from .llm_provider_handlers.typesafe_passthrough_logging_handler import ( TypeSafePassthroughLoggingHandler, diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 3f56141b21c..39fb237917c 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -108,7 +108,7 @@ from .config import ( ComplexityRouterConfig, ComplexityTier, CustomDimension, - JevClassifierConfig, + OpenSourceClassifierConfig, TierDefinition, ) from .jev_classifier import ( @@ -1308,10 +1308,22 @@ class ComplexityRouter(CustomLogger): """ @staticmethod - def _build_jev_client(config: JevClassifierConfig) -> JevClassifierClient: + def _build_jev_client(config: OpenSourceClassifierConfig) -> JevClassifierClient: + if config.provider == "laya": + from litellm.llms.laya.common_utils import laya_connection + + connection: Final = laya_connection(config.api_base, config.api_key) + return HttpJevClassifierClient( + api_key=connection.api_key, + api_base=connection.api_base, + http_client=get_async_httpx_client(httpxSpecialProvider.PassThroughEndpoint), + provider="laya", + ) api_key: Final = config.api_key or get_secret_str("TYPESAFE_API_KEY") if not api_key: - raise ValueError("jev_classifier_config.api_key or TYPESAFE_API_KEY is required for classifier_type 'jev'") + raise ValueError( + "opensource_classifier_config.api_key or TYPESAFE_API_KEY is required for classifier_type 'oss_classifier'" + ) api_base: Final = config.api_base or get_secret_str("TYPESAFE_API_BASE") or "https://api.typesafe.ai" return HttpJevClassifierClient( api_key=api_key, @@ -1354,12 +1366,12 @@ class ComplexityRouter(CustomLogger): if default_model: self.config.default_model = default_model - jev_config: Final = self.config.jev_classifier_config + jev_config: Final = self.config.opensource_classifier_config self._jev_client: JevClassifierClient | None = ( jev_client if jev_client is not None else self._build_jev_client(jev_config) - if self.config.classifier_type == "jev" and jev_config is not None + if self.config.classifier_type == "oss_classifier" and jev_config is not None else None ) @@ -1459,7 +1471,11 @@ class ComplexityRouter(CustomLogger): and self.config.classifier_llm_config.circuit_breaker_enabled ) else jev_config.circuit_breaker_cooldown_seconds - if (self.config.classifier_type == "jev" and jev_config is not None and jev_config.circuit_breaker_enabled) + if ( + self.config.classifier_type == "oss_classifier" + and jev_config is not None + and jev_config.circuit_breaker_enabled + ) else None ) self._classifier_circuit_breaker: _ClassifierCircuitBreaker | None = ( @@ -1909,7 +1925,7 @@ class ComplexityRouter(CustomLogger): return self._classify_with_heuristic_v2(prompt) if self.config.classifier_type == "custom": return await self._classify_with_plugin(prompt, system_prompt, request_kwargs, raw_messages) - if self.config.classifier_type == "jev": + if self.config.classifier_type == "oss_classifier": return await self._jev_classifier_outcome(prompt, system_prompt, request_kwargs, messages) if self.config.classifier_type in ("heuristic_first", "hybrid") and _encrypted_classifier_task( request_kwargs, self._reminder_markers_for_request(request_kwargs or EMPTY_MAPPING) @@ -2161,7 +2177,7 @@ class ComplexityRouter(CustomLogger): request_kwargs: Mapping[str, object] | None, messages: Sequence[Mapping[str, object]] | None, ) -> ClassificationOutcome: - config: Final = self.config.jev_classifier_config + config: Final = self.config.opensource_classifier_config client: Final = self._jev_client if config is None or client is None: return self._classifier_failure_outcome("jev classifier is not configured", prompt, system_prompt) @@ -2212,12 +2228,14 @@ class ComplexityRouter(CustomLogger): if not self._tier_pools().get(tier_name): raise ValueError(f"Jev classifier returned tier {tier_name!r}, which has no models configured") model: Final = response.model or config.model + accounting_provider: Final = "laya" if config.provider == "laya" else "typesafe" verdict: Final = JevVerdict( label=answer.choice, probabilities=answer.probabilities, confidence=answer.confidence, model=model, - cost=jev_classifier_cost(response, config.model), + cost=jev_classifier_cost(response, config.model, accounting_provider), + provider=accounting_provider, ) if breaker is not None and permit is not None: breaker.record_success(permit) @@ -2225,8 +2243,8 @@ class ComplexityRouter(CustomLogger): tier=tier, score=None, signals=( - f"jev-classifier:{tier_name}", - f"jev-confidence={answer.confidence:.6f}", + f"{'laya' if config.provider == 'laya' else 'jev'}-classifier:{tier_name}", + f"{'laya' if config.provider == 'laya' else 'jev'}-confidence={answer.confidence:.6f}", *( f"tier-probability:{label}={probability:.6f}" for label, probability in answer.probabilities.items() @@ -4757,7 +4775,7 @@ class ComplexityRouter(CustomLogger): tier_litellm_params: Final = self._litellm_params_for_model(tier, routed_model) classifier_model: Final = ( - f"typesafe/{outcome.jev_verdict.model}" + f"{outcome.jev_verdict.provider}/{outcome.jev_verdict.model}" if outcome.cause == "jev_classifier" and outcome.jev_verdict is not None else self.config.classifier_llm_config.model if outcome.cause in ("llm_classifier", "capability_classifier", "llm_v2_classifier", "llm_v2_fallback") diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 00cff661d2f..41f389db7d8 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -9,6 +9,7 @@ import math import re import warnings from collections.abc import Iterable, Mapping +from dataclasses import dataclass from enum import Enum from types import MappingProxyType from typing import Annotated, Final, Literal, NamedTuple @@ -19,6 +20,7 @@ from pydantic import ( Field, SkipValidation, StrictFloat, + TypeAdapter, field_serializer, field_validator, model_validator, @@ -674,14 +676,34 @@ class CapabilityClassifierConfig(BaseModel): return self -class JevClassifierConfig(BaseModel): +def normalize_classifier_config_aliases(config: Mapping[str, object]) -> Mapping[str, object]: + if "jev_classifier_config" in config and "opensource_classifier_config" in config: + return config + normalized: Final = dict(config) + if "jev_classifier_config" in normalized: + normalized["opensource_classifier_config"] = normalized.pop("jev_classifier_config") + if normalized.get("classifier_type") == "jev": + normalized["classifier_type"] = "oss_classifier" + classifier: Final = normalized.get("opensource_classifier_config") + if isinstance(classifier, Mapping): + classifier_fields: Final = TypeAdapter(Mapping[str, object]).validate_python(classifier) + if classifier_fields.get("provider") == "typesafe": + normalized["opensource_classifier_config"] = { + **classifier_fields, + "provider": "jev", + } + return normalized + + +class OpenSourceClassifierConfig(BaseModel): model_config = ConfigDict(extra="forbid", frozen=True) + provider: Literal["jev", "laya"] = "jev" model: str = "jev-latest" - api_key: str | None = Field(default=None, description="TypeSafe API key, falling back to TYPESAFE_API_KEY") + api_key: str | None = Field(default=None, description="Provider API key; optional for self-hosted Laya") api_base: str | None = Field( default=None, - description="TypeSafe API base, falling back to TYPESAFE_API_BASE and then https://api.typesafe.ai", + description="Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider", ) timeout_ms: int = Field(default=3000, ge=1) instructions: str | None = Field( @@ -691,30 +713,112 @@ class JevClassifierConfig(BaseModel): circuit_breaker_enabled: bool = True circuit_breaker_cooldown_seconds: float = Field(default=30.0, gt=0.0) + @field_validator("provider", mode="before") + @classmethod + def _normalize_provider_alias(cls, value: object) -> object: + return "jev" if value == "typesafe" else value + @field_validator("instructions") @classmethod def _reject_blank_instructions(cls, value: str | None) -> str | None: if value is not None and not value.strip(): - raise ValueError("jev_classifier_config.instructions must be non-empty; omit it to use the default") + raise ValueError("opensource_classifier_config.instructions must be non-empty; omit it to use the default") return value @field_validator("api_key") @classmethod def _reject_blank_api_key(cls, value: str | None) -> str | None: if value is not None and not value.strip(): - raise ValueError("jev_classifier_config.api_key must be non-empty; omit it to use TYPESAFE_API_KEY") + raise ValueError("opensource_classifier_config.api_key must be non-empty; omit it to use TYPESAFE_API_KEY") return value @model_validator(mode="after") - def _keep_the_environment_key_on_the_environment_base(self) -> "JevClassifierConfig": + def _keep_the_environment_key_on_the_environment_base(self) -> "OpenSourceClassifierConfig": + if self.provider == "laya": + from litellm.llms.laya.common_utils import validate_laya_api_base, validate_laya_model + + _ = validate_laya_model(self.model) + if self.api_base is not None: + _ = validate_laya_api_base(self.api_base) + return self if self.api_base is not None and self.api_key is None: raise ValueError( - "jev_classifier_config.api_base requires jev_classifier_config.api_key: TYPESAFE_API_KEY is only sent " + "opensource_classifier_config.api_base requires opensource_classifier_config.api_key: TYPESAFE_API_KEY is only sent " "to TYPESAFE_API_BASE or https://api.typesafe.ai" ) return self +JevClassifierConfig = OpenSourceClassifierConfig + + +@dataclass(frozen=True, slots=True) +class ComplexityRouterConfigWrite: + submitted: Mapping[str, object] | None + effective: Mapping[str, object] | None + + @property + def supplied_connection_fields(self) -> frozenset[str]: + classifier: Final = self.submitted.get("opensource_classifier_config") if self.submitted is not None else None + return frozenset( + field for field in ("api_base", "api_key") if isinstance(classifier, Mapping) and field in classifier + ) + + +def resolve_complexity_router_config_write( + incoming: Mapping[str, object] | None, stored: Mapping[str, object] | None +) -> ComplexityRouterConfigWrite: + if incoming is None: + return ComplexityRouterConfigWrite(submitted=None, effective=stored) + return _resolve_normalized_complexity_router_config_write( + normalize_classifier_config_aliases(incoming), + normalize_classifier_config_aliases(stored) if stored is not None else None, + ) + + +def _resolve_normalized_complexity_router_config_write( + incoming: Mapping[str, object], stored: Mapping[str, object] | None +) -> ComplexityRouterConfigWrite: + if ( + stored is None + or incoming.get("classifier_type") != "oss_classifier" + or stored.get("classifier_type") != "oss_classifier" + ): + return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming) + incoming_classifier: Final = incoming.get("opensource_classifier_config") + stored_classifier: Final = stored.get("opensource_classifier_config") + if not isinstance(incoming_classifier, Mapping) or not isinstance(stored_classifier, Mapping): + return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming) + existing: Final = TypeAdapter(dict[str, object]).validate_python(stored_classifier) + supplied: Final = TypeAdapter(dict[str, object]).validate_python(incoming_classifier) + classifier: Final = ( + MappingProxyType({**supplied, "provider": existing["provider"]}) + if "provider" not in supplied and "provider" in existing + else supplied + ) + same_provider: Final = classifier.get("provider", "jev") == existing.get("provider", "jev") + same_base: Final = "api_base" not in classifier or ( + classifier["api_base"] is not None and classifier["api_base"] == existing.get("api_base") + ) + transport: Final = MappingProxyType( + { + key: value + for key, value in existing.items() + if same_provider and key in ("api_key", "api_base") and (key != "api_key" or same_base) + } + ) + return ComplexityRouterConfigWrite( + submitted=MappingProxyType({**incoming, "opensource_classifier_config": classifier}), + effective={ + **incoming, + "opensource_classifier_config": { + **transport, + **classifier, + }, + }, + ) + + MAX_CUSTOM_PATTERN_REPEAT: Final[int] = 64 MAX_CUSTOM_PATTERN_WORK: Final[int] = 2048 MAX_CUSTOM_DIMENSIONS_WORK: Final[int] = 8192 @@ -846,6 +950,20 @@ class ContextCompactionConfig(BaseModel): class ComplexityRouterConfig(BaseModel): """Configuration for the ComplexityRouter.""" + @model_validator(mode="before") + @classmethod + def _normalize_classifier_aliases(cls, value: object) -> object: + if not isinstance(value, Mapping): + return value + config: Final = TypeAdapter(dict[str, object]).validate_python(value) + if "jev_classifier_config" in config and "opensource_classifier_config" in config: + raise ValueError("Use only opensource_classifier_config; do not also supply jev_classifier_config") + return normalize_classifier_config_aliases(config) + + @property + def jev_classifier_config(self) -> OpenSourceClassifierConfig | None: + return self.opensource_classifier_config + # string = pin; list = random pick when adaptive=False, soft-floor home pool when adaptive=True tiers: dict[str, str | list[str]] = Field( default_factory=lambda: DEFAULT_TIER_MODELS.copy(), @@ -880,7 +998,7 @@ class ComplexityRouterConfig(BaseModel): "becomes that tier's rubric bullet; entries named after a built-in tier may omit the " "description and inherit the built-in criteria. List order is ascending severity and " "decides which tier wins when several keyword_tier_rules match. Requires classifier_type " - "'llm', 'jev' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, " + "'llm', 'oss_classifier' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, " "adaptive selection, session affinity, plugins, tier_labels, and the calibration-example " "rubric presets are unavailable with a custom tier set: the first four are built on the " "built-in tier ladder, and the last two rename or exemplify tiers the set replaces." @@ -1024,7 +1142,7 @@ class ComplexityRouterConfig(BaseModel): "custom", "heuristic_first", "hybrid", - "jev", + "oss_classifier", ] = Field( default="heuristic", description=( @@ -1032,7 +1150,7 @@ class ComplexityRouterConfig(BaseModel): "an LLM tier-selection call, a Switchyard-compatible capability forecast, a joint Fuse V2 forecast, " "a custom classifier plugin, 'heuristic_first', which scores locally and only pays for the LLM classifier when the " "local scorer does not confidently land a cheap tier, or 'hybrid', which trusts the local scorer " - "everywhere except when its score lands near a tier boundary, or 'jev', a TypeSafe AI Jev structured choice call" + "everywhere except when its score lands near a tier boundary, or 'oss_classifier', a structured choice call using Jev or Laya" ), ) llm_v2_config: LLMV2Config | None = Field( @@ -1073,7 +1191,7 @@ class ComplexityRouterConfig(BaseModel): "and otherwise routes to capable_tier" ), ) - jev_classifier_config: JevClassifierConfig | None = None + opensource_classifier_config: OpenSourceClassifierConfig | None = None heuristic_first_max_tier: str | None = Field( default=None, description=( @@ -1639,14 +1757,16 @@ class ComplexityRouterConfig(BaseModel): return self @model_validator(mode="after") - def _validate_jev_classifier_config(self) -> "ComplexityRouterConfig": - jev: Final = self.jev_classifier_config - if self.classifier_type != "jev": + def _validate_opensource_classifier_config(self) -> "ComplexityRouterConfig": + jev: Final = self.opensource_classifier_config + if self.classifier_type != "oss_classifier": if jev is not None: - raise ValueError("jev_classifier_config requires classifier_type 'jev'; otherwise it has no effect") + raise ValueError( + "opensource_classifier_config requires classifier_type 'oss_classifier'; otherwise it has no effect" + ) return self if jev is None: - raise ValueError("jev_classifier_config is required when classifier_type is 'jev'") + raise ValueError("opensource_classifier_config is required when classifier_type is 'oss_classifier'") return self @model_validator(mode="after") @@ -1962,9 +2082,9 @@ class ComplexityRouterConfig(BaseModel): "enable_non_reasoning_tier cannot be combined with tier_definitions: a custom tier set " f"replaces the built-in ladder, so name a tier {non_reasoning_key} in tier_definitions instead" ) - if self.classifier_type not in ("llm", "custom", "jev"): + if self.classifier_type not in ("llm", "custom", "oss_classifier"): raise ValueError( - f"enable_non_reasoning_tier requires classifier_type 'llm', 'jev' or 'custom', got " + f"enable_non_reasoning_tier requires classifier_type 'llm', 'oss_classifier' or 'custom', got " f"{self.classifier_type!r}: the heuristic scorers only produce the four tiers from SIMPLE up, " f"so nothing would ever classify as {non_reasoning_key}" ) @@ -1997,7 +2117,7 @@ class ComplexityRouterConfig(BaseModel): raise ValueError(f"tier_definitions names must be unique (case-insensitive): {', '.join(duplicated)}") if self.classifier_type in ("heuristic", "heuristic_v2", "capability", "heuristic_first", "hybrid"): raise ValueError( - "tier_definitions requires classifier_type 'llm', 'jev' or 'custom': the heuristic scorer only " + "tier_definitions requires classifier_type 'llm', 'oss_classifier' or 'custom': the heuristic scorer only " "produces the built-in tiers from SIMPLE up, as does heuristic_v2" ) conflicts: Final = self._tier_definition_conflicts() @@ -2164,7 +2284,9 @@ class ComplexityRouterConfig(BaseModel): ) -COMPLEXITY_ROUTER_CONFIG_KEYS: Final[frozenset[str]] = frozenset(ComplexityRouterConfig.model_fields) +COMPLEXITY_ROUTER_CONFIG_KEYS: Final[frozenset[str]] = frozenset(ComplexityRouterConfig.model_fields) | frozenset( + ("jev_classifier_config",) +) """Every setting name this config owns, derived from the model so a field added later is covered. These names are disjoint from the OpenAI request params, from ``all_litellm_params``, and from the diff --git a/litellm/router_strategy/complexity_router/jev_classifier.py b/litellm/router_strategy/complexity_router/jev_classifier.py index a2f03b07e3a..073d87c25a6 100644 --- a/litellm/router_strategy/complexity_router/jev_classifier.py +++ b/litellm/router_strategy/complexity_router/jev_classifier.py @@ -18,6 +18,7 @@ from litellm.litellm_core_utils.internal_call_metadata import ( from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.laya.common_utils import laya_response_model from litellm.proxy.pass_through_endpoints.llm_provider_handlers.typesafe_passthrough_logging_handler import ( TypeSafePassthroughLoggingHandler, ) @@ -78,10 +79,17 @@ class JevClassifierClient(Protocol): class HttpJevClassifierClient: - def __init__(self, api_key: str, api_base: str, http_client: AsyncHTTPHandler) -> None: + def __init__( + self, + api_key: str | None, + api_base: str, + http_client: AsyncHTTPHandler, + provider: Literal["typesafe", "laya"] = "typesafe", + ) -> None: self._api_key = api_key self._api_base = api_base.rstrip("/") self._http_client = http_client + self._provider = provider async def evaluate( self, @@ -90,26 +98,30 @@ class HttpJevClassifierClient: request_kwargs: Mapping[str, object] | None = None, ) -> JevSystemOneResponse: start_time: Final = datetime.now(timezone.utc) + authorization: Final[Mapping[str, str]] = ( + MappingProxyType({"Authorization": f"Bearer {self._api_key}"}) if self._api_key else MappingProxyType({}) + ) response: Final = await self._http_client.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler has a dynamic post signature f"{self._api_base}/v1/systemone", json=request.model_dump(mode="json"), - headers=MappingProxyType( - { - "Authorization": f"Bearer {self._api_key}", - "Content-Type": "application/json", - } - ), # pyright: ignore[reportArgumentType] # HTTP headers are not mutated by AsyncHTTPHandler + headers=MappingProxyType({**authorization, "Content-Type": "application/json"}), # pyright: ignore[reportArgumentType] # HTTP headers are not mutated by AsyncHTTPHandler timeout=timeout_s, ) response.raise_for_status() + body: Final = TypeAdapter(dict[str, object]).validate_json(response.content) + normalized_body: Final = ( + MappingProxyType({**body, "model": laya_response_model(body, request.model)}) + if self._provider == "laya" + else body + ) try: self._log_response(request, response, request_kwargs, start_time) except Exception as exc: # noqa: BLE001 # logging integrations must not discard a provider verdict verbose_router_logger.warning("JEV response logging failed (%s)", type(exc).__name__) - return TypeAdapter(JevSystemOneResponse).validate_python(response.json()) + return TypeAdapter(JevSystemOneResponse).validate_python(normalized_body) - @staticmethod def _log_response( + self, request: JevSystemOneRequest, response: httpx.Response, request_kwargs: Mapping[str, object] | None, @@ -139,7 +151,7 @@ class HttpJevClassifierClient: "turn_off_message_logging": effective_turn_off_message_logging(request_kwargs), } logging_obj: Final = Logging( - model=f"typesafe/{request.model}", + model=f"{self._provider}/{request.model}", messages=[{"role": "user", "content": request.state}], stream=False, call_type="pass_through_endpoint", @@ -150,7 +162,7 @@ class HttpJevClassifierClient: kwargs=params, ) logging_obj.update_environment_variables( - model=f"typesafe/{request.model}", + model=f"{self._provider}/{request.model}", user=parent_user if isinstance(parent_user := parent.get("user"), str) else None, optional_params={}, litellm_params=params, @@ -165,7 +177,7 @@ class HttpJevClassifierClient: end_time=end_time, cache_hit=False, request_body=MappingProxyType({"model": request.model}), - custom_llm_provider="typesafe", + custom_llm_provider=self._provider, litellm_params=params, ) success_handlers: Final = logging_obj.dispatch_success_handlers( @@ -189,6 +201,7 @@ class JevVerdict(NamedTuple): confidence: float model: str cost: float | None + provider: Literal["typesafe", "laya"] = "typesafe" class _RegistryPricing(BaseModel): @@ -211,12 +224,14 @@ def build_jev_request( return JevSystemOneRequest(state=state, model=model, questions=MappingProxyType({"tier": question})) -def jev_classifier_cost(response: JevSystemOneResponse, configured_model: str) -> float | None: +def jev_classifier_cost( + response: JevSystemOneResponse, configured_model: str, provider: Literal["typesafe", "laya"] = "typesafe" +) -> float | None: usage: Final = response.usage if usage is None: return None model: Final = response.model or configured_model - model_key: Final = f"typesafe/{model}" + model_key: Final = f"{provider}/{model}" if model_key not in litellm.model_cost: # pyright: ignore[reportUnknownMemberType] # registry is dynamically typed return None try: diff --git a/litellm/router_utils/auto_router_model_naming.py b/litellm/router_utils/auto_router_model_naming.py index 6b589c3bfc0..3423836e4fd 100644 --- a/litellm/router_utils/auto_router_model_naming.py +++ b/litellm/router_utils/auto_router_model_naming.py @@ -19,6 +19,7 @@ from litellm.router_strategy.complexity_router.config import ( COMPLEXITY_ROUTER_CONFIG_KEYS, DEFAULT_JEV_INSTRUCTIONS, LLM_CLASSIFIER_TYPES, + normalize_classifier_config_aliases, ) AUTO_ROUTER_MODEL_PREFIX: Final = "auto_router/" @@ -151,8 +152,11 @@ def strategy_router_dependencies( ) ) ) - complexity: Final = _mapping(litellm_params.get("complexity_router_config")) + complexity: Final = normalize_classifier_config_aliases(_mapping(litellm_params.get("complexity_router_config"))) classifier: Final = _mapping(complexity.get("classifier_llm_config")) + decision_classifier: Final = _mapping(complexity.get("opensource_classifier_config")) + decision_provider: Final = decision_classifier.get("provider", "jev") + accounting_provider: Final = "typesafe" if decision_provider == "jev" else decision_provider return tuple( dict.fromkeys( tuple(dep for tier in _mapping(complexity.get("tiers")).values() for dep in _pool(tier, "tier")) @@ -165,10 +169,10 @@ def strategy_router_dependencies( ) + ( _named( - f"typesafe/{_mapping(complexity.get('jev_classifier_config')).get('model', 'jev-latest')}", + f"{accounting_provider}/{decision_classifier.get('model', 'jev-latest')}", "evaluation", ) - if complexity.get("classifier_type") == "jev" + if complexity.get("classifier_type") == "oss_classifier" else () ) + ( @@ -206,9 +210,9 @@ def defines_custom_classifier_prompt(complexity_router_config: object) -> bool: Scoped to the classifier types that actually call an LLM, which is also where the config validator accepts these fields: the heuristic scorers never read them. """ - config: Final = _mapping(complexity_router_config) - if config.get("classifier_type") == "jev": - instructions: Final = _mapping(config.get("jev_classifier_config")).get("instructions") + config: Final = normalize_classifier_config_aliases(_mapping(complexity_router_config)) + if config.get("classifier_type") == "oss_classifier": + instructions: Final = _mapping(config.get("opensource_classifier_config")).get("instructions") return isinstance(instructions, str) and instructions != DEFAULT_JEV_INSTRUCTIONS if config.get("classifier_type") not in LLM_CLASSIFIER_TYPES: return False @@ -272,6 +276,9 @@ _OPERATOR_PROMPT_FIELDS_SQL: Final = " OR ".join( f"{{config}} ->> '{field}' IS NOT NULL" for field in OPERATOR_CLASSIFIER_PROMPT_FIELDS ) _DEFAULT_JEV_INSTRUCTIONS_SQL: Final = DEFAULT_JEV_INSTRUCTIONS.replace("'", "''") +_OPENSOURCE_CLASSIFIER_CONFIG_SQL: Final = ( + "COALESCE({config} -> 'opensource_classifier_config', {config} -> 'jev_classifier_config')" +) CUSTOMIZATION_CAPABILITY: Final = GatedAutoRouterCapability( key="tier_or_classifier_prompt", @@ -286,9 +293,9 @@ CUSTOMIZATION_CAPABILITY: Final = GatedAutoRouterCapability( f"({{config}} ->> 'classifier_type' IN ({_LLM_CLASSIFIER_TYPES_SQL}) AND (" "{config} -> 'classifier_llm_config' ->> 'system_prompt' IS NOT NULL OR " f"{_OPERATOR_PROMPT_FIELDS_SQL})) OR " - "({config} ->> 'classifier_type' = 'jev' AND " - "jsonb_typeof({config} -> 'jev_classifier_config' -> 'instructions') = 'string' AND " - f"{{config}} -> 'jev_classifier_config' ->> 'instructions' <> '{_DEFAULT_JEV_INSTRUCTIONS_SQL}')" + "({config} ->> 'classifier_type' IN ('oss_classifier', 'jev') AND " + f"jsonb_typeof({_OPENSOURCE_CLASSIFIER_CONFIG_SQL} -> 'instructions') = 'string' AND " + f"{_OPENSOURCE_CLASSIFIER_CONFIG_SQL} ->> 'instructions' <> '{_DEFAULT_JEV_INSTRUCTIONS_SQL}')" ), ) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 6a548e7a82d..3985a232f62 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -72622,6 +72622,45 @@ "supports_audio_input": true, "supports_video_input": true }, + "laya/english": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, + "laya/multilingual": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, + "laya/typed-decisions": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, "typesafe/jev-1.13.0": { "input_cost_per_token": 4.2e-08, "litellm_provider": "typesafe", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 9cbd326277e..d18f8d2e6d1 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -6,6 +6,7 @@ "url": "Link to provider documentation", "endpoints": { "chat_completions": "Supports /chat/completions endpoint", + "systemone": "Supports native System One typed decisions", "messages": "Supports /messages endpoint (Anthropic format)", "responses": "Supports /responses endpoint (OpenAI/Anthropic unified)", "embeddings": "Supports /embeddings endpoint", @@ -1476,6 +1477,13 @@ "rerank": false } }, + "laya": { + "display_name": "Laya (`laya`)", + "url": "https://docs.litellm.ai/docs/auto_router/decision_classifiers", + "endpoints": { + "systemone": true + } + }, "lambda_ai": { "display_name": "Lambda AI (`lambda_ai`)", "url": "https://docs.litellm.ai/docs/providers/lambda_ai", @@ -3354,6 +3362,13 @@ "provider_json_field": "skills", "url": "https://docs.litellm.ai/docs/skills" }, + "systemone": { + "docs_label": "systemone", + "display_name": "System One Decision API", + "leftnav_label": "/laya/v1/systemone", + "provider_json_field": "systemone", + "url": "https://docs.litellm.ai/docs/auto_router/decision_classifiers" + }, "text_completion": { "docs_label": "text_completion", "display_name": "OpenAI Completions API", diff --git a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py index 6c57e59f7e3..7fb23223845 100644 --- a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py +++ b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py @@ -63,7 +63,8 @@ def mock_request(): self.method = method self.request_body = request_body or {} # Add url attribute that the actual code expects - self.url = "http://localhost:8000/test" + self.url = httpx.URL("http://localhost:8000/test") + self.scope = {"type": "http", "method": method, "path": "/test"} # Add state attribute that FastAPI requests have self.state = type("State", (), {})() @@ -414,6 +415,7 @@ PROTOCOL_CONSTRAINED_PASS_THROUGH_ROUTES = { "/transcribe": {"POST"}, "/transcribe/{operation}": {"POST"}, "/tinyfish/{endpoint:path}": {"GET", "POST"}, + "/laya/v1/systemone": {"POST"}, } diff --git a/tests/unit/llms/laya/__init__.py b/tests/unit/llms/laya/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/tests/unit/llms/laya/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/unit/llms/laya/test_common_utils.py b/tests/unit/llms/laya/test_common_utils.py new file mode 100644 index 00000000000..c9ee0062cd2 --- /dev/null +++ b/tests/unit/llms/laya/test_common_utils.py @@ -0,0 +1,60 @@ +from collections.abc import Mapping +from typing import Final + +import pytest + +from litellm.llms.laya.common_utils import laya_connection, laya_response_model + + +@pytest.mark.parametrize( + ("base", "key", "expected_base", "expected_key"), + [ + (None, None, "http://laya.test/root", "laya-env-key"), + ("http://custom.test/", None, "http://custom.test", None), + ("http://custom.test/", "explicit-key", "http://custom.test", "explicit-key"), + ], +) +def test_laya_credentials_stay_with_their_configured_destination( + monkeypatch: pytest.MonkeyPatch, + base: str | None, + key: str | None, + expected_base: str, + expected_key: str | None, +) -> None: + monkeypatch.setenv("LAYA_API_BASE", "http://laya.test/root/") + monkeypatch.setenv("LAYA_API_KEY", "laya-env-key") + monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-this") + connection: Final = laya_connection(base, key) + assert (connection.api_base, connection.api_key) == (expected_base, expected_key) + assert "key" not in repr(connection) + + +@pytest.mark.parametrize( + "base", + ["", "ftp://laya.test", "http://user:password@laya.test", "https://laya.test?key=x", "http://laya.test/#x"], +) +def test_laya_rejects_ambiguous_server_urls(base: str) -> None: + with pytest.raises(ValueError, match="Laya"): + laya_connection(base) + + +def test_laya_missing_server_does_not_fall_back_to_typesafe(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("LAYA_API_BASE", raising=False) + monkeypatch.setenv("TYPESAFE_API_BASE", "https://typesafe.test") + with pytest.raises(ValueError, match="LAYA_API_BASE"): + laya_connection() + + +@pytest.mark.parametrize( + ("routing", "requested", "expected"), + [ + ({"model": "multilingual"}, "english", "multilingual"), + (None, "english", "english"), + ({"model": 42}, "english", "english"), + (None, None, "unknown"), + ], +) +def test_laya_identity_tracks_the_checkpoint_not_the_shared_agent_name( + routing: Mapping[str, object] | None, requested: str | None, expected: str +) -> None: + assert laya_response_model({"model": "laya-rl-agent", "routing": routing}, requested) == expected diff --git a/tests/unit/proxy/auth/test_auth_utils.py b/tests/unit/proxy/auth/test_auth_utils.py index 83ac56c4c85..6cf4456a0ff 100644 --- a/tests/unit/proxy/auth/test_auth_utils.py +++ b/tests/unit/proxy/auth/test_auth_utils.py @@ -8,7 +8,7 @@ from typing import Optional from unittest.mock import MagicMock, patch import pytest -from fastapi import Request +from fastapi import HTTPException, Request from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( @@ -463,6 +463,24 @@ def test_get_model_from_request_no_request_extracts_model(): ) +@pytest.mark.parametrize("model", ["english", "multilingual", "typed-decisions"]) +@pytest.mark.parametrize("route", ["/laya/v1/systemone", "/laya/v1/systemone/"]) +def test_laya_native_model_uses_the_classifier_permission_identity(model: str, route: str) -> None: + assert get_model_from_request(request_data={"model": model}, route=route) == f"laya/{model}" + + +@pytest.mark.parametrize("model", [None, "", "auto", "laya/english", "unknown", ["english"], 7]) +def test_laya_native_model_cannot_implicitly_select_an_unauthorized_checkpoint(model: object) -> None: + with pytest.raises(HTTPException) as denied: + get_model_from_request(request_data={"model": model}, route="/laya/v1/systemone") + assert denied.value.status_code == 400 + + +def test_laya_model_normalization_does_not_change_other_provider_routes() -> None: + assert get_model_from_request(request_data={"model": "jev-latest"}, route="/typesafe/v1/systemone") == "jev-latest" + assert get_model_from_request(request_data={}, route="/laya/health") is None + + def _cache_prediction_router(): from litellm.router import Router diff --git a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py index 5f7807650e1..7eda03c560b 100644 --- a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py @@ -8,6 +8,8 @@ from typing import Dict, Final, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest +from fastapi import HTTPException +from fastapi.encoders import jsonable_encoder from fastapi.testclient import TestClient from litellm._uuid import uuid @@ -7515,6 +7517,96 @@ class TestTeamMemberAutoRouterWrites: "model_info": {"id": "allowed-id"}, }]) + @staticmethod + def _classifier_config(classifier: Mapping[str, object], legacy: bool) -> Mapping[str, object]: + return { + "classifier_type": "jev" if legacy else "oss_classifier", + "tiers": {"SIMPLE": "allowed"}, + "jev_classifier_config" if legacy else "opensource_classifier_config": classifier, + } + + @pytest.mark.asyncio + @pytest.mark.parametrize("team_id", [None, "member-team"]) + @pytest.mark.parametrize( + "legacy,provider,model", + [(True, "typesafe", "jev-latest"), (False, "jev", "jev-latest"), (True, "laya", "english"), (False, "laya", "english")], + ) + async def test_classifier_create_stores_only_canonical_configuration( + self, team_id: str | None, legacy: bool, provider: str, model: str + ) -> None: + from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model + + row: Final = self._row() + database: Final = self._database(self._team(), row) + classifier: Final = { + "provider": provider, "model": model, + "api_base": "https://decision.test", "api_key": "stored-secret", + } + deployment: Final = Deployment( + model_name="new-classifier-router", + litellm_params=LiteLLM_Params( + model="auto_router/complexity_router", + complexity_router_config=self._classifier_config(classifier, legacy), + ), + model_info=ModelInfo(id=row.model_id, team_id=team_id), + ) + actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + with ( + self._environment(database, row), + patch("litellm.proxy.proxy_server.proxy_config.add_deployment", new=AsyncMock(return_value=ReconcileOutcome( # test-quality-ok: [TQ008] model reload I/O boundary + still_desired=frozenset((row.model_id,)), live_after=frozenset((row.model_id,)) + ))), + patch("litellm.proxy.management_endpoints.model_management_endpoints.append_team_models", new=AsyncMock()), # test-quality-ok: [TQ008] team allowlist persistence boundary + ): + await add_new_model(deployment, actor) + written: Final = database.db.litellm_proxymodeltable.create.await_args.kwargs["data"] + saved: Final = json.loads(written["litellm_params"])["complexity_router_config"] + assert saved == { + "classifier_type": "oss_classifier", + "tiers": {"SIMPLE": "allowed"}, + "opensource_classifier_config": {**classifier, "provider": "laya" if provider == "laya" else "jev"}, + } + + @pytest.mark.asyncio + @pytest.mark.parametrize("endpoint", ["create", "patch", "legacy"]) + @pytest.mark.parametrize("legacy_config", [None, {"provider": "laya", "model": "english"}]) + async def test_ambiguous_classifier_blocks_are_rejected_before_persistence( + self, endpoint: str, legacy_config: Mapping[str, object] | None + ) -> None: + from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model + + row: Final = self._row() + database: Final = self._database(self._team(), row) + config: Final = { + **self._classifier_config({"provider": "laya", "model": "english"}, False), + "jev_classifier_config": legacy_config, + } + actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + request: Final = updateDeployment( + litellm_params=updateLiteLLMParams(complexity_router_config=config), model_info=ModelInfo(id=row.model_id), + ) + operation: Final = ( + add_new_model( + Deployment( + model_name="ambiguous-classifier-router", + litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=config), + model_info=ModelInfo(id=row.model_id), + ), + actor, + ) + if endpoint == "create" + else patch_model(row.model_id, request, actor) + if endpoint == "patch" + else update_model(request, actor) + ) + with self._environment(database, row), pytest.raises(ProxyException) as denied: + await operation + assert denied.value.code == "400" + assert "opensource_classifier_config" in denied.value.message + assert "jev_classifier_config" in denied.value.message + database.db.litellm_proxymodeltable.create.assert_not_awaited() + database.db.litellm_proxymodeltable.update.assert_not_awaited() + @pytest.mark.asyncio @pytest.mark.parametrize("endpoint,change", [("patch", "config"), ("legacy", "strategy"), ("patch", "unrelated")]) async def test_admin_router_changes_release_member_scope(self, endpoint: str, change: str) -> None: @@ -7546,15 +7638,16 @@ class TestTeamMemberAutoRouterWrites: @pytest.mark.asyncio @pytest.mark.parametrize("endpoint", ["patch", "legacy"]) + @pytest.mark.parametrize("stored_legacy,supplied_legacy", [(True, True), (True, False), (False, True), (False, False)]) @pytest.mark.parametrize("change", ["save", "rotate", "move", "move-without-key", "reset", "heuristic"]) - async def test_jev_dashboard_save_preserves_server_transport(self, endpoint: str, change: str) -> None: + async def test_jev_dashboard_save_preserves_server_transport( + self, endpoint: str, change: str, stored_legacy: bool, supplied_legacy: bool + ) -> None: original: Final = self._row() transport: Final = {"api_key": "synthetic-original-jev-key", "api_base": "https://jev.example.com"} - stored_config: Final = { - "classifier_type": "jev", - "tiers": {"SIMPLE": "allowed"}, - "jev_classifier_config": {**transport, "instructions": "Old instructions", "timeout_ms": 6100}, - } + stored_config: Final = self._classifier_config( + {**transport, "instructions": "Old instructions", "timeout_ms": 6100}, stored_legacy + ) row: Final = original.model_copy( update={ "litellm_params": { @@ -7572,11 +7665,11 @@ class TestTeamMemberAutoRouterWrites: "reset": {"api_key": None, "api_base": None}, "heuristic": {}, }[change] - config: Final = { - "tiers": {"SIMPLE": "allowed"}, - "classifier_type": "heuristic" if change == "heuristic" else "jev", - **({} if change == "heuristic" else {"jev_classifier_config": {"timeout_ms": 8100, **overrides}}), - } + config: Final = ( + {"tiers": {"SIMPLE": "allowed"}, "classifier_type": "heuristic"} + if change == "heuristic" + else self._classifier_config({"timeout_ms": 8100, **overrides}, supplied_legacy) + ) request: Final = updateDeployment( litellm_params=updateLiteLLMParams(complexity_router_config=config), model_info=ModelInfo(id=row.model_id), @@ -7597,12 +7690,179 @@ class TestTeamMemberAutoRouterWrites: expected: Final = ( config if change == "heuristic" - else {**config, "jev_classifier_config": {**transport, "timeout_ms": 8100, **overrides}} + else { + "classifier_type": "oss_classifier", + "tiers": {"SIMPLE": "allowed"}, + "opensource_classifier_config": {**transport, "timeout_ms": 8100, **overrides}, + } ) assert saved == expected assert row.litellm_params["complexity_router_config"] == stored_config assert request.litellm_params.complexity_router_config == config + @pytest.mark.asyncio + @pytest.mark.parametrize("endpoint", ["patch", "legacy"]) + @pytest.mark.parametrize("stored_legacy,supplied_legacy", [(True, True), (True, False), (False, True), (False, False)]) + @pytest.mark.parametrize( + "stored_provider,stored_base,supplied,expected_transport", + [ + ("laya", "https://decision.test", {"provider": "laya", "model": "english"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}), + ("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": "https://decision.test"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}), + ( + "laya", + "https://decision.test", + {"provider": "laya", "model": "english", "api_key": None}, + {"api_base": "https://decision.test"}, + ), + ("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": "https://new.test"}, {}), + ("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": None}, {}), + ("laya", None, {"provider": "laya", "model": "english", "api_base": None}, {}), + ("laya", "https://decision.test", {"provider": "typesafe", "model": "jev-latest"}, {}), + ( + "laya", "https://decision.test", {"model": "english", "timeout_ms": 8100}, + {"provider": "laya", "api_base": "https://decision.test", "api_key": "stored-secret"}, + ), + ("typesafe", "https://decision.test", {"provider": "laya", "model": "english"}, {}), + ( + "typesafe", "https://decision.test", {"provider": "jev", "model": "jev-latest"}, + {"api_base": "https://decision.test", "api_key": "stored-secret"}, + ), + ( + "jev", "https://decision.test", {"provider": "typesafe", "model": "jev-latest"}, + {"api_base": "https://decision.test", "api_key": "stored-secret"}, + ), + ], + ) + async def test_decision_provider_changes_cannot_reuse_a_stored_key( + self, endpoint: str, stored_provider: str, stored_base: str | None, + supplied: Mapping[str, object], expected_transport: Mapping[str, object], + stored_legacy: bool, supplied_legacy: bool, + ) -> None: + original: Final = self._row() + row: Final = original.model_copy(update={"litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": self._classifier_config( + { + "provider": stored_provider, "model": "english" if stored_provider == "laya" else "jev-latest", + "api_base": stored_base, "api_key": "stored-secret", + }, + stored_legacy, + ), + }}) + database: Final = self._database(self._team(), row) + config: Final = self._classifier_config(supplied, supplied_legacy) + request: Final = updateDeployment( + litellm_params=updateLiteLLMParams(complexity_router_config=config), model_info=ModelInfo(id=row.model_id), + ) + actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + with self._environment(database, row): + await (patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor)) + written: Final = database.db.litellm_proxymodeltable.update.await_args.kwargs["data"] + saved: Final = json.loads(written["litellm_params"])["complexity_router_config"] + expected_provider: Final = supplied.get("provider", stored_provider) + assert saved == { + "classifier_type": "oss_classifier", + "tiers": {"SIMPLE": "allowed"}, + "opensource_classifier_config": { + **expected_transport, **supplied, + "provider": "jev" if expected_provider == "typesafe" else expected_provider, + }, + } + + @pytest.mark.asyncio + @pytest.mark.parametrize("endpoint", ["patch", "legacy"]) + @pytest.mark.parametrize( + "string_params,reset_field,config_shape", + [ + (False, None, "full"), (True, None, "full"), (False, "api_key", "full"), + (False, "api_base", "full"), (False, None, "omit-provider"), + (False, None, "omit-config"), (False, None, "null-config"), + ], + ) + async def test_member_save_protects_stored_classifier_connection( + self, endpoint: str, string_params: bool, reset_field: str | None, config_shape: str + ) -> None: + original: Final = self._row() + config: Final = { + "classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, + "jev_classifier_config": {"provider": "laya", "model": "english"}, + } + secret_params: Final = { + "model": "auto_router/complexity_router", + "complexity_router_config": { + **config, "jev_classifier_config": { + **config["jev_classifier_config"], "api_key": "retained-laya-secret", "api_base": "https://laya.test", + }, + }, + } + row: Final = original.model_copy(update={"litellm_params": secret_params}) + team: Final = self._team().model_copy(update={"models": ["allowed", "laya/english"]}) + database: Final = self._database(team, row) + database.transaction.litellm_proxymodeltable.update.return_value = row.model_copy( + update={"litellm_params": json.dumps(secret_params) if string_params else secret_params} + ) + supplied_config: Final = { + **config, "jev_classifier_config": { + **{ + key: value for key, value in config["jev_classifier_config"].items() + if key != "provider" or config_shape != "omit-provider" + }, + **({reset_field: None} if reset_field is not None else {}), + }, + } + patch_params: Final = ( + {"complexity_router_default_model": "allowed"} + if config_shape == "omit-config" + else {"complexity_router_config": None, "complexity_router_default_model": "allowed"} + if config_shape == "null-config" + else {"complexity_router_config": supplied_config} + ) + request: Final = updateDeployment( + litellm_params=updateLiteLLMParams.model_validate(patch_params), + model_info=ModelInfo(id=row.model_id, team_id="member-team"), + ) + actor: Final = UserAPIKeyAuth( + user_id="owner", user_role=LitellmUserRoles.INTERNAL_USER, models=["allowed", "laya/english"], config={"timeout": 60}, + ) + with self._environment(database, row): + if reset_field is not None: + expected_error: Final = HTTPException if endpoint == "patch" else ProxyException + with pytest.raises(expected_error, match="Team members cannot change classifier connections") as denied: + await ( + patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor) + ) + assert ( + denied.value.status_code if isinstance(denied.value, HTTPException) else int(denied.value.code) + ) == 403 + database.transaction.litellm_proxymodeltable.update.assert_not_awaited() + assert row.litellm_params == secret_params + return + response: Final = await (patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor)) + written: Final = database.transaction.litellm_proxymodeltable.update.await_args.kwargs["data"] + saved_config: Final = json.loads(written["litellm_params"])["complexity_router_config"] + untouched: Final = config_shape in ("omit-config", "null-config") + saved: Final = saved_config["jev_classifier_config" if untouched else "opensource_classifier_config"] + assert saved == secret_params["complexity_router_config"]["jev_classifier_config"] + assert saved_config["classifier_type"] == ("jev" if untouched else "oss_classifier") + if untouched: + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper + + assert decrypt_value_helper( + json.loads(written["litellm_params"])["complexity_router_default_model"], + key="complexity_router_default_model", return_original_value=True, + ) == "allowed" + response_payload: Final = jsonable_encoder(response) + assert "retained-laya-secret" not in json.dumps(response_payload) + response_params: Final = json.loads(response_payload["litellm_params"]) if string_params else response_payload["litellm_params"] + assert response_params == { + **secret_params, "complexity_router_config": { + **config, "jev_classifier_config": { + **config["jev_classifier_config"], "api_key": "REDACTED", "api_base": "https://laya.test", + }, + }, + } + assert "retained-laya-secret" in row.model_dump_json() + @pytest.mark.asyncio @pytest.mark.parametrize("endpoint", ["patch", "legacy"]) @pytest.mark.parametrize("access", ["owner", "peer", "limited-key"]) diff --git a/tests/unit/proxy/management_helpers/test_auto_router_permissions.py b/tests/unit/proxy/management_helpers/test_auto_router_permissions.py index e16271a5189..b60fd4ac7ad 100644 --- a/tests/unit/proxy/management_helpers/test_auto_router_permissions.py +++ b/tests/unit/proxy/management_helpers/test_auto_router_permissions.py @@ -1,3 +1,4 @@ +import json from collections.abc import Mapping from dataclasses import dataclass from typing import Final @@ -137,33 +138,48 @@ def test_tier_config_is_normalized_and_unknown_router_extras_are_rejected() -> N @pytest.mark.parametrize( ("jev_override", "rejected_at"), [ - ({"api_base": "https://collector.invalid"}, "jev_classifier_config"), + ({"api_base": "https://collector.invalid"}, "opensource_classifier_config"), ({"api_key": "sk-member"}, "api_key"), ({"api_base": "https://collector.invalid", "api_key": "sk-member"}, "api_key"), - ({"api_base": "https://collector.invalid", "api_key": ""}, "jev_classifier_config.api_key"), + ({"api_base": "https://collector.invalid", "api_key": ""}, "opensource_classifier_config.api_key"), + ({"provider": "laya", "model": "english", "api_base": "https://collector.invalid"}, "api_base"), + ({"provider": "laya", "model": "english", "api_key": "sk-member"}, "api_key"), ], ) +@pytest.mark.parametrize("legacy", [False, True]) def test_members_cannot_move_the_jev_classifier_off_the_proxys_typesafe_account( - jev_override: Mapping[str, str], rejected_at: str + jev_override: Mapping[str, str], rejected_at: str, legacy: bool ) -> None: with pytest.raises(HTTPException) as denied: validate_member_auto_router_config( - {"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": jev_override} + { + "tiers": {"SIMPLE": "allowed"}, + "classifier_type": "jev" if legacy else "oss_classifier", + "jev_classifier_config" if legacy else "opensource_classifier_config": jev_override, + } ) assert denied.value.status_code == 400 assert denied.value.detail == f"Invalid member auto-router configuration at {rejected_at}." -def test_members_can_still_tune_the_jev_classifier() -> None: +@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-preview"), ("laya", "english")]) +@pytest.mark.parametrize("legacy", [False, True]) +def test_members_can_still_tune_the_jev_classifier(provider: str, model: str, legacy: bool) -> None: validated: Final = validate_member_auto_router_config( { "tiers": {"SIMPLE": "allowed"}, - "classifier_type": "jev", - "jev_classifier_config": {"model": "jev-preview", "timeout_ms": 500}, + "classifier_type": "jev" if legacy else "oss_classifier", + "jev_classifier_config" if legacy else "opensource_classifier_config": { + "provider": provider, "model": model, "timeout_ms": 500, + }, } ) assert validated.jev_classifier_config is not None - assert (validated.jev_classifier_config.model, validated.jev_classifier_config.timeout_ms) == ("jev-preview", 500) + assert ( + validated.jev_classifier_config.provider, + validated.jev_classifier_config.model, + validated.jev_classifier_config.timeout_ms, + ) == ("jev" if provider == "typesafe" else provider, model, 500) assert validate_member_auto_router_config(validated.model_dump()).jev_classifier_config is not None @@ -217,6 +233,92 @@ async def test_member_updates_restrict_fields_and_preserve_an_inherited_default( assert granted.default_model == "allowed" +@pytest.mark.asyncio +@pytest.mark.parametrize( + "nested,expected_identity,restricted", + [ + ("omit-config", "laya/english", False), + ("omit-config", "laya/english", True), + ("omit-block", None, False), + (None, None, False), + ({}, None, False), + ({"timeout_ms": 500}, None, False), + ({"model": "english", "timeout_ms": 500}, "laya/english", False), + ({"model": "english", "timeout_ms": 500}, "laya/english", True), + ({"model": "multilingual"}, "laya/multilingual", False), + ({"provider": "typesafe", "model": "jev-latest"}, "typesafe/jev-latest", False), + ({"provider": "typesafe", "model": "jev-latest"}, "typesafe/jev-latest", True), + ], +) +async def test_member_authorization_and_persistence_resolve_the_same_classifier( + catalog: Router, monkeypatch: pytest.MonkeyPatch, nested: object, expected_identity: str | None, restricted: bool +) -> None: + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + update_db_model, + ) + from litellm.types.management_endpoints.auto_router_endpoints import RequestComplexityRouterConfig + + monkeypatch.setenv("LITELLM_SALT_KEY", "member-router-test-salt") + stored_config: Final = { + "classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, + "jev_classifier_config": { + "provider": "laya", "model": "english", "timeout_ms": 12000, + "api_base": "https://laya.test", "api_key": "stored-classifier-key", + }, + } + existing: Final = Deployment( + model_name="member-router", + litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=stored_config), + model_info=ModelInfo(id="router-a", team_id="team-a"), created_by="owner", + ) + incoming_config: Final = ( + None if nested == "omit-config" else { + "classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, + **({} if nested == "omit-block" else {"jev_classifier_config": nested}), + } + ) + patch: Final = updateDeployment.model_validate({"litellm_params": { + "complexity_router_config": incoming_config, "complexity_router_default_model": "allowed", + }}) + operation: Final = authorize_member_auto_router_write( + incoming=patch, existing=existing, user_api_key_dict=_actor( + models=["allowed"] if restricted or expected_identity is None else ["allowed", expected_identity], + ), + team=_team(models=["allowed", "laya/english", "laya/multilingual", "typesafe/jev-latest"]), + premium_user=True, prisma_client=_Client(), llm_router=catalog, + ) + violation: Final = _strategy_router_write_violation(patch.litellm_params, existing.litellm_params) + if expected_identity is None: + assert violation is not None + with pytest.raises(HTTPException) as rejected: + await operation + assert rejected.value.status_code == 400 + return + assert violation is None + if restricted: + with pytest.raises(ProxyException, match=expected_identity): + await operation + return + grant: Final = await operation + persisted: Final = update_db_model(existing, patch) + saved: Final = RequestComplexityRouterConfig.model_validate( + json.loads(persisted["litellm_params"])["complexity_router_config"] + ) + assert grant.config == saved + assert saved.jev_classifier_config is not None + assert ( + "typesafe" if saved.jev_classifier_config.provider == "jev" else saved.jev_classifier_config.provider + ) + f"/{saved.jev_classifier_config.model}" == expected_identity + assert saved.jev_classifier_config.api_key == ( + "stored-classifier-key" if expected_identity.startswith("laya/") else None + ) + assert saved.jev_classifier_config.timeout_ms == ( + 12000 if nested == "omit-config" else 500 if nested == {"model": "english", "timeout_ms": 500} else 3000 + ) + assert existing.litellm_params.complexity_router_config == stored_config + + @pytest.mark.asyncio @pytest.mark.parametrize("target", ["missing", "nested"]) async def test_member_dependencies_require_plain_configured_models(target: str) -> None: @@ -246,13 +348,17 @@ async def test_member_dependencies_require_plain_configured_models(target: str) @pytest.mark.asyncio @pytest.mark.parametrize("restricted", ["key", "team", None]) +@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")]) async def test_jev_evaluation_requires_model_access_but_no_completion_deployment( - catalog: Router, restricted: str | None + catalog: Router, restricted: str | None, provider: str, model: str ) -> None: - permitted: Final = ["allowed", "typesafe/jev-latest"] + permitted: Final = ["allowed", f"{provider}/{model}"] operation: Final = authorize_member_auto_router_dependencies( config=validate_member_auto_router_config( - {"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": {}} + { + "tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", + "jev_classifier_config": {"provider": provider, "model": model}, + } ), default_model=None, user_api_key_dict=_actor(models=["allowed"] if restricted == "key" else permitted), @@ -261,17 +367,20 @@ async def test_jev_evaluation_requires_model_access_but_no_completion_deployment llm_router=catalog, ) if restricted is not None: - with pytest.raises(ProxyException, match="jev-latest"): + with pytest.raises(ProxyException, match=model): await operation return await operation - assert not catalog.get_model_list("typesafe/jev-latest") + assert not catalog.get_model_list(f"{provider}/{model}") @pytest.mark.asyncio @pytest.mark.parametrize("restricted", ["member", "project", "organization", None]) -async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restricted: str | None) -> None: - allowed: Final = ["allowed", "typesafe/jev-latest"] +@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")]) +async def test_jev_evaluation_obeys_each_containing_scope( + catalog: Router, restricted: str | None, provider: str, model: str +) -> None: + allowed: Final = ["allowed", f"{provider}/{model}"] membership: Final = LiteLLM_TeamMembership.model_validate( { "user_id": "owner", @@ -293,7 +402,10 @@ async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restr ) operation: Final = authorize_member_auto_router_dependencies( config=validate_member_auto_router_config( - {"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": {}} + { + "tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", + "jev_classifier_config": {"provider": provider, "model": model}, + } ), default_model=None, user_api_key_dict=_actor(models=allowed, project_id="project-a"), @@ -303,8 +415,8 @@ async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restr dependency_objects=MemberAutoRouterDependencyObjects(membership, organization, project), ) if restricted is not None: - with pytest.raises(ProxyException, match="jev-latest"): + with pytest.raises(ProxyException, match=model): await operation return await operation - assert not catalog.get_model_list("typesafe/jev-latest") + assert not catalog.get_model_list(f"{provider}/{model}") diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py index e0a5ef063e8..acf05dcdfde 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py @@ -1,10 +1,12 @@ from datetime import datetime +from typing import Final from unittest.mock import MagicMock import httpx import pytest import litellm +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.proxy.pass_through_endpoints.llm_provider_handlers.typesafe_passthrough_logging_handler import ( TypeSafePassthroughLoggingHandler, ) @@ -137,6 +139,84 @@ def test_success_handler_dispatches_to_typesafe_handler(): assert normalized["kwargs"]["model"] == "typesafe/jev-1.13.0" +@pytest.mark.asyncio +@pytest.mark.parametrize("guardrail_cost", [0.0, 0.25]) +@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"]) +@pytest.mark.parametrize("routing_model", ["multilingual", None]) +async def test_laya_gateway_accounts_for_checkpoint_usage_and_registered_cost( + monkeypatch: pytest.MonkeyPatch, routing_model: str | None, metadata_slot: str, guardrail_cost: float +) -> None: + checkpoint: Final = routing_model or "english" + model: Final = f"laya/{checkpoint}" + input_rate: Final = 0.002 + output_rate: Final = 0.005 + monkeypatch.setitem(litellm.model_cost, model, { + "input_cost_per_token": input_rate, "output_cost_per_token": output_rate, + "litellm_provider": "laya", "mode": "evaluation", + }) + start: Final = datetime.now() + logging_obj: Final = Logging( + model="english", messages=[], stream=False, call_type="pass_through_endpoint", + start_time=start, litellm_call_id="laya-accounting", function_id="laya-accounting", kwargs={}, + ) + from fastapi import Request + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import HttpPassThroughEndpointHelpers + + request: Final = Request({ + "type": "http", "method": "POST", "path": "/laya/v1/systemone", + "headers": [], "query_string": b"", + }) + auth: Final = UserAPIKeyAuth( + api_key="laya-budget-key", token="laya-budget-key", + model_max_budget={"laya/english": {"budget_limit": 0.01, "time_period": "1d"}}, + ) + request_body: Final = {"model": "english", metadata_slot: {"model_group": "unbounded-client-choice"}} + logging_kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=request, user_api_key_dict=auth, logging_obj=logging_obj, + passthrough_logging_payload={"url": "https://laya.test/v1/systemone"}, _parsed_body=request_body, + ) + logging_kwargs["litellm_params"]["metadata"]["standard_logging_guardrail_information"] = [ + {"guardrail_name": "trusted-hook", "guardrail_cost": guardrail_cost}, + ] + logging_obj.update_environment_variables( + model="english", user="unknown", optional_params={}, + litellm_params=logging_kwargs["litellm_params"], call_type="pass_through_endpoint", + ) + body: Final = { + "model": "laya-rl-agent", "usage": {"input_tokens": 10, "output_tokens": 3}, + **({"routing": {"model": routing_model}} if routing_model else {}), + } + normalized: Final = PassThroughEndpointLogging().normalize_llm_passthrough_logging_payload( + httpx_response=httpx.Response(200, request=httpx.Request("POST", "https://laya.test/v1/systemone"), json=body), + response_body=body, request_body={"model": "english"}, logging_obj=logging_obj, + url_route="https://laya.test/v1/systemone", result="{}", start_time=start, + end_time=datetime.now(), cache_hit=False, custom_llm_provider="laya", **logging_kwargs, + ) + logged: Final = normalized["kwargs"] + expected_cost: Final = 10 * input_rate + 3 * output_rate + assert (logged["model"], logged["custom_llm_provider"]) == (model, "laya") + assert logged["response_cost"] == pytest.approx(expected_cost) + assert logged["combined_usage_object"].model_dump(exclude_none=True) == { + "prompt_tokens": 10, "completion_tokens": 3, "total_tokens": 13, + } + assert logging_obj.model_call_details["model"] == model + assert logging_obj.model_call_details["response_cost"] == pytest.approx(expected_cost) + assert logged["standard_logging_object"]["model"] == model + assert logged["standard_logging_object"]["model_group"] == "laya/english" + assert logged["standard_logging_object"]["response_cost"] == pytest.approx(expected_cost + guardrail_cost) + + from litellm.caching.caching import DualCache + from litellm.exceptions import BudgetExceededError + from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter + + budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) + assert await budget_limiter.is_key_within_model_budget(auth, "laya/english") + await budget_limiter.async_log_success_event(logged, None, start, datetime.now()) + with pytest.raises(BudgetExceededError): + await budget_limiter.is_key_within_model_budget(auth, "laya/english") + + def test_openrouter_decisions_response_is_priced_from_request_model_registry_row(): logging_obj = _logging_obj() model_cost = litellm.model_cost["openrouter/typesafe/jev-1.13"] diff --git a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index a0772c7d4f3..22171f4ffb0 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -23,6 +23,8 @@ from starlette.datastructures import FormData import litellm +from litellm.caching.caching import DualCache +from litellm.types.utils import CallTypesLiteral from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS @@ -7407,6 +7409,152 @@ class TestTypeSafePassthroughRoute: ) +class TestLayaPassthroughRoute: + @pytest.fixture + def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: + from litellm.proxy.proxy_server import app + + monkeypatch.setenv("LAYA_API_BASE", "http://laya.test/base") + monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key") + monkeypatch.delenv("LAYA_API_KEY", raising=False) + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: UserAPIKeyAuth(api_key="sk-virtual")) + yield TestClient(app) + + @pytest.mark.parametrize("api_key", [None, "laya-provider-key"]) + def test_laya_forwards_native_decisions_without_gateway_or_typesafe_credentials( + self, client: TestClient, monkeypatch: pytest.MonkeyPatch, api_key: str | None + ) -> None: + if api_key is not None: + monkeypatch.setenv("LAYA_API_KEY", api_key) + body: Final = { + "model": "english", + "state": "refund", + "questions": {"department": {"type": "choice", "criteria": {"billing": "refunds"}}}, + } + answer: Final = {"model": "laya-rl-agent", "routing": {"model": "english"}, "answers": {}} + with respx.mock(assert_all_called=True) as upstream: + route: Final = upstream.post("http://laya.test/base/v1/systemone?trace=yes").respond(200, json=answer) + response: Final = client.post( + "/laya/v1/systemone?trace=yes", + json=body, + headers={"Authorization": "Bearer sk-virtual", "x-pass-authorization": "Bearer attacker"}, + ) + + assert (response.status_code, response.json()) == (200, answer) + sent: Final = route.calls.last.request + assert sent.headers.get("authorization") == (f"Bearer {api_key}" if api_key else None) + assert json.loads(sent.content) == body + + def test_laya_missing_server_fails_without_contacting_another_provider( + self, client: TestClient, monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.delenv("LAYA_API_BASE") + with respx.mock(assert_all_called=False) as upstream: + response: Final = client.post("/laya/v1/systemone", json={"model": "english"}) + assert response.status_code == 503 + assert "LAYA_API_BASE" in response.text + assert len(upstream.calls) == 0 + + def test_laya_does_not_forward_unsupported_endpoints(self, client: TestClient) -> None: + with respx.mock(assert_all_called=False) as upstream: + response: Final = client.post("/laya/v1/evaluate", json={"model": "english"}) + assert response.status_code == 404 + assert len(upstream.calls) == 0 + + @pytest.mark.parametrize("model", [None, "auto", "jev-latest"]) + def test_laya_rejects_implicit_checkpoint_selection(self, client: TestClient, model: str | None) -> None: + with respx.mock(assert_all_called=False) as upstream: + response: Final = client.post("/laya/v1/systemone", json={"model": model}) + assert response.status_code == 400 + assert len(upstream.calls) == 0 + + @pytest.mark.parametrize( + "controls", + [{"custom_body": {"model": "multilingual", "state": "refund"}}, {"stream": True}, {"stream": "true"}], + ) + def test_laya_rejects_controls_that_change_authorized_body_or_usage_accounting( + self, client: TestClient, controls: Mapping[str, object] + ) -> None: + with respx.mock(assert_all_called=False) as upstream: + route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}}) + response: Final = client.post("/laya/v1/systemone", json={"model": "english", **controls}) + assert response.status_code == 400 + assert not route.called + + + @pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"]) + def test_laya_hooks_enforce_canonical_model_limits_and_keep_native_wire_body( + self, client: TestClient, monkeypatch: pytest.MonkeyPatch, metadata_slot: str + ) -> None: + from litellm.integrations.custom_logger import CustomLogger + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.utils import InternalUsageCache + from litellm.proxy.proxy_server import app + + cache: Final = DualCache() + limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) + auth: Final = UserAPIKeyAuth( + api_key="laya-native-rpm", metadata={"model_rpm_limit": {"laya/english": 1}}, + ) + def authenticated_key() -> UserAPIKeyAuth: + return auth + + monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, authenticated_key) + + class LimitHook(CustomLogger): + async def async_pre_call_hook( + self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, + data: dict[str, object], call_type: CallTypesLiteral, + ) -> dict[str, object]: + assert data["model"] == "laya/english" + metadata: Final = data.get(metadata_slot) + assert isinstance(metadata, dict) + assert "standard_logging_guardrail_information" not in metadata + assert metadata["customer_label"] == "retained" + await limiter.async_pre_call_hook(user_api_key_dict, cache, data, call_type) + return data + + monkeypatch.setattr(litellm, "callbacks", [LimitHook()]) + body: Final = { + "model": "english", "state": "refund", + metadata_slot: { + "customer_label": "retained", "model_group": "unbounded-client-choice", + "standard_logging_guardrail_information": [{"guardrail_cost": 25.0}], + }, + } + with respx.mock(assert_all_called=True) as upstream: + route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}}) + first: Final = client.post("/laya/v1/systemone", json=body) + second: Final = client.post("/laya/v1/systemone", json=body) + assert first.status_code == 200, first.text + assert second.status_code == 429, second.text + assert route.call_count == 1 + assert json.loads(route.calls.last.request.content) == {"model": "english", "state": "refund"} + + def test_laya_preserves_trusted_hook_checkpoint_changes( + self, client: TestClient, monkeypatch: pytest.MonkeyPatch + ) -> None: + from litellm.integrations.custom_logger import CustomLogger + + class CheckpointHook(CustomLogger): + async def async_pre_call_hook( + self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, + data: dict[str, object], call_type: CallTypesLiteral, + ) -> dict[str, object]: + assert data["model"] == "laya/english" + return {**data, "model": "laya/multilingual"} + + monkeypatch.setattr(litellm, "callbacks", [CheckpointHook()]) + with respx.mock(assert_all_called=True) as upstream: + route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}}) + response: Final = client.post("/laya/v1/systemone", json={"model": "english", "state": "refund"}) + assert response.status_code == 200, response.text + assert json.loads(route.calls.last.request.content) == {"model": "multilingual", "state": "refund"} + + class TestFalAIPassthroughRoute: @pytest.fixture def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 52acdf93f35..72d76e54a5d 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -1470,7 +1470,7 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs(): # Create mock request mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/api/endpoint" + mock_request.url = httpx.URL("http://test-proxy.com/api/endpoint") mock_request.body = AsyncMock(return_value=b'{"message": "test request"}') mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -1575,7 +1575,7 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream(): mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/v1/messages" + mock_request.url = httpx.URL("http://test-proxy.com/v1/messages") mock_request.body = AsyncMock(return_value=b'{"model": "claude-3", "stream": true}') mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -1637,7 +1637,7 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream(): mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/v1/messages" + mock_request.url = httpx.URL("http://test-proxy.com/v1/messages") mock_request.body = AsyncMock(return_value=b'{"model": "claude-3"}') mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -2507,7 +2507,7 @@ async def test_pass_through_request_query_params_forwarding(): # Create mock request with query parameters (Azure API version) mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://localhost:4000/azure-assistant/openai/assistants" + mock_request.url = httpx.URL("http://localhost:4000/azure-assistant/openai/assistants") mock_request.body = AsyncMock(return_value=json.dumps(test_body).encode()) mock_request.headers = Headers({"Content-Type": "application/json"}) @@ -3016,7 +3016,7 @@ async def test_bedrock_router_passthrough_metadata_initialization(): # Create mock request with headers mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://localhost:4000/bedrock/model/my-model/invoke" + mock_request.url = httpx.URL("http://localhost:4000/bedrock/model/my-model/invoke") mock_request.headers = Headers( { "content-type": "application/json", @@ -3850,7 +3850,7 @@ def _lit3538_request(): r = MagicMock() r.method = "POST" r.query_params = {} - r.url = "http://testserver/mock/echo" + r.url = httpx.URL("http://testserver/mock/echo") r.state = SimpleNamespace() headers = MagicMock() headers.copy.return_value = {} @@ -3983,7 +3983,7 @@ async def test_pass_through_request_non_streaming_upstream_error_returned_unchan mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/mock-upstream/api/denied" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/denied") mock_request.body = AsyncMock(return_value=b'{"action": "read"}') mock_request.headers = Headers({"content-type": "application/json"}) mock_request.query_params = QueryParams({}) @@ -4069,7 +4069,7 @@ async def test_pass_through_request_upstream_error_failure_hook_exception_is_swa mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/mock-upstream/api/denied" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/denied") mock_request.body = AsyncMock(return_value=b'{"action": "read"}') mock_request.headers = Headers({"content-type": "application/json"}) mock_request.query_params = QueryParams({}) @@ -4118,7 +4118,7 @@ async def test_pass_through_request_streaming_upstream_error_returned_unchanged( mock_request = MagicMock(spec=Request) mock_request.method = "GET" - mock_request.url = "http://test-proxy.com/mock-upstream/api/stream-denied" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/stream-denied") mock_request.body = AsyncMock(return_value=b"") mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -4169,7 +4169,7 @@ class _UpstreamErrorBodyStream(httpx.AsyncByteStream): def _upstream_error_request() -> MagicMock: mock_request: Final = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent") mock_request.body = AsyncMock(return_value=b'{"contents": []}') mock_request.headers = Headers({"content-type": "application/json"}) mock_request.query_params = QueryParams({}) @@ -4966,7 +4966,7 @@ async def test_pass_through_request_non_streaming_success_unchanged(): mock_request = MagicMock(spec=Request) mock_request.method = "GET" - mock_request.url = "http://test-proxy.com/mock-upstream/api/success" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/success") mock_request.body = AsyncMock(return_value=b"") mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -5029,7 +5029,7 @@ async def test_pass_through_request_claims_the_budget_reservation_only_when_its_ mock_get_client.return_value = MagicMock(client=async_client) mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/mock-upstream/api/generate" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/generate") mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}') mock_request.headers = Headers({"content-type": "application/json"}) mock_request.query_params = QueryParams({}) @@ -5081,7 +5081,7 @@ async def test_pass_through_request_leaves_the_budget_reservation_for_the_reques mock_get_client.return_value = MagicMock(client=async_client) mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/mock-upstream/api/generate" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/generate") mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}') mock_request.headers = Headers({"content-type": "application/json"}) mock_request.query_params = QueryParams({}) @@ -5112,7 +5112,7 @@ async def test_pass_through_request_internal_failure_still_raises_proxy_exceptio mock_request = MagicMock(spec=Request) mock_request.method = "GET" - mock_request.url = "http://test-proxy.com/mock-upstream/api/success" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/success") mock_request.body = AsyncMock(return_value=b"") mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -5213,7 +5213,7 @@ def _enter_relay_logging_mocks(stack, parsed_body): def _relay_client_request(method="GET"): mock_request = MagicMock(spec=Request) mock_request.method = method - mock_request.url = "http://localhost:4000/passthrough-relay/results" + mock_request.url = httpx.URL("http://localhost:4000/passthrough-relay/results") mock_request.body = AsyncMock(return_value=b"") mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -6650,7 +6650,7 @@ def _passthrough_kwargs_for_reservation( ) -> dict: mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent") mock_request.headers = Headers({}) mock_request.scope = {"endpoint": _marked_pass_through_endpoint()} if user_defined_route else {} @@ -6797,7 +6797,7 @@ async def _drive_streaming_pass_through(upstream_content_type, chunk_delay_secon mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/v1/messages" + mock_request.url = httpx.URL("http://test-proxy.com/v1/messages") mock_request.body = AsyncMock( return_value=b'{"model": "claude-3", "stream": true}' if client_asked_for_stream @@ -6985,36 +6985,82 @@ def _marked_pass_through_endpoint(): return _endpoint -def test_user_defined_passthrough_is_neither_tracked_nor_enforced(): - """ - `get_model_from_request` returns None for a user-defined pass-through on - purpose: the body is forwarded verbatim, so its `model` names an UPSTREAM - model rather than a LiteLLM-managed one, and enforcing key/team allowlists - against it would reject valid requests. Enforcement is therefore skipped - on those routes. +@pytest.mark.asyncio +@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"]) +async def test_user_defined_passthrough_is_neither_tracked_nor_enforced(metadata_slot: str) -> None: + from datetime import datetime - Attaching the budget metadata anyway would charge a counter that nothing on - that route can refuse, and would attribute the spend to a budget the operator - scoped to a LiteLLM model that merely shares the name. Tracking and - enforcement have to agree: both on for the built-in provider routes, both off - here. - """ - kwargs = _passthrough_kwargs_for_reservation( - UserAPIKeyAuth( - token="hash", - user_id="u-1", - model_max_budget={"claude-opus-4-8": {"budget_limit": 1.0, "time_period": "18h"}}, - ), - user_defined_route=True, + from litellm.caching.caching import DualCache + from litellm.proxy.auth.auth_utils import get_model_from_request + from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter + + budget: Final = {"managed-model": {"budget_limit": 0.1, "time_period": "1d"}} + limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) + auth: Final = UserAPIKeyAuth( + api_key="custom-key", token="custom-key", team_id="shared-team", team_model_max_budget=budget, ) + endpoint: Final = create_pass_through_route( + endpoint="/custom-budget-test", target="https://upstream.test/echo", custom_headers={}, cost_per_request=0.25, + ) + request: Final = Request({ + "type": "http", "method": "POST", "path": "/custom-budget-test", "headers": [], + "query_string": b"", "endpoint": endpoint, + }) + body: Final = { + "model": "upstream-only-model", metadata_slot: { + "model_group": "managed-model", "customer_label": "retained", + "user_api_key_team_model_max_budget": budget, + }, + } + assert get_model_from_request(body, "/custom-budget-test", request=request) is None + assert await limiter.is_team_within_model_budget("shared-team", budget, None, "managed-model") + start: Final = datetime.now() + logging_obj: Final = LiteLLMLoggingObj( + model="upstream-only-model", messages=[], stream=False, call_type="pass_through_endpoint", + start_time=start, litellm_call_id="custom-budget", function_id="custom-budget", kwargs={}, + dynamic_async_success_callbacks=[limiter], + ) + payload: Final = { + "url": "https://upstream.test/echo", "request_body": body, "request_method": "POST", "cost_per_request": 0.25, + } + kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=request, user_api_key_dict=auth, passthrough_logging_payload=payload, logging_obj=logging_obj, + _parsed_body=body, litellm_call_id="custom-budget", + ) + logging_obj.update_environment_variables( + model="upstream-only-model", user="unknown", optional_params={}, + litellm_params=kwargs["litellm_params"], call_type="pass_through_endpoint", + ) + response: Final = httpx.Response( + 200, request=httpx.Request("POST", "https://upstream.test/echo"), json={"ok": True}, + ) + await PassThroughEndpointLogging().pass_through_async_success_handler( + httpx_response=response, response_body={"ok": True}, request_body=body, logging_obj=logging_obj, + url_route="https://upstream.test/echo", result=response.text, start_time=start, end_time=datetime.now(), + cache_hit=False, **kwargs, + ) + assert logging_obj.model_call_details["response_cost"] == 0.25 + assert await limiter.is_team_within_model_budget("shared-team", budget, None, "managed-model") + metadata: Final = kwargs["litellm_params"]["metadata"] + assert (metadata["model_group"], metadata["customer_label"]) == ("managed-model", "retained") + assert metadata.keys().isdisjoint({ + "user_api_key_model_max_budget", "user_api_key_team_model_max_budget", + "user_api_key_user_model_max_budget", "user_api_key_end_user_model_max_budget", + }) - metadata = kwargs["litellm_params"]["metadata"] - for field in ( - "user_api_key_model_max_budget", - "user_api_key_user_model_max_budget", - "user_api_key_end_user_model_max_budget", - ): - assert field not in metadata, f"{field} was attached on a route that never enforces it" + +@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"]) +def test_builtin_passthrough_pins_model_group_to_the_resolved_model(metadata_slot: str) -> None: + request: Final = Request({ + "type": "http", "method": "POST", "path": "/gemini/v1beta/models/gemini-2.5-flash:generateContent", + "headers": [], "query_string": b"", + }) + kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=request, user_api_key_dict=UserAPIKeyAuth(token="hash", user_id="u-1"), + passthrough_logging_payload=MagicMock(), logging_obj=MagicMock(), + _parsed_body={"contents": [], metadata_slot: {"model_group": "unbounded-client-choice"}}, + ) + assert kwargs["litellm_params"]["metadata"]["model_group"] == "gemini-2.5-flash" @pytest.mark.parametrize( @@ -7344,7 +7390,7 @@ def test_passthrough_client_cannot_forge_session_id_omission(client_metadata_key mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent") mock_request.headers = Headers({}) mock_request.scope = {} @@ -7377,7 +7423,7 @@ def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_bo the call to (LIT-1761: passthrough successes carried model_id="").""" mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent" + mock_request.url = httpx.URL("http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent") mock_request.headers = Headers({}) mock_request.scope = {} mock_request.state = SimpleNamespace( @@ -7409,7 +7455,7 @@ _PROXY_SERVER_REQUEST: Final = TypeAdapter(dict[str, object]) def _split_pass_through_body(body: str) -> _PassThroughSplit: mock_request: Final = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent") mock_request.headers = Headers() mock_request.scope = MappingProxyType({}) @@ -7665,7 +7711,7 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token() mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/anthropic/v1/messages" + mock_request.url = httpx.URL("http://0.0.0.0:4000/anthropic/v1/messages") mock_request.headers = Headers({}) mock_request.scope = {} session = UserAPIKeyAuth( diff --git a/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py b/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py index 73927e92c15..09987b2781c 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py +++ b/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py @@ -18,7 +18,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest -from fastapi import HTTPException +from fastapi import HTTPException, Request pytest.importorskip("opentelemetry") @@ -81,15 +81,16 @@ def _user_api_key_dict(): return d -def _mock_request(): - r = MagicMock() - r.method = "POST" - r.query_params = {} - r.url = "http://testserver/mock/echo" - headers = MagicMock() - headers.copy.return_value = {} - r.headers = headers - return r +def _mock_request() -> Request: + return Request({ + "type": "http", + "method": "POST", + "scheme": "http", + "server": ("testserver", 80), + "path": "/mock/echo", + "headers": [], + "query_string": b"", + }) def _httpx_response(text: str) -> httpx.Response: diff --git a/tests/unit/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py b/tests/unit/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py index 29a635e9b27..d6b69c7c010 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py +++ b/tests/unit/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py @@ -1,8 +1,8 @@ +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import Request -from starlette.datastructures import Headers, State from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( @@ -771,12 +771,15 @@ async def test_vertex_passthrough_attributes_the_call_to_the_resolved_deployment """The router deployment that rewrote the upstream URL is the one the logging kwargs must name, so the Prometheus model_id label (and SpendLogs.model_id) on a Vertex passthrough success reads the deployment's id instead of "" (LIT-1761).""" - mock_request = MagicMock(spec=Request) - mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent" - mock_request.headers = Headers({}) - mock_request.scope = {} - mock_request.state = State() + mock_request: Final = Request({ + "type": "http", + "method": "POST", + "scheme": "http", + "server": ("0.0.0.0", 4000), + "path": "/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent", + "headers": [], + "query_string": b"", + }) mock_handler = MagicMock() mock_handler.get_default_base_target_url.return_value = "https://aiplatform.googleapis.com" diff --git a/tests/unit/router_strategy/complexity_router/test_jev_classifier.py b/tests/unit/router_strategy/complexity_router/test_jev_classifier.py index 45070dfd3a7..affcfdc789c 100644 --- a/tests/unit/router_strategy/complexity_router/test_jev_classifier.py +++ b/tests/unit/router_strategy/complexity_router/test_jev_classifier.py @@ -8,6 +8,7 @@ from unittest.mock import create_autospec import httpx import pytest +import respx import litellm from litellm._logging import verbose_router_logger @@ -30,14 +31,15 @@ from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN class _UsageRecorder(CustomLogger): - def __init__(self) -> None: + def __init__(self, model_key: str = "typesafe/jev-accounting") -> None: super().__init__() + self.model_key = model_key self.calls: tuple[Mapping[str, object], ...] = () async def async_log_success_event( self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime ) -> None: - if str(kwargs.get("model", "")).removeprefix("typesafe/") != "jev-accounting": + if str(kwargs.get("model", "")) != self.model_key: return self.calls = (*self.calls, kwargs) @@ -167,8 +169,9 @@ async def test_jev_invalid_usage_never_reaches_spend_callbacks( @pytest.mark.asyncio @pytest.mark.parametrize("answer", ["SIMPLE", "UNAVAILABLE", "malformed"]) @pytest.mark.parametrize("private", [False, True]) +@pytest.mark.parametrize("legacy", [False, True]) async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fails( - monkeypatch: pytest.MonkeyPatch, answer: str, private: bool + monkeypatch: pytest.MonkeyPatch, answer: str, private: bool, legacy: bool ) -> None: recorder: Final = _UsageRecorder() monkeypatch.setattr(litellm, "_async_success_callback", [recorder]) @@ -196,7 +199,15 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail router: Final = ComplexityRouter( "jev-router", litellm.Router(model_list=[]), - {"classifier_type": "jev", "jev_classifier_config": {}, "tiers": {"SIMPLE": "cheap"}}, + { + "classifier_type": "jev" if legacy else "oss_classifier", + "jev_classifier_config" if legacy else "opensource_classifier_config": { + "provider": "typesafe" if legacy else "jev", + }, + "tiers": {"SIMPLE": "cheap"}, + "session_affinity": False, + "deployment_affinity": False, + }, jev_client=provider, derive_savings_baseline=False, ) @@ -209,8 +220,9 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail "user_api_key_budget_reservation": {"reservation_id": "parent-reservation"}, "user_api_key_auth": {"budget_reservation": {"reservation_id": "parent-reservation"}}, } - outcome: Final = await router.aclassify( - "private current ask", + result: Final = await router.async_pre_routing_hook( + model="jev-router", + messages=[{"role": "user", "content": "private current ask"}], request_kwargs={ "metadata": metadata, "litellm_session_id": "session-a", @@ -221,7 +233,15 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail await GLOBAL_LOGGING_WORKER.flush() await handler.client.aclose() - assert (outcome.cause == "jev_classifier") is (answer == "SIMPLE") + assert result is not None and result.model == "cheap" + assert result.routing_decision is not None + decision: Final = result.routing_decision + assert (decision["cause"] == "jev_classifier") is (answer == "SIMPLE") + if answer == "SIMPLE": + assert decision["classifier_model"] == "typesafe/jev-accounting" + assert decision["classifier_cost"] == pytest.approx(0.007) + assert "jev-classifier:SIMPLE" in decision["signals"] + assert "jev-confidence=1.000000" in decision["signals"] assert len(recorder.calls) == 1 event: Final = recorder.calls[0] assert event["response_cost"] == pytest.approx(0.007) @@ -416,10 +436,101 @@ def _answer(choice: str = "SIMPLE") -> JevChoiceAnswer: def test_jev_config_requires_classifier_config() -> None: - with pytest.raises(ValueError, match="jev_classifier_config is required"): + with pytest.raises(ValueError, match="opensource_classifier_config is required"): ComplexityRouterConfig.model_validate({"classifier_type": "jev"}) +@pytest.mark.parametrize( + ("classifier_type", "config_key"), + [ + ("oss_classifier", "opensource_classifier_config"), + ("jev", "jev_classifier_config"), + ("oss_classifier", "jev_classifier_config"), + ("jev", "opensource_classifier_config"), + ], +) +@pytest.mark.parametrize( + ("provider", "model", "canonical_provider"), + [(None, "jev-latest", "jev"), ("typesafe", "jev-latest", "jev"), ("jev", "jev-latest", "jev"), ("laya", "english", "laya")], +) +def test_classifier_aliases_load_and_serialize_one_canonical_config( + classifier_type: str, config_key: str, provider: str | None, model: str, canonical_provider: str +) -> None: + incoming: Final = { + "classifier_type": classifier_type, + config_key: {"model": model, "api_key": None, **({"provider": provider} if provider is not None else {})}, + } + original: Final = deepcopy(incoming) + config: Final = ComplexityRouterConfig.model_validate(incoming) + assert config.classifier_type == "oss_classifier" + assert config.opensource_classifier_config is not None + assert config.opensource_classifier_config.provider == canonical_provider + assert config.opensource_classifier_config.model == model + assert config.opensource_classifier_config.api_key is None + assert "api_key" in config.opensource_classifier_config.model_fields_set + assert "api_base" not in config.opensource_classifier_config.model_fields_set + assert "jev_classifier_config" not in config.model_dump() + assert config.jev_classifier_config is config.opensource_classifier_config + assert incoming == original + + +@pytest.mark.parametrize("config", [{"provider": "laya"}, {"provider": "laya", "model": " "}]) +def test_laya_requires_its_own_checkpoint(config: Mapping[str, object]) -> None: + with pytest.raises(ValueError, match="Laya model must be"): + JevClassifierConfig.model_validate(config) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("custom_base", [False, True]) +@pytest.mark.parametrize("legacy", [False, True]) +async def test_laya_routes_with_its_own_credentials_and_accounts_the_checkpoint( + monkeypatch: pytest.MonkeyPatch, custom_base: bool, legacy: bool +) -> None: + monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key") + monkeypatch.setenv("LAYA_API_BASE", "https://laya.test") + monkeypatch.setenv("LAYA_API_KEY", "laya-env-key") + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setitem(litellm.model_cost, "laya/english", {"input_cost_per_token": 0.01}) + recorder: Final = _UsageRecorder("laya/english") + monkeypatch.setattr(litellm, "_async_success_callback", [recorder]) + router: Final = ComplexityRouter( + "laya-route", + litellm.Router(model_list=[]), + { + "classifier_type": "jev" if legacy else "oss_classifier", + "jev_classifier_config" if legacy else "opensource_classifier_config": { + "provider": "laya", + "model": "english", + **({"api_base": "https://laya.test"} if custom_base else {}), + }, + "tiers": {"SIMPLE": "cheap"}, + }, + derive_savings_baseline=False, + ) + with respx.mock(assert_all_called=True) as upstream: + route: Final = upstream.post("https://laya.test/v1/systemone").respond( + 200, + json={ + "model": "laya-rl-agent", + "routing": {"model": "english"}, + "answers": {"tier": _answer().model_dump()}, + "usage": {"input_tokens": 31, "output_tokens": 0}, + }, + ) + outcome: Final = await router.aclassify("choose a tier") + await GLOBAL_LOGGING_WORKER.flush() + + assert outcome.cause == "jev_classifier" + assert outcome.jev_verdict is not None + assert (outcome.jev_verdict.provider, outcome.jev_verdict.model) == ("laya", "english") + assert outcome.classifier_cost == pytest.approx(0.31) + sent: Final = route.calls.last.request + assert sent.headers.get("authorization") == (None if custom_base else "Bearer laya-env-key") + assert json.loads(sent.content)["model"] == "english" + assert len(recorder.calls) == 1 + assert recorder.calls[0]["response_cost"] == pytest.approx(0.31) + + def test_jev_config_is_rejected_for_other_classifier_types() -> None: with pytest.raises(ValueError, match="has no effect"): ComplexityRouterConfig.model_validate( @@ -437,7 +548,7 @@ def test_jev_instructions_reject_blank_values() -> None: @pytest.mark.parametrize( ("missing_key", "rejection"), [ - ({}, r"api_base requires jev_classifier_config\.api_key"), + ({}, r"api_base requires opensource_classifier_config\.api_key"), ({"api_key": ""}, r"api_key must be non-empty"), ({"api_key": " "}, r"api_key must be non-empty"), ], diff --git a/tests/unit/router_utils/test_auto_router_model_naming.py b/tests/unit/router_utils/test_auto_router_model_naming.py index 645f9e5e62a..4881b850f2a 100644 --- a/tests/unit/router_utils/test_auto_router_model_naming.py +++ b/tests/unit/router_utils/test_auto_router_model_naming.py @@ -6,11 +6,11 @@ import pytest from litellm.router_strategy.complexity_router.fuse_presets import get_fuse_presets from litellm.router_strategy.complexity_router.jev_classifier import DEFAULT_JEV_INSTRUCTIONS from litellm.router_utils.auto_router_model_naming import ( - carries_complexity_router_settings, - classify_strategy_router_model, GATED_AUTO_ROUTER_CAPABILITIES, capability_limit_violation, + carries_complexity_router_settings, claimed_capability, + classify_strategy_router_model, count_capability_routers, gated_capability_of, strategy_router_dependencies, @@ -23,27 +23,59 @@ COMPLEXITY_FIELDS = frozenset({"complexity_router_config"}) SEMANTIC_FIELDS = frozenset({"auto_router_config", "auto_router_default_model", "auto_router_embedding_model"}) -@pytest.mark.parametrize("model", ["jev-latest", "jev-preview"]) -def test_jev_enumerates_a_paid_evaluation_without_a_completion_classifier(model: str) -> None: - found = strategy_router_dependencies( +@pytest.mark.parametrize( + ("classifier_type", "config_key"), + [ + ("jev", "jev_classifier_config"), + ("oss_classifier", "opensource_classifier_config"), + ("jev", "opensource_classifier_config"), + ("oss_classifier", "jev_classifier_config"), + ], +) +@pytest.mark.parametrize( + ("provider", "model", "accounting_provider"), + [ + (None, "jev-latest", "typesafe"), + ("typesafe", "jev-preview", "typesafe"), + ("jev", "jev-preview", "typesafe"), + ("laya", "english", "laya"), + ], +) +def test_open_source_classifier_enumerates_its_accounting_model( + classifier_type: str, config_key: str, provider: str | None, model: str, accounting_provider: str +) -> None: + found: Final = strategy_router_dependencies( { "model": "auto_router/complexity_router", "complexity_router_config": { - "classifier_type": "jev", - "jev_classifier_config": {"model": model}, + "classifier_type": classifier_type, + config_key: {"model": model, **({"provider": provider} if provider else {})}, "tiers": {"SIMPLE": "cheap"}, }, } ) assert tuple((dep.model_name, dep.role) for dep in found) == ( ("cheap", "tier"), - (f"typesafe/{model}", "evaluation"), + (f"{accounting_provider}/{model}", "evaluation"), ) @pytest.mark.parametrize("instructions", [None, DEFAULT_JEV_INSTRUCTIONS, "Route conservatively"]) -def test_only_non_default_jev_instructions_claim_the_shared_customization_slot(instructions: str | None) -> None: - capability = claimed_capability({"classifier_type": "jev", "jev_classifier_config": {"instructions": instructions}}) +@pytest.mark.parametrize( + ("classifier_type", "config_key"), + [ + ("jev", "jev_classifier_config"), + ("oss_classifier", "opensource_classifier_config"), + ("jev", "opensource_classifier_config"), + ("oss_classifier", "jev_classifier_config"), + ], +) +def test_only_non_default_open_source_instructions_claim_the_shared_customization_slot( + instructions: str | None, classifier_type: str, config_key: str +) -> None: + capability: Final = claimed_capability( + {"classifier_type": classifier_type, config_key: {"instructions": instructions}} + ) assert (capability.key if capability else None) == ( "tier_or_classifier_prompt" if instructions == "Route conservatively" else None ) @@ -123,6 +155,21 @@ VALID_TIERS = { } +@pytest.mark.parametrize("legacy_config", [None, {}, {"provider": "laya", "model": "english"}]) +def test_dual_classifier_blocks_return_a_write_validation_error(legacy_config: Mapping[str, object] | None) -> None: + violation: Final = validate_complexity_router_config_write( + { + "tiers": VALID_TIERS, + "classifier_type": "oss_classifier", + "opensource_classifier_config": {"provider": "laya", "model": "english"}, + "jev_classifier_config": legacy_config, + } + ) + assert violation is not None + assert "opensource_classifier_config" in violation + assert "jev_classifier_config" in violation + + @pytest.mark.parametrize( "keyword_tier_rules,expected_fragment", [ @@ -408,6 +455,8 @@ def test_complexity_embedding_model_is_a_dependency_only_when_semantic_matching_ ("token_thresholds", "dimension_weights"), ("reasoning_override_min_score",), ("tiers",), + ("jev_classifier_config",), + ("opensource_classifier_config",), ], ) def test_placement_rejects_settings_written_beside_the_config(misplaced): @@ -447,7 +496,7 @@ def test_placement_guards_every_setting_the_config_owns(): ComplexityRouterConfig, ) - assert COMPLEXITY_ROUTER_CONFIG_KEYS == frozenset(ComplexityRouterConfig.model_fields) + assert COMPLEXITY_ROUTER_CONFIG_KEYS == frozenset(ComplexityRouterConfig.model_fields) | {"jev_classifier_config"} assert {"tier_boundaries", "token_thresholds", "dimension_weights"} <= COMPLEXITY_ROUTER_CONFIG_KEYS diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 4c36f99d2c8..ab7ecfcda05 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -1020,6 +1020,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "/v1/videos", "/vertex_ai/live", "/v1/listen", + "/v1/systemone", "/v1beta/interactions", ], }, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 43098982a5e..1d50f37e573 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -8830,6 +8830,23 @@ export interface paths { patch: operations["langfuse_proxy_route_langfuse__endpoint__patch"]; trace?: never; }; + "/laya/v1/systemone": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Laya Proxy Route */ + post: operations["laya_proxy_route_laya_v1_systemone_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/lazy/warm/{name}": { parameters: { query?: never; @@ -33152,44 +33169,6 @@ export interface components { /** Updated By */ updated_by?: string | null; }; - /** JevClassifierConfig */ - JevClassifierConfig: { - /** - * Api Base - * @description TypeSafe API base, falling back to TYPESAFE_API_BASE and then https://api.typesafe.ai - */ - api_base?: string | null; - /** - * Api Key - * @description TypeSafe API key, falling back to TYPESAFE_API_KEY - */ - api_key?: string | null; - /** - * Circuit Breaker Cooldown Seconds - * @default 30 - */ - circuit_breaker_cooldown_seconds: number; - /** - * Circuit Breaker Enabled - * @default true - */ - circuit_breaker_enabled: boolean; - /** - * Instructions - * @description Replaces the built-in Jev question instructions - */ - instructions?: string | null; - /** - * Model - * @default jev-latest - */ - model: string; - /** - * Timeout Ms - * @default 3000 - */ - timeout_ms: number; - }; /** Job */ Job: { /** @@ -39134,6 +39113,50 @@ export interface components { */ type: "openIdConnect"; }; + /** OpenSourceClassifierConfig */ + OpenSourceClassifierConfig: { + /** + * Api Base + * @description Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider + */ + api_base?: string | null; + /** + * Api Key + * @description Provider API key; optional for self-hosted Laya + */ + api_key?: string | null; + /** + * Circuit Breaker Cooldown Seconds + * @default 30 + */ + circuit_breaker_cooldown_seconds: number; + /** + * Circuit Breaker Enabled + * @default true + */ + circuit_breaker_enabled: boolean; + /** + * Instructions + * @description Replaces the built-in Jev question instructions + */ + instructions?: string | null; + /** + * Model + * @default jev-latest + */ + model: string; + /** + * Provider + * @default jev + * @enum {string} + */ + provider: "jev" | "laya"; + /** + * Timeout Ms + * @default 3000 + */ + timeout_ms: number; + }; /** * OperationCreateFile * @description Instruction describing how to create a file via the apply_patch tool. @@ -41934,11 +41957,11 @@ export interface components { classifier_plugin_timeout_ms: number; /** * Classifier Type - * @description Classification strategy: local regex/keyword scoring, the bundled trained four-tier heuristic, an LLM tier-selection call, a Switchyard-compatible capability forecast, a joint Fuse V2 forecast, a custom classifier plugin, 'heuristic_first', which scores locally and only pays for the LLM classifier when the local scorer does not confidently land a cheap tier, or 'hybrid', which trusts the local scorer everywhere except when its score lands near a tier boundary, or 'jev', a TypeSafe AI Jev structured choice call + * @description Classification strategy: local regex/keyword scoring, the bundled trained four-tier heuristic, an LLM tier-selection call, a Switchyard-compatible capability forecast, a joint Fuse V2 forecast, a custom classifier plugin, 'heuristic_first', which scores locally and only pays for the LLM classifier when the local scorer does not confidently land a cheap tier, or 'hybrid', which trusts the local scorer everywhere except when its score lands near a tier boundary, or 'oss_classifier', a structured choice call using Jev or Laya * @default heuristic * @enum {string} */ - classifier_type: "heuristic" | "heuristic_v2" | "llm" | "capability" | "llm_v2" | "custom" | "heuristic_first" | "hybrid" | "jev"; + classifier_type: "heuristic" | "heuristic_v2" | "llm" | "capability" | "llm_v2" | "custom" | "heuristic_first" | "hybrid" | "oss_classifier"; /** * Code Keywords * @description Keywords indicating code-related content @@ -42037,7 +42060,6 @@ export interface components { * @description How close to a tier boundary a heuristic score has to land before the LLM classifier breaks the tie; required when classifier_type is 'hybrid' and rejected otherwise. Everything further than this from every active boundary routes on the scorer's own tier with no classifier call, at any tier, which is what separates 'hybrid' from 'heuristic_first' and its cheap-tier ceiling. A prompt where no dimension fired still goes to the classifier, since the scorer has no opinion to be near a boundary with. 0 escalates only scores sitting exactly on a boundary. */ hybrid_boundary_margin?: number | null; - jev_classifier_config?: components["schemas"]["JevClassifierConfig"] | null; /** * Keyword Tier Rules * @description Rules that force a specific tier when their keywords match the prompt @@ -42069,6 +42091,7 @@ export interface components { * @default false */ modality_routing: boolean; + opensource_classifier_config?: components["schemas"]["OpenSourceClassifierConfig"] | null; /** * Plan Mode Min Tier * @description When set, requests carrying a coding-agent plan-mode sentinel (Claude Code plan mode, VS Code Copilot Plan mode, Copilot CLI's exit_plan_mode tool) are routed to at least this tier: the classified tier still wins when it is higher, and the floor also overrides a session-affinity pin to a lower tier for exactly the turns carrying the sentinel, without rewriting the pin -- the first turn after plan mode exits routes as if plan mode had never happened. Names a built-in tier, or with tier_definitions set, one of the defined tier names (list order is ascending severity, same as keyword_tier_rules). Unset disables detection entirely. The sentinels ride in client-injected prompt text, so a caller who pastes one can spend up to this tier's models -- never down, and never outside the configured pools. @@ -42166,7 +42189,7 @@ export interface components { }; /** * Tier Definitions - * @description Operator-defined tier set replacing the built-in SIMPLE/MEDIUM/COMPLEX/REASONING. Each entry's name becomes a value the LLM classifier can return and its description becomes that tier's rubric bullet; entries named after a built-in tier may omit the description and inherit the built-in criteria. List order is ascending severity and decides which tier wins when several keyword_tier_rules match. Requires classifier_type 'llm', 'jev' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, adaptive selection, session affinity, plugins, tier_labels, and the calibration-example rubric presets are unavailable with a custom tier set: the first four are built on the built-in tier ladder, and the last two rename or exemplify tiers the set replaces. + * @description Operator-defined tier set replacing the built-in SIMPLE/MEDIUM/COMPLEX/REASONING. Each entry's name becomes a value the LLM classifier can return and its description becomes that tier's rubric bullet; entries named after a built-in tier may omit the description and inherit the built-in criteria. List order is ascending severity and decides which tier wins when several keyword_tier_rules match. Requires classifier_type 'llm', 'oss_classifier' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, adaptive selection, session affinity, plugins, tier_labels, and the calibration-example rubric presets are unavailable with a custom tier set: the first four are built on the built-in tier ladder, and the last two rename or exemplify tiers the set replaces. */ tier_definitions?: components["schemas"]["TierDefinition"][] | null; /** @@ -61639,6 +61662,26 @@ export interface operations { }; }; }; + laya_proxy_route_laya_v1_systemone_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; warm_lazy_warm__name__post: { parameters: { query?: never; From 8aaa7667343c0f48ce7ecded3f01106e82f2774a Mon Sep 17 00:00:00 2001 From: tin-berri Date: Fri, 2 Oct 2026 11:57:42 -0700 Subject: [PATCH 017/139] fix(auto-router): count usage savings by selected UTC request day (#44115) The auto-router usage view summed the lifetime savings of every session overlapping the date range, so it disagreed with the Overall savings view, which sums daily rollups by request day. Record auto-routed money per UTC request day and router in one new table, written in the same statement as the session rollup and corrected in the same transaction as late baseline estimates. The all-router headline reads the same daily rows and filters as Overall; savings no router day row accounts for are reported as unattributed and void the baseline comparison. Session shape and caching stay whole-session and are labelled so; the savings-per-session tile is removed. Co-authored-by: Claude Opus 5.5 --- .../migration.sql | 17 ++ .../litellm_proxy_extras/schema.prisma | 21 ++ litellm/proxy/db/autorouter_session_rollup.py | 144 ++++++++---- litellm/proxy/db/baseline_accounting.py | 25 +++ .../db_transaction_queue/spend_log_cleanup.py | 22 +- .../auto_router_endpoints.py | 154 +++++++++---- litellm/proxy/schema.prisma | 21 ++ .../auto_router_endpoints.py | 29 ++- schema.prisma | 21 ++ .../spend/test_autorouter_session_rollup.py | 211 ++++++++++-------- .../spend/test_baseline_accounting.py | 7 + .../test_auto_router_endpoints.py | 132 ++++++++--- tests/unit/proxy/test_spend_log_cleanup.py | 17 +- .../AutoRouterBenchmarksTab.test.tsx | 47 ++-- .../_components/AutoRouterBenchmarksTab.tsx | 44 ++-- .../_components/TierTurnsChart.test.tsx | 6 +- .../_components/autoRouterBenchmarks.test.ts | 1 - .../user_info_view.integration.test.tsx | 1 - ...KeyAutoRouterUsageTab.integration.test.tsx | 1 - ui/litellm-dashboard/src/lib/http/schema.d.ts | 102 ++++++--- 20 files changed, 739 insertions(+), 284 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20261001200000_add_autorouter_daily_spend/migration.sql diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261001200000_add_autorouter_daily_spend/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261001200000_add_autorouter_daily_spend/migration.sql new file mode 100644 index 00000000000..ce166b4df45 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261001200000_add_autorouter_daily_spend/migration.sql @@ -0,0 +1,17 @@ +CREATE TABLE IF NOT EXISTS "LiteLLM_AutoRouterDailySpend" ( + "date" TEXT NOT NULL, + "api_key" TEXT NOT NULL, + "user_id" TEXT NOT NULL, + "router_name" TEXT NOT NULL, + "router_type" TEXT NOT NULL, + "turns" INTEGER NOT NULL DEFAULT 0, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0, + "saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0, + "savings_estimated_turns" INTEGER NOT NULL DEFAULT 0, + "savings_estimated_actual_spend" DOUBLE PRECISION NOT NULL DEFAULT 0, + "savings_estimated_saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0, + "classifier_cost" DOUBLE PRECISION NOT NULL DEFAULT 0, + "classifier_cost_recorded_turns" INTEGER NOT NULL DEFAULT 0, + + CONSTRAINT "LiteLLM_AutoRouterDailySpend_pkey" PRIMARY KEY ("date", "api_key", "user_id", "router_name", "router_type") +); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 6f285e9dc39..aba89526cf6 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1744,6 +1744,27 @@ model LiteLLM_AutoRouterUserSession { @@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn") } +// Auto-routed requests per UTC request day and router: the selected-day money behind the +// auto-router usage view. Written in the same statement as the session rollup, so a day row +// and its session row never disagree; corrected in the same transaction as late baselines. +model LiteLLM_AutoRouterDailySpend { + date String + api_key String + user_id String + router_name String + router_type String + turns Int @default(0) + spend Float @default(0) + saved_spend Float @default(0) + savings_estimated_turns Int @default(0) + savings_estimated_actual_spend Float @default(0) + savings_estimated_saved_spend Float @default(0) + classifier_cost Float @default(0) + classifier_cost_recorded_turns Int @default(0) + + @@id([date, api_key, user_id, router_name, router_type]) +} + // Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in // either direction. forward duplicates the requests the keys did not route through the // router through it, answering whether they should adopt it; reverse duplicates the diff --git a/litellm/proxy/db/autorouter_session_rollup.py b/litellm/proxy/db/autorouter_session_rollup.py index b762a40f344..cdc949a6dca 100644 --- a/litellm/proxy/db/autorouter_session_rollup.py +++ b/litellm/proxy/db/autorouter_session_rollup.py @@ -4,11 +4,12 @@ Per-session auto-router benchmarks rollup. At request time the spend writer builds one AutoRouterTurnTransaction per successful auto-routed request (a request whose metadata carries a routing_decision) and queues it on the prisma client. The spend-log flush job drains the queue into -key and user session rollups with one atomic statement per turn: each upsert classifies +key and user session rollups, plus the per-day router rollup, with one atomic statement +per turn: each upsert classifies the turn (same model, first visit, return to a model the session already used, out of order) against the row's own columns, so nothing is read before the write and concurrent -pods compose. The benchmarks endpoint aggregates these rows and can recover matching historical -costs from retained spend logs when estimate coverage predates these columns. +pods compose. The benchmarks endpoint reads session shape from the session rows and money from the +day rows, so spend and savings count only requests on the selected UTC days. """ from __future__ import annotations @@ -71,45 +72,82 @@ tier_maps AS ( GROUP BY router_name, router_type, kv.key ) per_tier GROUP BY router_name, router_type +), +sessions AS ( + SELECT + router_name, + router_type, + COUNT(*)::int AS sessions, + SUM(turns)::int AS session_turns, + SUM(unordered_turns)::int AS unordered_turns, + SUM(covered_turns)::int AS covered_turns, + SUM(cache_hits)::int AS cache_hits, + SUM(same_model_turns)::int AS same_model_turns, + SUM(same_model_hits)::int AS same_model_hits, + SUM(first_visit_turns)::int AS first_visit_turns, + SUM(first_visit_hits)::int AS first_visit_hits, + SUM(return_turns)::int AS return_turns, + SUM(return_hits)::int AS return_hits, + SUM(return_expired_misses)::int AS return_expired_misses, + SUM(return_within_ttl_misses)::int AS return_within_ttl_misses, + SUM(ttl_5m_turns)::int AS ttl_5m_turns, + SUM(ttl_1h_turns)::int AS ttl_1h_turns, + SUM(total_tokens)::bigint AS total_tokens, + SUM(EXTRACT(EPOCH FROM (last_turn_at - first_turn_at)))::float8 AS session_seconds + FROM windowed + GROUP BY router_name, router_type +), +days AS ( + SELECT + router_name, + router_type, + SUM(turns)::int AS turns, + SUM(spend)::float8 AS spend, + SUM(saved_spend)::float8 AS saved_spend, + SUM(savings_estimated_turns)::int AS savings_estimated_turns, + SUM(savings_estimated_actual_spend)::float8 AS savings_estimated_actual_spend, + SUM(savings_estimated_saved_spend)::float8 AS savings_estimated_saved_spend, + SUM(classifier_cost)::float8 AS classifier_cost, + SUM(classifier_cost_recorded_turns)::int AS classifier_cost_recorded_turns + FROM "LiteLLM_AutoRouterDailySpend" + WHERE date >= $5 AND date <= $6 + AND ($3::text IS NULL OR api_key = $3::text) + AND ($4::text IS NULL OR user_id = $4::text) + GROUP BY router_name, router_type ) -SELECT - agg.*, - COALESCE(tier_maps.tier_turns, '{{}}'::jsonb) AS tier_turns -FROM ( SELECT router_name, router_type, - COUNT(*)::int AS sessions, - COALESCE(SUM(turns), 0)::int AS turns, - COALESCE(SUM(unordered_turns), 0)::int AS unordered_turns, - COALESCE(SUM(covered_turns), 0)::int AS covered_turns, - COALESCE(SUM(cache_hits), 0)::int AS cache_hits, - COALESCE(SUM(same_model_turns), 0)::int AS same_model_turns, - COALESCE(SUM(same_model_hits), 0)::int AS same_model_hits, - COALESCE(SUM(first_visit_turns), 0)::int AS first_visit_turns, - COALESCE(SUM(first_visit_hits), 0)::int AS first_visit_hits, - COALESCE(SUM(return_turns), 0)::int AS return_turns, - COALESCE(SUM(return_hits), 0)::int AS return_hits, - COALESCE(SUM(return_expired_misses), 0)::int AS return_expired_misses, - COALESCE(SUM(return_within_ttl_misses), 0)::int AS return_within_ttl_misses, - COALESCE(SUM(ttl_5m_turns), 0)::int AS ttl_5m_turns, - COALESCE(SUM(ttl_1h_turns), 0)::int AS ttl_1h_turns, - COALESCE(SUM(total_tokens), 0)::bigint AS total_tokens, - COALESCE(SUM(spend), 0)::float8 AS spend, - COALESCE(SUM(saved_spend), 0)::float8 AS saved_spend, - COALESCE(SUM(savings_estimated_turns), 0)::int AS savings_estimated_turns, - COALESCE(SUM(savings_estimated_actual_spend), 0)::float8 AS savings_estimated_actual_spend, - CASE WHEN BOOL_AND(savings_estimated_turns = turns AND classifier_cost_recorded_turns = turns) - THEN SUM(classifier_cost)::float8 END AS savings_estimated_classifier_cost, - COALESCE(SUM(savings_estimated_saved_spend), 0)::float8 AS savings_estimated_saved_spend, - COALESCE(SUM(classifier_cost), 0)::float8 AS classifier_cost, - COALESCE(SUM(classifier_cost_recorded_turns), 0)::int AS classifier_cost_recorded_turns, - COALESCE(SUM(EXTRACT(EPOCH FROM (last_turn_at - first_turn_at))), 0)::float8 AS session_seconds -FROM windowed -GROUP BY router_name, router_type -) agg + COALESCE(tier_maps.tier_turns, '{{}}'::jsonb) AS tier_turns, + COALESCE(sessions.sessions, 0) AS sessions, + COALESCE(sessions.session_turns, 0) AS session_turns, + COALESCE(sessions.unordered_turns, 0) AS unordered_turns, + COALESCE(sessions.covered_turns, 0) AS covered_turns, + COALESCE(sessions.cache_hits, 0) AS cache_hits, + COALESCE(sessions.same_model_turns, 0) AS same_model_turns, + COALESCE(sessions.same_model_hits, 0) AS same_model_hits, + COALESCE(sessions.first_visit_turns, 0) AS first_visit_turns, + COALESCE(sessions.first_visit_hits, 0) AS first_visit_hits, + COALESCE(sessions.return_turns, 0) AS return_turns, + COALESCE(sessions.return_hits, 0) AS return_hits, + COALESCE(sessions.return_expired_misses, 0) AS return_expired_misses, + COALESCE(sessions.return_within_ttl_misses, 0) AS return_within_ttl_misses, + COALESCE(sessions.ttl_5m_turns, 0) AS ttl_5m_turns, + COALESCE(sessions.ttl_1h_turns, 0) AS ttl_1h_turns, + COALESCE(sessions.total_tokens, 0) AS total_tokens, + COALESCE(sessions.session_seconds, 0) AS session_seconds, + COALESCE(days.turns, 0) AS turns, + COALESCE(days.spend, 0) AS spend, + COALESCE(days.saved_spend, 0) AS saved_spend, + COALESCE(days.savings_estimated_turns, 0) AS savings_estimated_turns, + COALESCE(days.savings_estimated_actual_spend, 0) AS savings_estimated_actual_spend, + COALESCE(days.savings_estimated_saved_spend, 0) AS savings_estimated_saved_spend, + COALESCE(days.classifier_cost, 0) AS classifier_cost, + COALESCE(days.classifier_cost_recorded_turns, 0) AS classifier_cost_recorded_turns +FROM sessions +FULL OUTER JOIN days USING (router_name, router_type) LEFT JOIN tier_maps USING (router_name, router_type) -ORDER BY agg.spend DESC +ORDER BY spend DESC, router_name, router_type """ @@ -391,15 +429,43 @@ ON CONFLICT ({user_column}api_key, session_id, router_name) DO UPDATE SET """ +_DAY_UPSERT_SQL: Final = f""" +day_rollup AS ( + INSERT INTO "LiteLLM_AutoRouterDailySpend" AS d ( + date, api_key, user_id, router_name, router_type, turns, spend, saved_spend, savings_estimated_turns, + savings_estimated_actual_spend, savings_estimated_saved_spend, classifier_cost, classifier_cost_recorded_turns + ) + VALUES ( + ({_TURN_AT}::timestamp)::date::text, {_p("api_key")}::text, {_p("user_id")}::text, {_p("router_name")}, + {_p("router_type")}, 1, {_p("spend")}::float8, {_p("saved_spend")}::float8, {_p("savings_estimated_turns")}::int, + {_p("savings_estimated_actual_spend")}::float8, {_p("savings_estimated_saved_spend")}::float8, + {_p("classifier_cost")}::float8, 1 + ) + ON CONFLICT (date, api_key, user_id, router_name, router_type) DO UPDATE SET + turns = d.turns + 1, + spend = d.spend + EXCLUDED.spend, + saved_spend = d.saved_spend + EXCLUDED.saved_spend, + savings_estimated_turns = d.savings_estimated_turns + EXCLUDED.savings_estimated_turns, + savings_estimated_actual_spend = d.savings_estimated_actual_spend + EXCLUDED.savings_estimated_actual_spend, + savings_estimated_saved_spend = d.savings_estimated_saved_spend + EXCLUDED.savings_estimated_saved_spend, + classifier_cost = d.classifier_cost + EXCLUDED.classifier_cost, + classifier_cost_recorded_turns = d.classifier_cost_recorded_turns + 1 + RETURNING 1 +) +""" + UPSERT_AUTOROUTER_SESSION_SQL: Final = f""" WITH key_rollup AS ( {_session_upsert_sql(user_scoped=False)} RETURNING 1 -) +), {_DAY_UPSERT_SQL} {_session_upsert_sql(user_scoped=True)} """ -UPSERT_AUTOROUTER_USER_SESSION_SQL: Final = _session_upsert_sql(user_scoped=True) +UPSERT_AUTOROUTER_USER_SESSION_SQL: Final = f""" +WITH {_DAY_UPSERT_SQL} +{_session_upsert_sql(user_scoped=True)} +""" def _as_sql_param(value: str | float | bool | datetime | None) -> str | float | None: diff --git a/litellm/proxy/db/baseline_accounting.py b/litellm/proxy/db/baseline_accounting.py index 05a4a989152..9536f8d740a 100644 --- a/litellm/proxy/db/baseline_accounting.py +++ b/litellm/proxy/db/baseline_accounting.py @@ -179,6 +179,8 @@ class _Change(BaseModel): actual_delta: float savings_delta: float daily: DailyBaselineAttribution | None + date: str | None = None + router_type: str | None = None class _TransactionManager(Protocol): @@ -303,6 +305,26 @@ WHERE {user_match}session.api_key = totals.api_key AND session.session_id = tota _UPDATE_SESSIONS: Final = _session_correction_sql(user_scoped=False) _UPDATE_USER_SESSIONS: Final = _session_correction_sql(user_scoped=True) +_UPDATE_DAYS: Final = """ +WITH totals AS ( + SELECT date, api_key, user_id, router_name, router_type, SUM(covered_delta)::int AS covered_delta, + SUM(actual_delta) AS actual_delta, SUM(savings_delta) AS savings_delta + FROM jsonb_to_recordset($1::jsonb) AS x( + date text, api_key text, user_id text, router_name text, router_type text, + covered_delta int, actual_delta float8, savings_delta float8 + ) + WHERE date IS NOT NULL + GROUP BY date, api_key, user_id, router_name, router_type +) +UPDATE "LiteLLM_AutoRouterDailySpend" AS day +SET saved_spend = day.saved_spend + totals.savings_delta, + savings_estimated_turns = day.savings_estimated_turns + totals.covered_delta, + savings_estimated_actual_spend = day.savings_estimated_actual_spend + totals.actual_delta, + savings_estimated_saved_spend = day.savings_estimated_saved_spend + totals.savings_delta +FROM totals +WHERE day.date = totals.date AND day.api_key = totals.api_key AND day.user_id = totals.user_id + AND day.router_name = totals.router_name AND day.router_type = totals.router_type +""" def _primary_transaction(client: PrismaClient) -> _TransactionManager: @@ -331,6 +353,8 @@ def _change(record: BaselineAccountingRecord, old: BaselinePublication | None, n savings_delta=(current.savings if current is not None else 0.0) - (previous.savings if previous is not None else 0.0), daily=record.daily, + date=record.turn.turn_at.date().isoformat() if record.turn is not None else None, + router_type=record.turn.router_type if record.turn is not None else None, ) @@ -373,6 +397,7 @@ async def _publish(db: SupportsRawQueries, changes: Sequence[_Change]) -> None: await db.execute_raw(_UPDATE_SESSIONS, serialized) if any(change.user_id for change in changes): await db.execute_raw(_UPDATE_USER_SESSIONS, serialized) + await db.execute_raw(_UPDATE_DAYS, serialized) for entity, table in DAILY_SPEND_TABLES.items(): if adjustments := tuple( change.daily.adjustment(target, change.savings_delta, change.request_id) diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py index 06e4d06fca4..679e286cedd 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py +++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py @@ -540,6 +540,18 @@ class SpendLogCleanup: deadline=deadline, ) + async def _delete_old_autorouter_daily_rows( + self, prisma_client: PrismaClient, cutoff_day: str, deadline: float + ) -> TableCleanupResult: + return await self._delete_old_rows_batched( + prisma_client, + cutoff_day, + table_name="LiteLLM_AutoRouterDailySpend", + key_columns=("date", "api_key", "user_id", "router_name", "router_type"), + time_column="date", + deadline=deadline, + ) + async def _delete_old_health_check_rows( self, prisma_client: PrismaClient, cutoff_date: datetime, deadline: float ) -> TableCleanupResult: @@ -623,16 +635,20 @@ class SpendLogCleanup: except Exception: # noqa: BLE001 # retained observations are retried by the next cleanup job verbose_proxy_logger.warning("Auto-router baseline retention remains pending") sessions_result: Final = await self._delete_old_autorouter_session_rows( - prisma_client, session_cutoff, self._group_deadline(deadline, 2) + prisma_client, session_cutoff, self._group_deadline(deadline, 3) ) verbose_proxy_logger.info("Deleted %s expired auto-router session rollup rows", sessions_result.rows_deleted) user_sessions_result: Final = await self._delete_old_autorouter_user_session_rows( - prisma_client, session_cutoff, deadline + prisma_client, session_cutoff, self._group_deadline(deadline, 2) ) verbose_proxy_logger.info( "Deleted %s expired auto-router user session rollup rows", user_sessions_result.rows_deleted ) - return (sessions_result, user_sessions_result) + days_result: Final = await self._delete_old_autorouter_daily_rows( + prisma_client, session_cutoff.date().isoformat(), deadline + ) + verbose_proxy_logger.info("Deleted %s expired auto-router daily rollup rows", days_result.rows_deleted) + return (sessions_result, user_sessions_result, days_result) async def _clean_health_checks( self, prisma_client: PrismaClient, retention_seconds: int, deadline: float diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 50145ed3923..35ff9186f5e 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -5,6 +5,8 @@ POST /auto_router/test_routing - Route one request through an unsaved complexity POST /auto_router/validate_complexity_router_config - Dry-run the complexity-router write gate without saving """ +import asyncio +import math from collections.abc import Mapping, Sequence from datetime import datetime, timedelta, timezone from itertools import chain, groupby @@ -40,6 +42,7 @@ from litellm.proxy.litellm_pre_call_utils import ( refresh_proxy_server_request_body_snapshot, ) from litellm.proxy.management.teams.access import is_team_admin +from litellm.proxy.management_endpoints.common_daily_activity import daily_activity_scope from litellm.proxy.management_helpers.auto_router_permissions import ( authorize_member_auto_router_dependencies, authorize_member_auto_router_team, @@ -47,6 +50,7 @@ from litellm.proxy.management_helpers.auto_router_permissions import ( ) from litellm.repositories.autorouter_session_repository import AutoRouterSessionRepository from litellm.repositories.base_repository import SupportsModelDump +from litellm.repositories.daily_activity_sql import build_where_clause from litellm.repositories.team_repository import TeamRepository from litellm.router_strategy.complexity_router import ComplexityRouter from litellm.router_utils.auto_router_model_naming import ( @@ -616,34 +620,37 @@ async def preview_auto_router_routing( class _SessionAggRow(BaseModel): + """One router's window: session shape from overlapping sessions, money from the selected days.""" + router_name: str router_type: str - tier_turns: Mapping[str, int] - sessions: int - turns: int - unordered_turns: int - covered_turns: int - cache_hits: int - same_model_turns: int - same_model_hits: int - first_visit_turns: int - first_visit_hits: int - return_turns: int - return_hits: int - return_expired_misses: int - return_within_ttl_misses: int - ttl_5m_turns: int - ttl_1h_turns: int - total_tokens: int - spend: float - saved_spend: float + tier_turns: Mapping[str, int] = MappingProxyType({}) + sessions: int = 0 + session_turns: int = 0 + unordered_turns: int = 0 + covered_turns: int = 0 + cache_hits: int = 0 + same_model_turns: int = 0 + same_model_hits: int = 0 + first_visit_turns: int = 0 + first_visit_hits: int = 0 + return_turns: int = 0 + return_hits: int = 0 + return_expired_misses: int = 0 + return_within_ttl_misses: int = 0 + ttl_5m_turns: int = 0 + ttl_1h_turns: int = 0 + total_tokens: int = 0 + session_seconds: float = 0.0 + turns: int = 0 + spend: float = 0.0 + saved_spend: float = 0.0 savings_estimated_turns: int = 0 savings_estimated_actual_spend: float = 0.0 savings_estimated_classifier_cost: float | None = None savings_estimated_saved_spend: float = 0.0 - classifier_cost: float - classifier_cost_recorded_turns: int - session_seconds: float + classifier_cost: float = 0.0 + classifier_cost_recorded_turns: int = 0 _SESSION_AGG_ROWS: Final = TypeAdapter(list[_SessionAggRow]) @@ -692,6 +699,13 @@ def _compared_row(row: _SessionAggRow) -> _SessionAggRow: ) +def _per_session(row: _SessionAggRow, total: float) -> float | None: + """Unknown, not zero, when routed requests have no session rows of their own to average over.""" + if row.sessions: + return total / row.sessions + return None if row.turns else 0.0 + + def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals: return_misses: Final = row.return_turns - row.return_hits saved_spend, baseline_spend = _savings_cohort( @@ -701,9 +715,9 @@ def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals: return AutoRouterBenchmarkTotals( sessions=sessions, turns=row.turns, - avg_turns_per_session=row.turns / sessions if sessions else 0.0, - avg_session_seconds=row.session_seconds / sessions if sessions else 0.0, - avg_tokens_per_session=row.total_tokens / sessions if sessions else 0.0, + avg_turns_per_session=_per_session(row, row.session_turns), + avg_session_seconds=_per_session(row, row.session_seconds), + avg_tokens_per_session=_per_session(row, row.total_tokens), spend=row.spend, savings_estimated_turns=row.savings_estimated_turns, savings_estimated_actual_spend=row.savings_estimated_actual_spend, @@ -712,9 +726,8 @@ def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals: classifier_cost=row.classifier_cost if row.classifier_cost_recorded_turns == row.turns else None, baseline_spend=baseline_spend, saved_pct=_pct(saved_spend, baseline_spend) if saved_spend is not None and baseline_spend is not None else None, - saved_per_session=(saved_spend / sessions if sessions else 0.0) if saved_spend is not None else None, cache=AutoRouterCacheStats( - coverage_pct=_pct(row.covered_turns, row.turns), + coverage_pct=_pct(row.covered_turns, row.session_turns), hit_rate_pct=_pct(row.cache_hits, row.covered_turns), same_model=_cache_bucket(row.same_model_turns, row.same_model_hits), first_visit=_cache_bucket(row.first_visit_turns, row.first_visit_hits), @@ -748,7 +761,6 @@ def _benchmark_group(row: _SessionAggRow) -> AutoRouterBenchmarkGroup: classifier_cost=totals.classifier_cost, baseline_spend=totals.baseline_spend, saved_pct=totals.saved_pct, - saved_per_session=totals.saved_per_session, cache=totals.cache, ) @@ -759,6 +771,7 @@ def _summed_agg_row(rows: Sequence[_SessionAggRow]) -> _SessionAggRow: router_type="", tier_turns=MappingProxyType({}), sessions=sum(row.sessions for row in rows), + session_turns=sum(row.session_turns for row in rows), turns=sum(row.turns for row in rows), unordered_turns=sum(row.unordered_turns for row in rows), covered_turns=sum(row.covered_turns for row in rows), @@ -790,6 +803,49 @@ def _summed_agg_row(rows: Sequence[_SessionAggRow]) -> _SessionAggRow: ) +async def _recorded_autorouter_savings( + prisma_client: "PrismaClient", start_day: str, end_day: str, api_key: str | None, user_id: str | None +) -> float: + """The selected days' auto-router savings exactly as the Overall view sums them: same table, same filters.""" + where, params = build_where_clause( + daily_activity_scope( + table="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id=user_id, + exclude_entity_ids=None, + api_key=api_key, + start_date=start_day, + end_date=end_day, + model=None, + timezone_offset_minutes=None, + ) + ) + rows: Final = await _query_raw( + prisma_client, + f'SELECT COALESCE(SUM(autorouter_savings_spend), 0)::float8 AS saved FROM "LiteLLM_DailyUserSpend" WHERE {where}', + *params, + ) + return float(rows[0]["saved"]) if rows else 0.0 + + +def _with_recorded_savings( + totals: AutoRouterBenchmarkTotals, rows: Sequence[_SessionAggRow], recorded: float +) -> AutoRouterBenchmarkTotals: + """The headline is the recorded total. Savings outside the compared routers void the cost comparison, + and the part no router's day rows account for is reported as unattributed.""" + if math.isclose(recorded, totals.saved_spend or 0.0, abs_tol=1e-9): + return totals + unattributed: Final = recorded - sum(row.saved_spend for row in rows) + return totals.model_copy( + update={ + "saved_spend": recorded, + "unattributed_saved_spend": None if math.isclose(unattributed, 0.0, abs_tol=1e-9) else unattributed, + "baseline_spend": None, + "saved_pct": None, + } + ) + + def _strategy_router_key(deployment: object) -> tuple[str, str] | None: """``(model_name, kind)`` for a deployment whose routing the session rollup records. @@ -849,7 +905,7 @@ async def get_auto_router_benchmarks( str | None, Query(description="YYYY-MM-DD UTC, inclusive (defaults to 30 days before end_date)") ] = None, end_date: Annotated[str | None, Query(description="YYYY-MM-DD UTC, inclusive (defaults to today)")] = None, - api_key: Annotated[str | None, Query(description="Filter to one virtual key token hash")] = None, + api_key: Annotated[str | None, Query(min_length=1, description="Filter to one virtual key token hash")] = None, user_id: Annotated[ str | None, Query(min_length=1, description="Filter to one canonical internal user recorded on each turn") ] = None, @@ -860,9 +916,10 @@ async def get_auto_router_benchmarks( 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 + internal user when written; older key-only history remains outside user views. Money counts + only requests on the selected UTC days, and the all-router savings headline is the same daily + total the Overall view reads. Session shape and caching cover every session that overlaps the + window, whole. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is over that bucket's turns. The rollup supplies the measures, never the list. Which routers appear comes from the @@ -885,24 +942,35 @@ async def get_auto_router_benchmarks( if end_day < start_day: raise HTTPException(status_code=400, detail="end_date must not be earlier than start_date") - raw_rows: Final = await _query_raw( - prisma_client, - AUTOROUTER_BENCHMARKS_SQL, - start_day.isoformat(), - (end_day + timedelta(days=1)).isoformat(), - api_key, - user_id, + first_day: Final = start_day.strftime("%Y-%m-%d") + last_day: Final = end_day.strftime("%Y-%m-%d") + raw_rows, recorded = await asyncio.gather( + _query_raw( + prisma_client, + AUTOROUTER_BENCHMARKS_SQL, + start_day.isoformat(), + (end_day + timedelta(days=1)).isoformat(), + api_key, + user_id, + first_day, + last_day, + ), + _recorded_autorouter_savings(prisma_client, first_day, last_day, api_key, user_id), ) rows: Final = tuple(_compared_row(row) for row in _SESSION_AGG_ROWS.validate_python(raw_rows or ())) + totals: Final = _with_recorded_savings(_benchmark_totals(_summed_agg_row(rows)), rows, recorded) + unattributed: Final = MappingProxyType( + {"baseline_spend": None, "saved_pct": None} if totals.unattributed_saved_spend is not None else {} + ) groups: Final = ( - *(_benchmark_group(row) for row in rows), + *(_benchmark_group(row).model_copy(update=unattributed) for row in rows), *_idle_router_groups(llm_router, frozenset((row.router_name, row.router_type) for row in rows)), ) return AutoRouterBenchmarksResponse( - start_date=start_day.strftime("%Y-%m-%d"), - end_date=end_day.strftime("%Y-%m-%d"), + start_date=first_day, + end_date=last_day, routers_in_scope=len(groups), - totals=_benchmark_totals(_summed_agg_row(rows)), + totals=totals, groups=groups, ) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 6f285e9dc39..aba89526cf6 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1744,6 +1744,27 @@ model LiteLLM_AutoRouterUserSession { @@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn") } +// Auto-routed requests per UTC request day and router: the selected-day money behind the +// auto-router usage view. Written in the same statement as the session rollup, so a day row +// and its session row never disagree; corrected in the same transaction as late baselines. +model LiteLLM_AutoRouterDailySpend { + date String + api_key String + user_id String + router_name String + router_type String + turns Int @default(0) + spend Float @default(0) + saved_spend Float @default(0) + savings_estimated_turns Int @default(0) + savings_estimated_actual_spend Float @default(0) + savings_estimated_saved_spend Float @default(0) + classifier_cost Float @default(0) + classifier_cost_recorded_turns Int @default(0) + + @@id([date, api_key, user_id, router_name, router_type]) +} + // Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in // either direction. forward duplicates the requests the keys did not route through the // router through it, answering whether they should adopt it; reverse duplicates the diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index 00083e01f54..ded971f6705 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -205,14 +205,18 @@ class AutoRouterCacheStats(BaseModel): class AutoRouterBenchmarkTotals(BaseModel): - """Session-shape and savings aggregates over auto-routed traffic in the window.""" + """Auto-routed traffic in the window. Turns, spend and savings count requests on the selected UTC days; + the session averages and cache stats describe every session overlapping the window, whole.""" - sessions: int - turns: int - avg_turns_per_session: float - avg_session_seconds: float - avg_tokens_per_session: float - spend: float = Field(description="What the routed traffic actually cost") + sessions: int = Field(description="Sessions overlapping the window, counted whole") + turns: int = Field(description="Auto-routed requests on the selected UTC days") + avg_turns_per_session: float | None = Field( + description="Lifetime turns per overlapping session; null when the window has routed requests but no session " + "rows for this router type, such as an alias whose router type changed mid-session" + ) + avg_session_seconds: float | None = Field(description="Lifetime seconds per overlapping session; null as above") + avg_tokens_per_session: float | None = Field(description="Lifetime tokens per overlapping session; null as above") + spend: float = Field(description="What the selected days' routed traffic actually cost") classifier_cost: float | None = Field( description="Recorded LLM classifier cost already included in spend; null when any session turns predate " "subtotal recording, and zero for an empty window" @@ -229,14 +233,19 @@ class AutoRouterBenchmarkTotals(BaseModel): "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" + description="Recorded savings on the selected UTC days; null when traffic has no recorded savings estimates. " + "On totals this is the same daily figure the Overall savings view reports" + ) + unattributed_saved_spend: float | None = Field( + default=None, + description="Part of saved_spend no router's daily rows account for, such as history recorded before " + "per-router daily tracking; when set, baseline_spend and saved_pct are null", ) 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 @@ -291,7 +300,7 @@ class AutoRouterSessionResponse(BaseModel): class AutoRouterBenchmarksResponse(BaseModel): - """Benchmarks for the auto-router dashboard, aggregated from the per-session rollup.""" + """Benchmarks for the auto-router dashboard, aggregated from the per-session and per-day rollups.""" start_date: str = Field(description="Window start day, YYYY-MM-DD UTC, inclusive") end_date: str = Field(description="Window end day, YYYY-MM-DD UTC, inclusive") diff --git a/schema.prisma b/schema.prisma index 6f285e9dc39..aba89526cf6 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1744,6 +1744,27 @@ model LiteLLM_AutoRouterUserSession { @@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn") } +// Auto-routed requests per UTC request day and router: the selected-day money behind the +// auto-router usage view. Written in the same statement as the session rollup, so a day row +// and its session row never disagree; corrected in the same transaction as late baselines. +model LiteLLM_AutoRouterDailySpend { + date String + api_key String + user_id String + router_name String + router_type String + turns Int @default(0) + spend Float @default(0) + saved_spend Float @default(0) + savings_estimated_turns Int @default(0) + savings_estimated_actual_spend Float @default(0) + savings_estimated_saved_spend Float @default(0) + classifier_cost Float @default(0) + classifier_cost_recorded_turns Int @default(0) + + @@id([date, api_key, user_id, router_name, router_type]) +} + // Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in // either direction. forward duplicates the requests the keys did not route through the // router through it, answering whether they should adopt it; reverse duplicates the diff --git a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py index 2c648f309f6..a5c6f5962a2 100644 --- a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py +++ b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py @@ -80,6 +80,25 @@ async def _turn( ) +async def _benchmark_rows( + db, start: datetime, end: datetime, key: str | None = None, user_id: str | None = None +) -> list[dict]: + return await db.query_raw( + AUTOROUTER_BENCHMARKS_SQL, + start.isoformat(), + end.isoformat(), + key, + user_id, + start.date().isoformat(), + (end - timedelta(days=1)).date().isoformat(), + ) + + +async def _days(db, key: str | None = None, user_id: str | None = None, router: str | None = None) -> list[dict]: + rows = await _benchmark_rows(db, T0 - timedelta(days=1), T0 + timedelta(days=2), key, user_id) + return [row for row in rows if row["turns"] and (router is None or row["router_name"] == router)] + + async def _row(db, key: str, session_id: str = "s1", router: str = "auto-1") -> dict: rows = await db.query_raw( 'SELECT * FROM "LiteLLM_AutoRouterSession" WHERE api_key = $1 AND session_id = $2 AND router_name = $3', @@ -225,18 +244,15 @@ async def test_subtotal_coverage_survives_legacy_and_rolling_writers(db, writers assert row["savings_estimated_turns"] == sum(writers) assert row["savings_estimated_actual_spend"] == pytest.approx(0.01 * sum(writers)) assert row["savings_estimated_saved_spend"] == pytest.approx(0.02 * sum(writers)) - groups: Final = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, T0.isoformat(), (T0 + timedelta(days=1)).isoformat(), key, None - ) - assert len(groups) == 1 - assert groups[0]["classifier_cost"] == row["classifier_cost"] - assert groups[0]["classifier_cost_recorded_turns"] == sum(writers) - assert groups[0]["turns"] == len(writers) - assert groups[0]["spend"] == row["spend"] - assert groups[0]["saved_spend"] == row["saved_spend"] - assert groups[0]["savings_estimated_turns"] == sum(writers) - assert groups[0]["savings_estimated_actual_spend"] == row["savings_estimated_actual_spend"] - assert groups[0]["savings_estimated_saved_spend"] == row["savings_estimated_saved_spend"] + days: Final = await _days(db, key) + assert len(days) == int(any(writers)) + for day in days: + assert day["classifier_cost"] == row["classifier_cost"] + assert day["classifier_cost_recorded_turns"] == day["turns"] == sum(writers) + assert day["spend"] == pytest.approx(0.01 * sum(writers)) + assert day["saved_spend"] == pytest.approx(0.02 * sum(writers)) + assert day["savings_estimated_actual_spend"] == row["savings_estimated_actual_spend"] + assert day["savings_estimated_saved_spend"] == row["savings_estimated_saved_spend"] async def test_unknown_and_legacy_turns_preserve_actual_spend_without_entering_the_estimated_cohort(db: Prisma) -> None: @@ -250,13 +266,10 @@ async def test_unknown_and_legacy_turns_preserve_actual_spend_without_entering_t row: Final = await _row(db, key) assert row["saved_spend"] == pytest.approx(-0.03) assert row["savings_estimated_baseline_models"] == {"opus": 1} - groups: Final = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, T0.isoformat(), (T0 + timedelta(days=1)).isoformat(), key, None - ) - assert len(groups) == 1 - for actual in (row, groups[0]): - assert actual["turns"] == 3 - assert actual["spend"] == pytest.approx(0.96) + (day,) = await _days(db, key) + assert (row["turns"], day["turns"]) == (3, 2) + assert (row["spend"], day["spend"]) == (pytest.approx(0.96), pytest.approx(0.95)) + for actual in (row, day): assert actual["savings_estimated_turns"] == 1 assert actual["savings_estimated_actual_spend"] == pytest.approx(0.25) assert actual["savings_estimated_saved_spend"] == pytest.approx(-0.05) @@ -281,25 +294,20 @@ async def test_the_benchmarks_aggregate_reads_only_overlapping_sessions(db): await _turn(db, key, "A", T0, session_id=in_window, router=router, saved=0.5, spend=0.25, classifier_cost=0.02) await _turn(db, key, "A", T0 - timedelta(days=40), session_id=out_of_window, router=router, classifier_cost=9.0) - rows = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, - (T0 - timedelta(days=1)).isoformat(), - (T0 + timedelta(days=1)).isoformat(), - None, - None, - ) + rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None) matching = [row for row in rows if row["router_name"] == router] assert len(matching) == 1 grouped = matching[0] assert grouped["router_type"] == "complexity" assert grouped["sessions"] == 1 - assert grouped["turns"] == 2 - assert grouped["spend"] == pytest.approx(0.5) - assert grouped["saved_spend"] == pytest.approx(1.0) - assert grouped["classifier_cost"] == pytest.approx(0.03) - assert grouped["classifier_cost_recorded_turns"] == 2 + assert grouped["session_turns"] == 2 assert grouped["unordered_turns"] == 1 assert grouped["session_seconds"] == pytest.approx(60.0) + (day,) = await _days(db, router=router) + assert (day["turns"], day["classifier_cost_recorded_turns"]) == (2, 2) + assert day["spend"] == pytest.approx(0.5) + assert day["saved_spend"] == pytest.approx(1.0) + assert day["classifier_cost"] == pytest.approx(0.03) async def test_the_benchmarks_aggregate_can_filter_to_one_key(db): @@ -309,32 +317,22 @@ async def test_the_benchmarks_aggregate_can_filter_to_one_key(db): await _turn(db, first_key, "A", T0, router=router, saved=0.5, classifier_cost=0.01) await _turn(db, second_key, "A", T0, router=router, saved=9.0, classifier_cost=0.09) - rows = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, - (T0 - timedelta(days=1)).isoformat(), - (T0 + timedelta(days=1)).isoformat(), - first_key, - None, - ) + rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), first_key, None) matching = [row for row in rows if row["router_name"] == router] assert len(matching) == 1 assert matching[0]["sessions"] == 1 - assert matching[0]["saved_spend"] == pytest.approx(0.5) - assert matching[0]["classifier_cost"] == pytest.approx(0.01) - assert matching[0]["classifier_cost_recorded_turns"] == 1 + (day,) = await _days(db, first_key, router=router) + assert day["saved_spend"] == pytest.approx(0.5) + assert day["classifier_cost"] == pytest.approx(0.01) + assert day["classifier_cost_recorded_turns"] == 1 - unknown_key_rows = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, - (T0 - timedelta(days=1)).isoformat(), - (T0 + timedelta(days=1)).isoformat(), - f"k-{uuid.uuid4()}", - None, - ) + unknown_key_rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), f"k-{uuid.uuid4()}", None) assert [row for row in unknown_key_rows if row["router_name"] == router] == [] class _BenchmarkRow(TypedDict): sessions: ReadOnly[int] + session_turns: ReadOnly[int] turns: ReadOnly[int] same_model_turns: ReadOnly[int] first_visit_turns: ReadOnly[int] @@ -350,14 +348,11 @@ class _BenchmarkRow(TypedDict): async def _scoped_benchmarks( db: Prisma, router: str, user_id: str | None = None, key: str | None = None ) -> tuple[_BenchmarkRow, ...]: - rows: Final = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, - (T0 - timedelta(days=1)).isoformat(), - (T0 + timedelta(days=1)).isoformat(), - key, - user_id, + rows: Final = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), key, user_id) + days: Final = await _days(db, key, user_id, router) + return tuple( + cast(_BenchmarkRow, {**row, **next(iter(days), {})}) for row in rows if row["router_name"] == router ) - return tuple(cast(_BenchmarkRow, row) for row in rows if row["router_name"] == router) async def test_users_keep_written_identity_across_shared_keys_and_keyless_sessions(db: Prisma) -> None: @@ -384,28 +379,33 @@ async def test_users_keep_written_identity_across_shared_keys_and_keyless_sessio intersection: Final = await _scoped_benchmarks(db, router, user_id=alice, key=first_key) assert len(alice_rows) == len(bob_rows) == len(global_rows) == len(key_rows) == len(intersection) == 1 assert (alice_rows[0]["sessions"], alice_rows[0]["turns"], alice_rows[0]["same_model_turns"]) == (3, 4, 1) + assert (alice_rows[0]["session_turns"], bob_rows[0]["session_turns"]) == (4, 2) assert (bob_rows[0]["sessions"], bob_rows[0]["turns"], bob_rows[0]["first_visit_turns"]) == (2, 2, 2) assert alice_rows[0]["spend"] == pytest.approx(0.05) assert bob_rows[0]["spend"] == pytest.approx(0.07) assert alice_rows[0]["tier_turns"] == {"simple": 1} assert bob_rows[0]["tier_turns"] == {"complex": 1} assert (alice_rows[0]["cache_hits"], bob_rows[0]["cache_hits"]) == (1, 0) - assert (global_rows[0]["sessions"], global_rows[0]["turns"]) == (4, 7) + assert (global_rows[0]["sessions"], global_rows[0]["session_turns"], global_rows[0]["turns"]) == (4, 7, 6) assert (alice_rows[0]["savings_estimated_turns"], bob_rows[0]["savings_estimated_turns"]) == (4, 2) assert global_rows[0]["savings_estimated_turns"] == 6 for scoped in (alice_rows[0], bob_rows[0]): assert scoped["savings_estimated_actual_spend"] == pytest.approx(scoped["spend"]) assert scoped["savings_estimated_saved_spend"] == pytest.approx(scoped["saved_spend"]) - assert global_rows[0]["spend"] == pytest.approx(alice_rows[0]["spend"] + bob_rows[0]["spend"] + 0.01) - assert global_rows[0]["saved_spend"] == pytest.approx(alice_rows[0]["saved_spend"] + bob_rows[0]["saved_spend"] + 0.02) + assert global_rows[0]["spend"] == pytest.approx(alice_rows[0]["spend"] + bob_rows[0]["spend"]) + assert global_rows[0]["saved_spend"] == pytest.approx(alice_rows[0]["saved_spend"] + bob_rows[0]["saved_spend"]) assert global_rows[0]["tier_turns"] == {"simple": 1, "complex": 1} - assert (key_rows[0]["sessions"], key_rows[0]["turns"]) == (1, 3) - assert key_rows[0]["spend"] == pytest.approx(0.05) + assert (key_rows[0]["sessions"], key_rows[0]["session_turns"], key_rows[0]["turns"]) == (1, 3, 2) + assert key_rows[0]["spend"] == pytest.approx(0.04) assert (intersection[0]["sessions"], intersection[0]["turns"]) == (1, 1) assert intersection[0]["spend"] == pytest.approx(0.01) assert await _scoped_benchmarks(db, router, user_id=bob, key=second_key) == () assert await _scoped_benchmarks(db, router, user_id=f"u-{uuid.uuid4()}") == () - assert await _scoped_benchmarks(db, router, user_id="") == () + assert [ + row + for row in await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, "") + if row["router_name"] == router + ] == [] async def test_a_failed_user_projection_rolls_back_the_keys_increment(db: Prisma) -> None: @@ -419,6 +419,7 @@ async def test_a_failed_user_projection_rolls_back_the_keys_increment(db: Prisma assert await _row(db, key) == before assert await db.query_raw('SELECT user_id FROM "LiteLLM_AutoRouterUserSession" WHERE user_id = $1', user_id) == [] + assert [day["turns"] for day in await _days(db, key)] == [1] first_user: Final = f"u-{uuid.uuid4()}" second_user: Final = f"u-{uuid.uuid4()}" @@ -463,6 +464,14 @@ async def test_a_failed_user_projection_rolls_back_the_keys_increment(db: Prisma assert (row["turns"], row["same_model_turns"], row["unordered_turns"], row["last_model"]) == (count, 1, 0, model) assert row["spend"] == pytest.approx(count * 0.01) assert row["saved_spend"] == pytest.approx(count * 0.02) + days: Final = await db.query_raw( + 'SELECT user_id, turns, saved_spend FROM "LiteLLM_AutoRouterDailySpend" WHERE api_key = $1', key + ) + assert {day["user_id"]: (day["turns"], day["saved_spend"]) for day in days} == { + "": (1, pytest.approx(0.02)), + first_user: (3, pytest.approx(0.06)), + second_user: (2, pytest.approx(0.04)), + } async def test_user_session_cleanup_keeps_another_users_recent_keyless_session(db: Prisma) -> None: @@ -490,13 +499,7 @@ async def test_a_reconfigured_alias_reports_each_router_type_as_its_own_group(db db, key, "A", T0 + timedelta(seconds=10), session_id=f"s-{uuid.uuid4()}", router=router, router_type="quality" ) - rows = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, - (T0 - timedelta(days=1)).isoformat(), - (T0 + timedelta(days=1)).isoformat(), - None, - None, - ) + rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None) matching = sorted( (row for row in rows if row["router_name"] == router), key=lambda row: row["router_type"], @@ -575,16 +578,10 @@ async def test_the_benchmarks_aggregate_sums_tier_turns_across_sessions(db): await _turn(db, key, "B", T0 + timedelta(seconds=20), session_id=f"s-{uuid.uuid4()}", router=router, tier="complex") await _turn(db, key, "C", T0 + timedelta(seconds=30), session_id=f"s-{uuid.uuid4()}", router=router, tier=None) - rows = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, - (T0 - timedelta(days=1)).isoformat(), - (T0 + timedelta(days=1)).isoformat(), - None, - None, - ) + rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None) grouped = next(row for row in rows if row["router_name"] == router) assert grouped["tier_turns"] == {"simple": 2, "complex": 1} - assert grouped["turns"] == 4 + assert grouped["session_turns"] == 4 async def test_tier_maps_stay_separate_per_router_type_on_a_reconfigured_alias(db): @@ -604,13 +601,7 @@ async def test_tier_maps_stay_separate_per_router_type_on_a_reconfigured_alias(d tier="2", ) - rows = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, - (T0 - timedelta(days=1)).isoformat(), - (T0 + timedelta(days=1)).isoformat(), - None, - None, - ) + rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None) by_type = {row["router_type"]: row["tier_turns"] for row in rows if row["router_name"] == router} assert by_type == {"complexity": {"medium": 1}, "quality": {"2": 1}} @@ -620,13 +611,7 @@ async def test_a_window_with_no_tiered_turns_aggregates_to_an_empty_map(db): router = f"r-{uuid.uuid4()}" await _turn(db, key, "A", T0, session_id=f"s-{uuid.uuid4()}", router=router, tier=None) - rows = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, - (T0 - timedelta(days=1)).isoformat(), - (T0 + timedelta(days=1)).isoformat(), - None, - None, - ) + rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None) grouped = next(row for row in rows if row["router_name"] == router) assert grouped["tier_turns"] == {} @@ -653,3 +638,49 @@ async def test_an_out_of_order_hit_still_counts_toward_the_overall_hit_rate(db): assert row["unordered_turns"] == 1 assert row["cache_hits"] == 1 assert row["same_model_hits"] + row["first_visit_hits"] + row["return_hits"] == 0 + + +async def test_a_cross_midnight_session_splits_its_money_by_request_day(db): + key = f"k-{uuid.uuid4()}" + router = f"auto-{uuid.uuid4()}" + midnight = datetime(2026, 9, 2) + await _turn(db, key, "A", midnight - timedelta(minutes=10), router=router, spend=1.0, saved=7.0, user_id="u1") + await _turn(db, key, "A", midnight + timedelta(minutes=10), router=router, spend=1.0, saved=3.0, user_id="u1") + await _turn(db, key, "B", midnight + timedelta(days=1), router=router, spend=1.0, saved=11.0, user_id="u1") + + assert (await _row(db, key, router=router))["saved_spend"] == 21.0 + days = await db.query_raw( + 'SELECT date, turns, saved_spend FROM "LiteLLM_AutoRouterDailySpend" WHERE api_key = $1 ORDER BY date', key + ) + assert [(d["date"], d["turns"], d["saved_spend"]) for d in days] == [ + ("2026-09-01", 1, 7.0), + ("2026-09-02", 1, 3.0), + ("2026-09-03", 1, 11.0), + ] + for user_id in (None, "u1"): + (selected,) = await _benchmark_rows(db, midnight, midnight + timedelta(days=1), key, user_id) + assert (selected["sessions"], selected["session_turns"]) == (1, 3) + assert (selected["turns"], selected["spend"], selected["saved_spend"]) == (1, 1.0, 3.0) + + +async def test_a_router_type_change_within_a_day_keeps_each_types_money_apart(db): + key = f"k-{uuid.uuid4()}" + router = f"auto-{uuid.uuid4()}" + await _turn(db, key, "A", T0, router=router, router_type="complexity", spend=1.0, saved=4.0) + await _turn(db, key, "A", T0 + timedelta(hours=1), router=router, router_type="quality", spend=2.0, saved=0.0) + + days = {day["router_type"]: (day["turns"], day["spend"], day["saved_spend"]) for day in await _days(db, key)} + assert days == {"complexity": (1, 1.0, 4.0), "quality": (1, 2.0, 0.0)} + + +async def test_a_router_type_change_mid_session_keeps_session_shape_with_the_sessions_type(db): + key = f"k-{uuid.uuid4()}" + router = f"auto-{uuid.uuid4()}" + await _turn(db, key, "A", T0, router=router, router_type="complexity", spend=1.0, saved=4.0) + await _turn(db, key, "B", T0 + timedelta(hours=1), router=router, router_type="quality", spend=2.0, saved=0.0) + + rows = {row["router_type"]: row for row in await _benchmark_rows(db, T0, T0 + timedelta(days=1), key)} + assert set(rows) == {"complexity", "quality"} + assert (rows["complexity"]["sessions"], rows["complexity"]["session_turns"], rows["complexity"]["turns"]) == (1, 2, 1) + assert (rows["quality"]["sessions"], rows["quality"]["session_turns"], rows["quality"]["turns"]) == (0, 0, 1) + assert rows["quality"]["spend"] == 2.0 diff --git a/tests/proxy_behavior/spend/test_baseline_accounting.py b/tests/proxy_behavior/spend/test_baseline_accounting.py index 3504751d132..8fb82d0c80e 100644 --- a/tests/proxy_behavior/spend/test_baseline_accounting.py +++ b/tests/proxy_behavior/spend/test_baseline_accounting.py @@ -151,6 +151,13 @@ async def test_late_replay_updates_all_projections_without_rebilling(db: Prisma, ): assert after_users["late-user"][field] == after[field] assert after_users["late-user"]["turns"] == 1 and after_users["late-user"]["spend"] == 0.17 + days: Final = await db.query_raw( + 'SELECT * FROM "LiteLLM_AutoRouterDailySpend" WHERE api_key=$1 ORDER BY user_id', late.api_key + ) + assert [(day["date"], day["user_id"]) for day in days] == [("1970-01-01", "early-user"), ("1970-01-01", "late-user")] + assert days[0]["saved_spend"] == days[0]["savings_estimated_turns"] == 0 + for field in ("saved_spend", "savings_estimated_turns", "savings_estimated_actual_spend", "savings_estimated_saved_spend"): + assert days[1][field] == after[field] for table in ("DailyUserSpend", "DailyTeamSpend", "DailyOrganizationSpend", "DailyEndUserSpend", "DailyAgentSpend", "DailyTagSpend"): rows: Final = await db.query_raw(f'SELECT spend,api_requests,autorouter_savings_spend FROM "LiteLLM_{table}" WHERE api_key=$1', late.api_key) assert rows[0]["spend"] == rows[0]["api_requests"] == 0 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 2cbba9da8b3..01e41e8b03f 100644 --- a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py @@ -619,6 +619,18 @@ def test_classifier_plugin_is_not_settable_over_http(): _request("what is 2+2", classifier_type="custom", classifier_plugin="my_module.instance") +def _benchmark_db(rows: Sequence[Mapping[str, object]], recorded: float | None = None) -> SimpleNamespace: + """The joined benchmark statement returns the rows as given; any other statement is the Overall total.""" + from litellm.proxy.db.autorouter_session_rollup import AUTOROUTER_BENCHMARKS_SQL + + total: Final = recorded if recorded is not None else sum(float(row.get("saved_spend") or 0.0) for row in rows) + + async def query_raw(sql: str, *params: object) -> Sequence[Mapping[str, object]]: + return rows if sql == AUTOROUTER_BENCHMARKS_SQL else ({"saved": total},) + + return SimpleNamespace(db=SimpleNamespace(query_raw=AsyncMock(side_effect=query_raw))) + + class TestAutoRouterBenchmarks: from litellm.proxy.management_endpoints.auto_router_endpoints import _SessionAggRow @@ -635,15 +647,12 @@ class TestAutoRouterBenchmarks: rows: Sequence[Mapping[str, object]], model_list: Sequence[object], api_key: str | None = None, + recorded: float | None = None, ) -> AutoRouterBenchmarksResponse: from litellm.proxy import proxy_server from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks - class _DB: - async def query_raw(self, sql: str, *params: object): - return rows - - monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})()) + monkeypatch.setattr(proxy_server, "prisma_client", _benchmark_db(rows, recorded)) monkeypatch.setattr(proxy_server, "llm_router", type("R", (), {"model_list": model_list})()) return await get_auto_router_benchmarks( user_api_key_dict=ADMIN, @@ -657,6 +666,7 @@ class TestAutoRouterBenchmarks: router_type="complexity", tier_turns={}, sessions=4, + session_turns=40, turns=40, unordered_turns=1, covered_turns=38, @@ -703,7 +713,6 @@ class TestAutoRouterBenchmarks: assert totals.baseline_spend == 40.0 assert totals.saved_pct == 75.0 assert totals.savings_estimated_classifier_cost == 0.4 - assert totals.saved_per_session == 7.5 assert totals.cache.coverage_pct == 95.0 assert totals.cache.hit_rate_pct == pytest.approx(73.7) assert totals.cache.same_model.hit_rate_pct == 95.0 @@ -742,7 +751,6 @@ class TestAutoRouterBenchmarks: 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)]) @@ -772,9 +780,50 @@ class TestAutoRouterBenchmarks: 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.unattributed_saved_spend is None + assert (totals.saved_spend, totals.baseline_spend, totals.saved_pct) == ( + (30.0, 40.0, 75.0) if saved == 0.0 else (32.0, None, None) + ) assert totals.savings_estimated_classifier_cost == 0.4 + @pytest.mark.asyncio + @pytest.mark.parametrize("recorded, unattributed", [(30.0, None), (33.0, 3.0), (27.0, -3.0)]) + async def test_the_headline_is_the_overall_daily_total_and_untracked_savings_void_the_baseline( + self, recorded: float, unattributed: float | None, monkeypatch: pytest.MonkeyPatch + ) -> None: + response: Final = await self._benchmarks( + monkeypatch, rows=[self.ROW.model_dump()], model_list=[], recorded=recorded + ) + totals: Final = response.totals + assert (totals.saved_spend, totals.unattributed_saved_spend) == (recorded, unattributed) + assert (totals.baseline_spend, totals.saved_pct) == ((40.0, 75.0) if unattributed is None else (None, None)) + group: Final = response.groups[0] + assert group.saved_spend == 30.0 + assert (group.baseline_spend, group.saved_pct) == ((40.0, 75.0) if unattributed is None else (None, None)) + + @pytest.mark.asyncio + async def test_a_window_holding_only_untracked_history_shows_no_router_baseline( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + history_only: Final = self.ROW.model_dump( + exclude={ + "turns", + "spend", + "saved_spend", + "savings_estimated_turns", + "savings_estimated_actual_spend", + "savings_estimated_classifier_cost", + "savings_estimated_saved_spend", + "classifier_cost", + "classifier_cost_recorded_turns", + } + ) + response: Final = await self._benchmarks(monkeypatch, rows=[history_only], model_list=[], recorded=3.0) + assert (response.totals.saved_spend, response.totals.unattributed_saved_spend) == (3.0, 3.0) + group: Final = response.groups[0] + assert (group.sessions, group.turns, group.saved_spend) == (4, 0, 0.0) + assert (group.baseline_spend, group.saved_pct) == (None, None) + def test_an_empty_window_folds_to_zeros(self): from litellm.proxy.management_endpoints.auto_router_endpoints import ( _benchmark_totals, @@ -805,7 +854,7 @@ class TestAutoRouterBenchmarks: "savings_estimated_classifier_cost": 0.0, } ) - summed = _summed_agg_row([self.ROW, other]) + summed = _summed_agg_row([self.ROW, other.model_copy(update={"session_turns": 10})]) totals = _benchmark_totals(summed) assert summed.sessions == 5 assert summed.turns == 50 @@ -868,6 +917,28 @@ class TestAutoRouterBenchmarks: assert response.status_code == 422 query.assert_not_awaited() + @pytest.mark.asyncio + async def test_an_empty_key_filter_is_rejected_before_querying_deployment_data( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + import httpx + from fastapi import FastAPI + + from litellm.proxy import proxy_server + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks + + query: Final = AsyncMock(return_value=[]) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(query_raw=query))) + app: Final = FastAPI() + app.get("/auto_router/benchmarks")(get_auto_router_benchmarks) + app.dependency_overrides[user_api_key_auth] = lambda: ADMIN + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client: + response: Final = await client.get("/auto_router/benchmarks", params={"api_key": ""}) + + assert response.status_code == 422 + query.assert_not_awaited() + @pytest.mark.asyncio async def test_a_reversed_window_is_rejected(self, monkeypatch: pytest.MonkeyPatch): from litellm.proxy import proxy_server @@ -891,15 +962,8 @@ class TestAutoRouterBenchmarks: from litellm.proxy import proxy_server from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks - captured: dict = {} - - class _DB: - async def query_raw(self, sql: str, *params: object): - captured["sql"] = sql - captured["params"] = params - return [TestAutoRouterBenchmarks.ROW.model_dump()] - - monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})()) + prisma_client: Final = _benchmark_db([TestAutoRouterBenchmarks.ROW.model_dump()]) + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) response = await get_auto_router_benchmarks( user_api_key_dict=UserAPIKeyAuth(user_role=role, api_key="sk-admin", user_id="viewer"), @@ -908,7 +972,11 @@ class TestAutoRouterBenchmarks: api_key="key-hash", user_id=user_id, ) - assert captured["params"] == ("2026-07-01T00:00:00", "2026-08-02T00:00:00", "key-hash", user_id) + params: Final = tuple(call.args[1:] for call in prisma_client.db.query_raw.await_args_list) + assert params == ( + ("2026-07-01T00:00:00", "2026-08-02T00:00:00", "key-hash", user_id, "2026-07-01", "2026-08-01"), + ("2026-07-01", "2026-08-01", *(([user_id],) if user_id else ()), ["key-hash"]), + ) assert response.routers_in_scope == 1 assert response.groups[0].router_name == "live-auto" assert response.groups[0].saved_pct == response.totals.saved_pct == 75.0 @@ -946,7 +1014,6 @@ class TestAutoRouterBenchmarks: assert response.totals.saved_spend == 29.5 assert response.totals.baseline_spend == 41.5 assert response.totals.saved_pct == 71.1 - assert response.totals.saved_per_session == 5.9 @pytest.mark.asyncio @pytest.mark.parametrize( @@ -958,11 +1025,9 @@ class TestAutoRouterBenchmarks: from litellm.proxy import proxy_server from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks - class _DB: - async def query_raw(self, sql: str, *params: object): - return [{**TestAutoRouterBenchmarks.ROW.model_dump(), "tier_turns": wire_value}] - - monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})()) + monkeypatch.setattr( + proxy_server, "prisma_client", _benchmark_db([{**TestAutoRouterBenchmarks.ROW.model_dump(), "tier_turns": wire_value}]) + ) response = await get_auto_router_benchmarks( user_api_key_dict=ADMIN, @@ -1006,7 +1071,7 @@ class TestAutoRouterBenchmarks: 0.0, 0.0, ) - assert (idle.saved_pct, idle.saved_per_session, idle.avg_turns_per_session) == (0.0, 0.0, 0.0) + assert (idle.saved_pct, idle.avg_turns_per_session) == (0.0, 0.0) assert (idle.cache.hit_rate_pct, idle.cache.coverage_pct) == (0.0, 0.0) assert idle.cache.same_model.turns == idle.cache.return_to_tier.hits == 0 assert idle.tier_turns == {} @@ -3724,3 +3789,18 @@ async def test_availability_waits_for_the_first_complete_catalog(monkeypatch): with pytest.raises(HTTPException) as error: await auto_router_endpoints.get_auto_router_availability(AutoRouterAvailabilityRequest(), ADMIN) assert error.value.status_code == 503 + + +class TestPerSessionAverages: + @pytest.mark.parametrize( + "sessions, turns, expected", + [(4, 40, (10.0, 100.0, 1000.0)), (0, 0, (0.0, 0.0, 0.0)), (0, 3, (None, None, None))], + ) + def test_requests_without_session_rows_have_unknown_averages_not_zero( + self, sessions: int, turns: int, expected: tuple[float | None, ...] + ) -> None: + from litellm.proxy.management_endpoints.auto_router_endpoints import _benchmark_totals + + row: Final = TestAutoRouterBenchmarks.ROW.model_copy(update={"sessions": sessions, "turns": turns}) + totals: Final = _benchmark_totals(row) + assert (totals.avg_turns_per_session, totals.avg_session_seconds, totals.avg_tokens_per_session) == expected diff --git a/tests/unit/proxy/test_spend_log_cleanup.py b/tests/unit/proxy/test_spend_log_cleanup.py index 46ac1234615..399c76d97c1 100644 --- a/tests/unit/proxy/test_spend_log_cleanup.py +++ b/tests/unit/proxy/test_spend_log_cleanup.py @@ -796,19 +796,23 @@ async def test_spend_logs_retention_alone_does_not_touch_the_session_rollup(): assert any('"LiteLLM_SpendLogs"' in sql for sql in tables) assert not any('"LiteLLM_AutoRouterSession"' in sql for sql in tables) assert not any('"LiteLLM_AutoRouterUserSession"' in sql for sql in tables) + assert not any('"LiteLLM_AutoRouterDailySpend"' in sql for sql in tables) assert not any('"LiteLLM_HealthCheckTable"' in sql for sql in tables) @pytest.mark.asyncio -async def test_session_retention_alone_cleans_both_session_rollups(): - client = _mock_prisma_for_retention([0, 0]) +async def test_session_retention_alone_cleans_both_session_rollups_and_the_daily_rollup(): + client = _mock_prisma_for_retention([0, 0, 0]) cleaner = SpendLogCleanup(general_settings={"maximum_autorouter_session_retention_period": "365d"}) cleaner.pod_lock_manager = None await cleaner.cleanup_old_spend_logs(client) - tables = [call[0][0] for call in client.db.execute_raw.call_args_list] - assert len(tables) == 2 + calls = client.db.execute_raw.call_args_list + tables = [call[0][0] for call in calls] + assert len(tables) == 3 assert '"LiteLLM_AutoRouterSession"' in tables[0] assert '"LiteLLM_AutoRouterUserSession"' in tables[1] + assert '"LiteLLM_AutoRouterDailySpend"' in tables[2] + assert calls[2][0][1] == calls[0][0][1].date().isoformat() @pytest.mark.asyncio @@ -852,7 +856,7 @@ async def test_spend_logs_retention_alone_keeps_daily_tag_spend_forever(): @pytest.mark.asyncio async def test_each_retention_key_cuts_off_at_its_own_horizon(): - client = _mock_prisma_for_retention([0, 0, 0, 0, 0]) + client = _mock_prisma_for_retention([0, 0, 0, 0, 0, 0]) cleaner = SpendLogCleanup( general_settings={ "maximum_spend_logs_retention_period": "7d", @@ -868,6 +872,8 @@ async def test_each_retention_key_cuts_off_at_its_own_horizon(): if '"LiteLLM_AutoRouterSession"' in call[0][0] else "LiteLLM_AutoRouterUserSession" if '"LiteLLM_AutoRouterUserSession"' in call[0][0] + else "LiteLLM_AutoRouterDailySpend" + if '"LiteLLM_AutoRouterDailySpend"' in call[0][0] else "LiteLLM_HealthCheckTable" if '"LiteLLM_HealthCheckTable"' in call[0][0] else "logs" @@ -878,6 +884,7 @@ async def test_each_retention_key_cuts_off_at_its_own_horizon(): assert (now - cutoffs["logs"]).days == 7 assert (now - cutoffs["LiteLLM_AutoRouterSession"]).days == 365 assert cutoffs["LiteLLM_AutoRouterUserSession"] == cutoffs["LiteLLM_AutoRouterSession"] + assert cutoffs["LiteLLM_AutoRouterDailySpend"] == cutoffs["LiteLLM_AutoRouterSession"].date().isoformat() assert (now - cutoffs["LiteLLM_HealthCheckTable"]).days == 30 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 b4b34dbfaf3..081eb7f6e09 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 @@ -75,7 +75,6 @@ const totals = (overrides: Partial = {}): Totals => ({ saved_spend: 2174.59, baseline_spend: 2534.45, saved_pct: 85.8, - saved_per_session: 23.13, cache: cache(), ...overrides, }); @@ -110,7 +109,6 @@ const zeroTotals: Totals = { saved_spend: 0, baseline_spend: 0, saved_pct: 0, - saved_per_session: 0, cache: zeroCache, }; @@ -173,7 +171,6 @@ describe("AutoRouterBenchmarksTab", () => { saved_spend: saved, baseline_spend: estimatedTurns ? actual + (saved ?? 0) : null, saved_pct: pct, - saved_per_session: null, }; mockHook({ data: response([], totals(comparison)), @@ -204,18 +201,15 @@ describe("AutoRouterBenchmarksTab", () => { } }); - it("leads with total estimated savings, before the four session-shape metrics", () => { + it("leads with total estimated savings, before the three session-shape metrics", () => { mockHook({ data: response([group(), group({ router_name: "gpt-auto" })]) }); renderTab(); const labels = screen - .getAllByText( - /Total estimated savings|Avg saved per session|Avg turns per session|Avg session length|Avg tokens per session/, - ) + .getAllByText(/Total estimated savings|Avg turns per session|Avg session length|Avg tokens per session/) .map((node) => node.textContent); expect(labels).toEqual([ "Total estimated savings", - "Avg saved per session", "Avg turns per session", "Avg session length", "Avg tokens per session", @@ -271,15 +265,40 @@ describe("AutoRouterBenchmarksTab", () => { }, ); - it("pairs the savings with the session count it was earned over, in its own tile", () => { + it("labels selected-day money apart from whole-session metrics, with no savings-per-session tile", () => { mockHook({ data: response([group(), group({ router_name: "gpt-auto" })]) }); renderTab(); - const tile = screen.getByText("Avg saved per session").closest('[data-slot="card"]'); - if (!tile) throw new Error("expected avg saved per session to render as a metric tile"); - - expect(within(tile).getByText("$23.13")).toBeInTheDocument(); + const tile = screen.getByText("Avg turns per session").closest('[data-slot="card"]'); + if (!tile) throw new Error("expected avg turns per session to render as a metric tile"); expect(within(tile).getByText("· 94 sessions")).toBeInTheDocument(); + expect(screen.queryByText("Avg saved per session")).not.toBeInTheDocument(); + expect(screen.getByText(/Savings and spend count requests on the selected UTC days/)).toBeInTheDocument(); + expect(screen.getByText(/Session metrics cover every session that overlaps the range/)).toBeInTheDocument(); + }); + + it("shows session averages as unavailable, not zero, when routed requests have no session rows", () => { + const noSessions = { + sessions: 0, + avg_turns_per_session: null, + avg_session_seconds: null, + avg_tokens_per_session: null, + }; + mockHook({ data: response([], totals(noSessions)) }); + renderTab(); + + expect(screen.getAllByText("Unavailable")).toHaveLength(3); + expect(screen.queryByText("0.0")).not.toBeInTheDocument(); + }); + + it.each([3, -3])("explains a %s gap between router records and recorded savings instead of comparing", (gap) => { + const residual = { saved_spend: 5, unattributed_saved_spend: gap, baseline_spend: null, saved_pct: null }; + mockHook({ data: response([], totals(residual)) }); + renderTab(); + + expect(screen.getByText("$5.00")).toBeInTheDocument(); + expect(screen.getByText(/Per-router records differ from recorded savings by \$3\.00/)).toBeInTheDocument(); + expect(screen.getByText("Estimated baseline spend").nextSibling?.textContent).toBe("Unavailable"); }); it("exposes each spend row as a term and its value, not as loose text", () => { @@ -431,7 +450,7 @@ describe("AutoRouterBenchmarksTab", () => { renderTab(); expect(screen.getByText("Total estimated savings")).toBeInTheDocument(); - expect(screen.getAllByText("$0.00")).toHaveLength(6); + expect(screen.getAllByText("$0.00")).toHaveLength(5); expect(screen.getByText("· 0 sessions")).toBeInTheDocument(); expect(screen.getByText("0s")).toBeInTheDocument(); expect(screen.getByText(/turns measured/)).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 24a97587e32..27e8df6db87 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 @@ -105,6 +105,12 @@ const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => { adaptive and quality routers are excluded

)} + {stats.unattributed_saved_spend != null && ( +

+ Per-router records differ from recorded savings by {usd(Math.abs(stats.unattributed_saved_spend))}, for + example history from before per-router tracking, so the baseline comparison is unavailable +

+ )}
@@ -297,22 +303,34 @@ const BenchmarksBody: React.FC = ({ isPending, error, data, -
- - - - -
+

+ Savings and spend count requests on the selected UTC days. 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. +

- 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. + Session metrics cover every session that overlaps the range, including its turns outside the range.

+
+ + + +
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx index e4417d77463..42444fd8f06 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx @@ -12,21 +12,21 @@ vi.mock("@/components/shared/charts", () => ({ })); import TierTurnsChart, { tierDisplayLabel } from "./TierTurnsChart"; -import type { AutoRouterBenchmarkGroup, BenchmarkView } from "./autoRouterBenchmarks"; +import type { AutoRouterBenchmarkGroup, AutoRouterBenchmarkTotals, BenchmarkView } from "./autoRouterBenchmarks"; -const totalsOnly = { +const totalsOnly: AutoRouterBenchmarkTotals = { sessions: 3, turns: 9, avg_turns_per_session: 3, avg_session_seconds: 60, avg_tokens_per_session: 100, spend: 1, + classifier_cost: 0, savings_estimated_turns: 9, savings_estimated_actual_spend: 1, saved_spend: 1, baseline_spend: 2, saved_pct: 50, - saved_per_session: 0.33, cache: { coverage_pct: 0, hit_rate_pct: 0, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts index 0586163e77e..62d0c0e4e5e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts @@ -43,7 +43,6 @@ const totals = (overrides: Partial = {}) => ({ saved_spend: 2174.59, baseline_spend: 2534.45, saved_pct: 85.8, - saved_per_session: 23.13, cache: cache(), ...overrides, }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx index 676bc7d29eb..4195a3c9913 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx @@ -347,7 +347,6 @@ const routerUsageResponse = (saved: number): AutoRouterBenchmarksResponse => ({ saved_spend: saved, baseline_spend: 10 + saved, saved_pct: (100 * saved) / (10 + saved), - saved_per_session: saved / 2, cache: { coverage_pct: 100, hit_rate_pct: 0, 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 5c7b9b28789..2ce585e895e 100644 --- a/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx @@ -39,7 +39,6 @@ const stats = { saved_spend: 8.75, baseline_spend: 10, saved_pct: 87.5, - saved_per_session: 4.375, cache, }; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 1d50f37e573..fc1a9c67946 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -1350,9 +1350,10 @@ export interface paths { * * 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 + * internal user when written; older key-only history remains outside user views. Money counts + * only requests on the selected UTC days, and the all-router savings headline is the same daily + * total the Overall view reads. Session shape and caching cover every session that overlaps the + * window, whole. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is * over that bucket's turns. * * The rollup supplies the measures, never the list. Which routers appear comes from the @@ -26169,12 +26170,21 @@ export interface components { * @description One auto-router's slice of the benchmarks. */ AutoRouterBenchmarkGroup: { - /** Avg Session Seconds */ - avg_session_seconds: number; - /** Avg Tokens Per Session */ - avg_tokens_per_session: number; - /** Avg Turns Per Session */ - avg_turns_per_session: number; + /** + * Avg Session Seconds + * @description Lifetime seconds per overlapping session; null as above + */ + avg_session_seconds: number | null; + /** + * Avg Tokens Per Session + * @description Lifetime tokens per overlapping session; null as above + */ + avg_tokens_per_session: number | null; + /** + * Avg Turns Per Session + * @description Lifetime turns per overlapping session; null when the window has routed requests but no session rows for this router type, such as an alias whose router type changed mid-session + */ + avg_turns_per_session: number | null; /** * Baseline Spend * @description Estimated single-model cost: compared actual spend plus recorded savings; null when traffic has no recorded savings @@ -26201,14 +26211,9 @@ export interface components { * @description Recorded savings over baseline_spend, as a percentage */ saved_pct: number | null; - /** - * Saved Per Session - * @description Recorded savings per session, including historical estimates - */ - saved_per_session: number | null; /** * Saved Spend - * @description Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates + * @description Recorded savings on the selected UTC days; null when traffic has no recorded savings estimates. On totals this is the same daily figure the Overall savings view reports */ saved_spend: number | null; /** @@ -26226,11 +26231,14 @@ export interface components { * @description Requests compared against the baseline: every request on complexity routers that recorded savings */ savings_estimated_turns: number; - /** Sessions */ + /** + * Sessions + * @description Sessions overlapping the window, counted whole + */ sessions: number; /** * Spend - * @description What the routed traffic actually cost + * @description What the selected days' routed traffic actually cost */ spend: number; /** @@ -26240,20 +26248,38 @@ export interface components { tier_turns?: { [key: string]: number; }; - /** Turns */ + /** + * Turns + * @description Auto-routed requests on the selected UTC days + */ turns: number; + /** + * Unattributed Saved Spend + * @description Part of saved_spend no router's daily rows account for, such as history recorded before per-router daily tracking; when set, baseline_spend and saved_pct are null + */ + unattributed_saved_spend?: number | null; }; /** * AutoRouterBenchmarkTotals - * @description Session-shape and savings aggregates over auto-routed traffic in the window. + * @description Auto-routed traffic in the window. Turns, spend and savings count requests on the selected UTC days; + * the session averages and cache stats describe every session overlapping the window, whole. */ AutoRouterBenchmarkTotals: { - /** Avg Session Seconds */ - avg_session_seconds: number; - /** Avg Tokens Per Session */ - avg_tokens_per_session: number; - /** Avg Turns Per Session */ - avg_turns_per_session: number; + /** + * Avg Session Seconds + * @description Lifetime seconds per overlapping session; null as above + */ + avg_session_seconds: number | null; + /** + * Avg Tokens Per Session + * @description Lifetime tokens per overlapping session; null as above + */ + avg_tokens_per_session: number | null; + /** + * Avg Turns Per Session + * @description Lifetime turns per overlapping session; null when the window has routed requests but no session rows for this router type, such as an alias whose router type changed mid-session + */ + avg_turns_per_session: number | null; /** * Baseline Spend * @description Estimated single-model cost: compared actual spend plus recorded savings; null when traffic has no recorded savings @@ -26270,14 +26296,9 @@ export interface components { * @description Recorded savings over baseline_spend, as a percentage */ saved_pct: number | null; - /** - * Saved Per Session - * @description Recorded savings per session, including historical estimates - */ - saved_per_session: number | null; /** * Saved Spend - * @description Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates + * @description Recorded savings on the selected UTC days; null when traffic has no recorded savings estimates. On totals this is the same daily figure the Overall savings view reports */ saved_spend: number | null; /** @@ -26295,19 +26316,30 @@ export interface components { * @description Requests compared against the baseline: every request on complexity routers that recorded savings */ savings_estimated_turns: number; - /** Sessions */ + /** + * Sessions + * @description Sessions overlapping the window, counted whole + */ sessions: number; /** * Spend - * @description What the routed traffic actually cost + * @description What the selected days' routed traffic actually cost */ spend: number; - /** Turns */ + /** + * Turns + * @description Auto-routed requests on the selected UTC days + */ turns: number; + /** + * Unattributed Saved Spend + * @description Part of saved_spend no router's daily rows account for, such as history recorded before per-router daily tracking; when set, baseline_spend and saved_pct are null + */ + unattributed_saved_spend?: number | null; }; /** * AutoRouterBenchmarksResponse - * @description Benchmarks for the auto-router dashboard, aggregated from the per-session rollup. + * @description Benchmarks for the auto-router dashboard, aggregated from the per-session and per-day rollups. */ AutoRouterBenchmarksResponse: { /** From 2584721ca3e88dbbb4eaf4a79c00d3ead1b52832 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Fri, 2 Oct 2026 11:58:32 -0700 Subject: [PATCH 018/139] fix(lens): preserve framework agent names and GenAI message content (#44218) * fix(lens): use recorded agent identities across framework traces * fix(lens): tighten agent identity and bound trace lookups * style(tracing): wrap framework agent identity test case --- .../crates/traces/query/list_traces.sql | 22 ++- .../crates/traces/src/normalize/mod.rs | 72 +++++++- litellm-rust/crates/traces/src/otlp/span.rs | 13 +- .../crates/traces/tests/migrations.rs | 166 ++++++++++++++++++ litellm/tracing/store.py | 10 +- litellm/tracing/types.py | 1 + tests/test_litellm/tracing/test_decode.py | 83 ++++++++- tests/test_litellm/tracing/test_store.py | 17 ++ .../TraceView/AgentTracesSection.test.tsx | 11 +- .../TraceView/AgentTracesSection.tsx | 6 +- .../view_logs/TraceView/AgentTracesTable.tsx | 6 +- .../view_logs/TraceView/traceUtils.test.ts | 13 ++ .../view_logs/TraceView/traceUtils.ts | 23 ++- ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 + 14 files changed, 419 insertions(+), 26 deletions(-) diff --git a/litellm-rust/crates/traces/query/list_traces.sql b/litellm-rust/crates/traces/query/list_traces.sql index c0c1b28aa7f..163d9b03d6b 100644 --- a/litellm-rust/crates/traces/query/list_traces.sql +++ b/litellm-rust/crates/traces/query/list_traces.sql @@ -1,11 +1,13 @@ +WITH page AS ( SELECT TraceId AS trace_id, hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) AS trace_ref, TeamId AS team_id, ApiKeyHash AS api_key_hash, ifNull(any(RootName), '') AS name, any(ServiceName) AS service, ifNull(any(RootInput), '') AS input_preview, ifNull(any(RootStatus), '') AS status, toUnixTimestamp64Milli(min(StartTs)) AS start_ms, + min(StartTs) AS trace_start, max(EndTs) AS trace_end, dateDiff('millisecond', min(StartTs), max(EndTs)) AS duration_ms, - sum(SpanCount) AS span_count, length(groupUniqArrayArray(AgentNames)) AS agent_count, + sum(SpanCount) AS span_count, sum(AgentCount) AS agent_invocations, sum(LlmCount) AS llm_calls, sum(ToolCount) AS tool_calls, sum(InputTokens) AS input_tokens, sum(OutputTokens) AS output_tokens, @@ -21,3 +23,21 @@ HAVING min(StartTs) >= fromUnixTimestamp64Milli({start_ms:Int64}) < ({cursor_ms:Int64}, {cursor_trace_id:String})) ORDER BY start_ms DESC, trace_ref DESC LIMIT {limit:UInt32} +) +SELECT page.* EXCEPT (trace_start, trace_end), + identities.agent_names AS agent_names, identities.agent_count AS agent_count +FROM page +LEFT JOIN ( + SELECT TeamId, ApiKeyHash, TraceId, + arraySort(groupUniqArrayIf(AgentName, AgentName != '')) AS agent_names, + uniqExactIf(if(AgentName = '', SpanName, AgentName), ObservationType = 'agent') AS agent_count + FROM otel_traces + WHERE Timestamp >= (SELECT min(trace_start) FROM page) + AND Timestamp <= (SELECT max(trace_end) FROM page) + AND TraceId IN (SELECT trace_id FROM page) + AND (TeamId, ApiKeyHash, TraceId) IN (SELECT team_id, api_key_hash, trace_id FROM page) + GROUP BY TeamId, ApiKeyHash, TraceId +) AS identities +ON page.team_id = identities.TeamId AND page.api_key_hash = identities.ApiKeyHash + AND page.trace_id = identities.TraceId +ORDER BY page.start_ms DESC, page.trace_ref DESC diff --git a/litellm-rust/crates/traces/src/normalize/mod.rs b/litellm-rust/crates/traces/src/normalize/mod.rs index b4d5a4a9d06..60dc8f816f2 100644 --- a/litellm-rust/crates/traces/src/normalize/mod.rs +++ b/litellm-rust/crates/traces/src/normalize/mod.rs @@ -1,7 +1,7 @@ use std::collections::BTreeMap; use crate::DecodeError; -use serde::Serialize; +use serde::{Deserialize, Serialize}; #[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)] #[serde(rename_all = "lowercase")] @@ -148,6 +148,60 @@ fn usage_tokens(attributes: &BTreeMap) -> Result<(u32, u32), Dec )) } +#[derive(Default, Deserialize)] +struct AgentMetadata { + #[serde(default)] + lc_agent_name: String, + #[serde(default)] + ls_integration: String, +} + +fn recorded_agent_name( + name: &str, + attributes: &BTreeMap, + span: &NormalizedSpan, +) -> String { + let explicit = [ + span.agent_name.as_str(), + attr(attributes, "gen_ai.agent.name"), + attr(attributes, "agent.name"), + attr(attributes, "openclaw.agent"), + ] + .into_iter() + .find(|value| !value.is_empty()); + if let Some(value) = explicit { + return value.to_owned(); + } + let metadata = + serde_json::from_str::(attr(attributes, "metadata")).unwrap_or_default(); + if !metadata.lc_agent_name.is_empty() { + return metadata.lc_agent_name; + } + if span.observation_type == ObservationType::Agent { + let node = attr(attributes, "graph.node.id"); + if !node.is_empty() { + return node.to_owned(); + } + if metadata.ls_integration == "langgraph" && name != "LangGraph" && !is_middleware(name) { + return name.to_owned(); + } + } + String::new() +} + +fn is_middleware(name: &str) -> bool { + [ + ".wrap_model_call", + ".wrap_tool_call", + ".before_agent", + ".after_agent", + ".before_model", + ".after_model", + ] + .iter() + .any(|suffix| name.ends_with(suffix)) +} + pub fn normalize( scope_name: &str, name: &str, @@ -163,8 +217,22 @@ pub fn normalize( .into_iter() .find(|normalizer| normalizer.matches(scope_name, attributes)) .expect("GenAI fallback always matches"); + let span = normalizer.normalize(name, parent_span_id, attributes)?; + let agent_name = recorded_agent_name(name, attributes, &span); + let observation_type = if !parent_span_id.is_empty() + && scope_name == "openinference.instrumentation.langchain" + && is_middleware(name) + { + ObservationType::Framework + } else { + span.observation_type + }; Ok(Normalization { - span: normalizer.normalize(name, parent_span_id, attributes)?, + span: NormalizedSpan { + agent_name, + observation_type, + ..span + }, consumed_attributes: normalizer.consumed_attributes(attributes), }) } diff --git a/litellm-rust/crates/traces/src/otlp/span.rs b/litellm-rust/crates/traces/src/otlp/span.rs index 1c1e53e756c..0ae5725ba47 100644 --- a/litellm-rust/crates/traces/src/otlp/span.rs +++ b/litellm-rust/crates/traces/src/otlp/span.rs @@ -133,7 +133,18 @@ fn decoded_span( &parent_span_id, &span_attributes, )?; - let normalized = normalization.span; + let resource_agent_name = resource_attributes + .get("gen_ai.agent.name") + .filter(|name| !name.is_empty()); + let agent_name = match (resource_agent_name, normalization.span.agent_name.as_str()) { + (Some(name), "") => name.clone(), + (Some(name), "hermes-agent") if scope_name.as_ref() == "hermes-otel-plugin" => name.clone(), + (_, name) => name.to_owned(), + }; + let normalized = crate::normalize::NormalizedSpan { + agent_name, + ..normalization.span + }; budget.consume( normalized.input.len() + normalized.output.len() diff --git a/litellm-rust/crates/traces/tests/migrations.rs b/litellm-rust/crates/traces/tests/migrations.rs index ccefa8b0b9f..c62e3538ddb 100644 --- a/litellm-rust/crates/traces/tests/migrations.rs +++ b/litellm-rust/crates/traces/tests/migrations.rs @@ -362,6 +362,143 @@ async fn keyed_rollup_keeps_same_trace_ids_separate_by_api_key( Ok(()) } +#[rstest] +#[tokio::test] +async fn listed_agent_names_preserve_scope_and_cursor( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + for (team, key, trace, agent, span, parent) in [ + ("alpha", "one", "shared", "research_agent", "root", ""), + ("alpha", "one", "shared", "reviewer", "child", "root"), + ("alpha", "one", "shared", "reviewer", "repeated", "root"), + ("alpha", "one", "shared", "", "unnamed", "root"), + ("alpha", "one", "second", "support_agent", "root", ""), + ("alpha", "two", "shared", "private_agent", "root", ""), + ("beta", "one", "shared", "other_agent", "root", ""), + ] { + insert_rows( + &database, + "otel_traces", + vec![serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp, "TraceId": trace, "SpanId": span, "ParentSpanId": parent, + "ServiceName": "shared-app", "SpanName": span, "AgentName": agent, + "ObservationType": "agent", + "ResourceAttributes": {"litellm.team_id": team, "litellm.api_key_hash": key} + }))?], + ) + .await?; + } + let historical_rows = (0..5000) + .map(|index| { + serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp - 86_400_000_000_000_i64, + "TraceId": "shared", "SpanId": format!("historical-{index}"), + "ParentSpanId": "", "SpanName": "historical", "AgentName": "private_agent", + "ObservationType": "agent", "ServiceName": "shared-app", + "ResourceAttributes": {"litellm.team_id": "alpha", "litellm.api_key_hash": "history"} + })) + }) + .collect::, _>>()?; + insert_rows(&database, "otel_traces", historical_rows).await?; + let connection = Connection::configured(&database.url, "trace_test", "default", "")?; + let parameters = BTreeMap::from([ + ("team_ids".into(), Parameter::Strings(vec!["alpha".into()])), + ("api_key_hash".into(), Parameter::Text("one".into())), + ( + "start_ms".into(), + Parameter::Integer(timestamp / 1_000_000 - 1000), + ), + ( + "end_ms".into(), + Parameter::Integer(timestamp / 1_000_000 + 1000), + ), + ("cursor_ms".into(), Parameter::Integer(0)), + ("cursor_trace_id".into(), Parameter::Text(String::new())), + ("limit".into(), Parameter::Integer(1)), + ]); + let first: serde_json::Value = serde_json::from_str( + &execute_named_read( + &database.client, + &connection, + ReadQuery::ListTraces, + ¶meters, + ) + .await?, + )?; + let cursor = first["data"][0]["trace_ref"] + .as_str() + .ok_or("missing cursor")?; + let next_parameters = parameters + .into_iter() + .chain([ + ( + "cursor_ms".into(), + Parameter::Integer(timestamp / 1_000_000), + ), + ("cursor_trace_id".into(), Parameter::Text(cursor.into())), + ]) + .collect(); + let second: serde_json::Value = serde_json::from_str( + &execute_named_read( + &database.client, + &connection, + ReadQuery::ListTraces, + &next_parameters, + ) + .await?, + )?; + assert_eq!( + first["data"].as_array().ok_or("missing first page")?.len(), + 1 + ); + assert_eq!( + second["data"] + .as_array() + .ok_or("missing second page")? + .len(), + 1 + ); + assert_ne!(first["data"][0]["trace_id"], second["data"][0]["trace_id"]); + let names = [&first["data"][0], &second["data"][0]] + .into_iter() + .map(|row| { + ( + row["trace_id"].as_str().unwrap(), + row["agent_names"].clone(), + ) + }) + .collect::>(); + assert_eq!( + names["shared"], + serde_json::json!(["research_agent", "reviewer"]) + ); + assert_eq!(names["second"], serde_json::json!(["support_agent"])); + let counts = [&first["data"][0], &second["data"][0]] + .into_iter() + .map(|row| { + ( + row["trace_id"].as_str().unwrap(), + row["agent_count"].as_u64(), + ) + }) + .collect::>(); + assert_eq!(counts["shared"], Some(3)); + assert_eq!(counts["second"], Some(1)); + for page in [&first, &second] { + assert!( + page["statistics"]["rows_read"] + .as_u64() + .ok_or("missing read statistics")? + < 5000 + ); + } + Ok(()) +} + #[rstest] #[tokio::test] async fn rollup_merges_spans_across_days_without_losing_root_fields( @@ -376,6 +513,7 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields( let root = serde_json::from_value(serde_json::json!({ "Timestamp": day_start - 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-root", "ParentSpanId": "", "ServiceName": "proxy", "SpanName": "root", "Input": "root input", + "AgentName": "lead", "ObservationType": "agent", "StatusCode": "STATUS_CODE_ERROR", "ResourceAttributes": {"litellm.team_id": "team-1"} }))?; @@ -383,6 +521,7 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields( let child = serde_json::from_value(serde_json::json!({ "Timestamp": day_start + 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-child", "ParentSpanId": "span-root", "ServiceName": "proxy", "SpanName": "child", + "AgentName": "researcher", "ObservationType": "agent", "StatusCode": "STATUS_CODE_UNSET", "ResourceAttributes": {"litellm.team_id": "team-1"} }))?; @@ -406,6 +545,33 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields( "RootStatus": "STATUS_CODE_ERROR", "SpanCount": 2 }]) ); + let connection = Connection::configured(&database.url, "trace_test", "default", "")?; + let parameters = BTreeMap::from([ + ("team_ids".into(), Parameter::Strings(vec!["team-1".into()])), + ("api_key_hash".into(), Parameter::Text(String::new())), + ( + "start_ms".into(), + Parameter::Integer(day_start / 1_000_000 - 2000), + ), + ("end_ms".into(), Parameter::Integer(day_start / 1_000_000)), + ("cursor_ms".into(), Parameter::Integer(0)), + ("cursor_trace_id".into(), Parameter::Text(String::new())), + ("limit".into(), Parameter::Integer(10)), + ]); + let listed: serde_json::Value = serde_json::from_str( + &execute_named_read( + &database.client, + &connection, + ReadQuery::ListTraces, + ¶meters, + ) + .await?, + )?; + assert_eq!( + listed["data"][0]["agent_names"], + serde_json::json!(["lead", "researcher"]) + ); + assert_eq!(listed["data"][0]["agent_count"], 2); Ok(()) } diff --git a/litellm/tracing/store.py b/litellm/tracing/store.py index 91420ffd025..edfe1285fc7 100644 --- a/litellm/tracing/store.py +++ b/litellm/tracing/store.py @@ -119,6 +119,7 @@ def trace_summary_from_row(row: dict[str, Any], spend_rows: Sequence[_SpendRow] trace_ref=row.get("trace_ref", ""), name=row["name"], service=row["service"], + agent_names=tuple(row.get("agent_names") or ()), input_preview=row["input_preview"], start_time=_iso(int(row["start_ms"])), duration_ms=float(row["duration_ms"]), @@ -169,8 +170,8 @@ def _parent_agent_of(span: Span, by_id: Mapping[str, Span]) -> str | None: if parent_id is None or parent_id not in by_id or parent_id == span["span_id"]: return None parent = by_id[parent_id] - if parent["type"] == "agent" and parent["name"] != span["name"]: - return parent["name"] + if parent["type"] == "agent" and (parent["agent"] or parent["name"]) != (span["agent"] or span["name"]): + return parent["agent"] or parent["name"] parent_id = parent["parent_span_id"] return None @@ -183,9 +184,9 @@ def agent_nodes(spans: Sequence[Span]) -> tuple[AgentNode, ...]: if span["type"] != "agent": continue node = agents.setdefault( - span["name"], + span["agent"] or span["name"], AgentNode( - name=span["name"], + name=span["agent"] or span["name"], parent_agent=_parent_agent_of(span, by_id), invocations=0, llm_calls=0, @@ -250,6 +251,7 @@ def trace_from_rows( trace_ref=trace_ref, name=root["name"], service=rows[0]["service"], + agent_names=tuple(sorted(frozenset(s["agent"] for s in spans if s["agent"]))), input_preview=root["input_preview"], start_time=_iso(trace_start_ns // NANOS_PER_MS), duration_ms=(trace_end_ns - trace_start_ns) / NANOS_PER_MS, diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py index ff965483013..a0e824982fa 100644 --- a/litellm/tracing/types.py +++ b/litellm/tracing/types.py @@ -56,6 +56,7 @@ class TraceSummary(TypedDict): trace_ref: ReadOnly[NotRequired[str]] name: ReadOnly[str] service: ReadOnly[str] + agent_names: ReadOnly[NotRequired[tuple[str, ...]]] input_preview: ReadOnly[str] start_time: ReadOnly[str] # ISO 8601 duration_ms: ReadOnly[float] diff --git a/tests/test_litellm/tracing/test_decode.py b/tests/test_litellm/tracing/test_decode.py index 21a79dd6b87..dda075863ac 100644 --- a/tests/test_litellm/tracing/test_decode.py +++ b/tests/test_litellm/tracing/test_decode.py @@ -55,13 +55,94 @@ def _kv(key: str, value: str | int) -> KeyValue: return KeyValue(key=key, value=AnyValue(string_value=value)) -def _export(*spans: Span, service: str = "svc", scope: str = "test") -> bytes: +def _export(*spans: Span, service: str = "svc", scope: str = "test", agent_name: str = "") -> bytes: resource_spans = ResourceSpans(scope_spans=[ScopeSpans(spans=list(spans))]) resource_spans.resource.attributes.append(_kv("service.name", service)) + if agent_name: + resource_spans.resource.attributes.append(_kv("gen_ai.agent.name", agent_name)) resource_spans.scope_spans[0].scope.name = scope return ExportTraceServiceRequest(resource_spans=[resource_spans]).SerializeToString() +@pytest.mark.parametrize( + ("name", "attributes"), + [ + ("research_agent", {"openinference.span.kind": "AGENT", "metadata": '{"lc_agent_name":"research_agent"}'}), + ("research_agent", {"openinference.span.kind": "AGENT", "metadata": '{"ls_integration":"langgraph"}'}), + ("research_agent._execute_core", {"openinference.span.kind": "AGENT", "graph.node.id": "research_agent"}), + ("agent", {"openinference.span.kind": "AGENT", "gen_ai.agent.name": "research_agent"}), + ("openclaw.harness.run", {"openclaw.agent": "research_agent"}), + ( + "invoke_agent research_agent", + {"gen_ai.operation.name": "invoke_agent", "gen_ai.agent.name": "research_agent"}, + ), + ], + ids=["deepagents", "langgraph", "crewai", "hermes", "openclaw", "genai"], +) +def test_framework_agent_identity_is_independent_of_service(name: str, attributes: dict[str, str]): + span = _span(name, b"\x02" * 8, **attributes) + row = decode_otlp(_export(span, service="shared-deployment"), "application/x-protobuf")[0] + assert row["AgentName"] == "research_agent" + assert row["ServiceName"] == "shared-deployment" + assert row["SpanName"] == name + + +@pytest.mark.parametrize("name", ["ClaudeAgentSDK.query", "FunctionAgent.run"]) +def test_resource_agent_name_labels_instrumentors_without_an_agent_attribute(name: str): + span = _span(name, b"\x02" * 8, openinference__span__kind="AGENT") + row = decode_otlp(_export(span, agent_name="research_agent"), "application/x-protobuf")[0] + assert row["AgentName"] == "research_agent" + + +def test_span_agent_name_takes_precedence_over_resource_default(): + span = _span("invoke_agent child", b"\x02" * 8, gen_ai__agent__name="child") + row = decode_otlp(_export(span, agent_name="research_agent"), "application/x-protobuf")[0] + assert row["AgentName"] == "child" + + +@pytest.mark.parametrize( + ("scope", "span_name", "configured_name", "expected"), + [ + ("hermes-otel-plugin", "hermes-agent", "research_agent", "research_agent"), + ("hermes-otel-plugin", "child", "research_agent", "child"), + ("hermes-otel-plugin", "hermes-agent", "", "hermes-agent"), + ("other-plugin", "hermes-agent", "research_agent", "hermes-agent"), + ], +) +def test_hermes_resource_name_replaces_only_its_plugin_default( + scope: str, span_name: str, configured_name: str, expected: str +): + span = _span("agent", b"\x02" * 8, gen_ai__agent__name=span_name) + row = decode_otlp(_export(span, scope=scope, agent_name=configured_name), "application/x-protobuf")[0] + assert row["AgentName"] == expected + + +@pytest.mark.parametrize("agent_name", ["research_agent", ""]) +def test_openinference_middleware_is_not_a_separate_agent(agent_name: str): + span = _span( + "PatchToolCallsMiddleware.before_agent", b"\x02" * 8, b"\x01" * 8, + openinference__span__kind="AGENT", metadata=json.dumps({"lc_agent_name": agent_name}), + ) + row = decode_otlp(_export(span, scope="openinference.instrumentation.langchain"), "application/x-protobuf")[0] + assert (row["ObservationType"], row["AgentName"]) == ("framework", agent_name) + + +@pytest.mark.parametrize("scope", ["test", "openinference.instrumentation.langchain"]) +@pytest.mark.parametrize("kind", ["CHAIN", "AGENT"]) +@pytest.mark.parametrize("metadata", ["not json", "[]", '{"lc_agent_name":null}', "{}"]) +def test_unnamed_framework_does_not_invent_an_agent_from_service(metadata: str, scope: str, kind: str): + span = _span("workflow", b"\x02" * 8, openinference__span__kind=kind, metadata=metadata) + row = decode_otlp(_export(span, scope=scope), "application/x-protobuf")[0] + assert row["AgentName"] == "" + + +@pytest.mark.parametrize("name,expected", [("support", "support"), ("LangGraph", "")]) +def test_langgraph_distinguishes_configured_graph_name_from_default(name: str, expected: str): + span = _span(name, b"\x02" * 8, openinference__span__kind="CHAIN", metadata='{"ls_integration":"langgraph"}') + row = decode_otlp(_export(span, scope="openinference.instrumentation.langchain"), "application/x-protobuf")[0] + assert row["AgentName"] == expected + + def _span(name: str, span_id: bytes, parent: bytes = b"", **attributes: str | int) -> Span: return Span( trace_id=bytes.fromhex(TRACE_ID), diff --git a/tests/test_litellm/tracing/test_store.py b/tests/test_litellm/tracing/test_store.py index 3f43e42842c..0dd615a96ac 100644 --- a/tests/test_litellm/tracing/test_store.py +++ b/tests/test_litellm/tracing/test_store.py @@ -221,6 +221,23 @@ def test_agent_nodes_ignores_spans_of_unknown_agents(): assert agent_nodes(spans) == () +def test_trace_groups_normalized_names_and_preserves_span_labels(): + rows = [ + _row("root", "", "invoke_agent research_agent", "agent", "research_agent"), + _row("r1", "root", "researcher._execute_core", "agent", "researcher"), + _row("r2", "r1", "invoke_agent researcher", "agent", "researcher"), + _row("llm", "r2", "chat", "llm", "researcher"), + ] + result = trace_from_rows("t1", rows) + assert result is not None + assert result["summary"]["agent_names"] == ("research_agent", "researcher") + assert result["summary"]["name"] == "invoke_agent research_agent" + agents = {agent["name"]: agent for agent in result["agents"]} + assert agents["researcher"]["parent_agent"] == "research_agent" + assert agents["researcher"]["invocations"] == 2 + assert agents["researcher"]["llm_calls"] == 1 + + # ---------------------------------------------------------------- list helpers diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx index 5db34ef9788..d9cb4ee0399 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx @@ -271,10 +271,13 @@ describe("AgentTracesSection", () => { expect(rows[0]).toHaveTextContent("Should we store OTEL agent spans"); }); - it("labels the OTEL service as the agent and filters runs by it", async () => { + it("uses recorded agent names for the column and filter even when services are shared", async () => { vi.mocked(agentTraceListCall).mockResolvedValue({ ...(traceList as TracePage), - data: [...runs.slice(1), { ...runs[0], service: "billing-agent" }], + data: [ + ...runs.slice(1).map((run) => ({ ...run, service: "shared-app", agent_names: ["research-agent"] })), + { ...runs[0], service: "shared-app", agent_names: ["billing-agent", "review-agent"] }, + ], }); const user = userEvent.setup(); renderSection(); @@ -288,6 +291,10 @@ describe("AgentTracesSection", () => { const rows = screen.getAllByTestId("agent-trace-row"); expect(rows).toHaveLength(1); expect(rows[0]).toHaveTextContent("billing-agent"); + expect(rows[0]).not.toHaveTextContent("shared-app"); + + await chooseSelectOption(user, agentFilter, "review-agent"); + expect(screen.getAllByTestId("agent-trace-row")).toHaveLength(1); await chooseSelectOption(user, agentFilter, "All agents"); expect(screen.getAllByTestId("agent-trace-row")).toHaveLength(runs.length); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx index 06094864230..4ec8fe7e0ec 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx @@ -9,7 +9,7 @@ import { AgentTracesTable } from "./AgentTracesTable"; import { RunDrawer } from "./RunDrawer"; import { ALL_AGENTS, RunsToolbar, type RunStatusFilter } from "./RunsToolbar"; import type { TraceSummary } from "./traceTypes"; -import { previewText } from "./traceUtils"; +import { previewText, traceAgentNames } from "./traceUtils"; import { TimeRangeControls } from "./TimeRangeControls"; import { TracesTimeline, type TimeWindow } from "./TracesTimeline"; import { ActiveDot } from "./ActiveDot"; @@ -27,7 +27,7 @@ export function filterRuns( return runs.filter((run) => { const haystack = [run.trace_id, previewText(run.input_preview), run.name].map((s) => s.toLowerCase()); const matchesQuery = !q || haystack.some((text) => text.includes(q)); - const matchesAgent = agent === ALL_AGENTS || run.service === agent; + const matchesAgent = agent === ALL_AGENTS || traceAgentNames(run).includes(agent); const failed = run.error_count > 0; const matchesStatus = status === "all" || (status === "error" ? failed : !failed); return matchesQuery && matchesAgent && matchesStatus; @@ -121,7 +121,7 @@ export function AgentTracesSection({ if (setup.disabledDetail == null) void history.refetch(); }; - const agents = useMemo(() => Array.from(new Set(traces.traces.map((t) => t.service))).sort(), [traces.traces]); + const agents = useMemo(() => Array.from(new Set(traces.traces.flatMap(traceAgentNames))).sort(), [traces.traces]); // Relative ranges end "now" (the list query uses Date.now() too); round to the minute so the histogram is stable. const endMs = isCustomDate ? moment(endTime).valueOf() : moment().endOf("minute").valueOf(); const range = useMemo( diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx index 60f6816bd7a..6f0c8eac5f0 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx @@ -8,7 +8,7 @@ import { cn } from "@/lib/cva.config"; import { StatusMark } from "./StatusMark"; import type { TraceSummary } from "./traceTypes"; -import { fmtMs, previewText, traceDisplayName } from "./traceUtils"; +import { fmtMs, previewText, traceDisplayName, traceAgentNames } from "./traceUtils"; interface AgentTracesTableProps { traces: TraceSummary[]; @@ -83,8 +83,8 @@ export function AgentTracesTable({ > {formatActivityTimestamp(run.start_time)} - - {run.service} + + {traceAgentNames(run).join(", ") || "—"}
diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts index 9353f2024c8..f1394cec980 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts @@ -287,6 +287,19 @@ describe("payload helpers", () => { expect(messageText(image)).toBe(image); expect(messageText("[not json")).toBe("[not json"); }); + + it("reads GenAI message parts and native content arrays without crashing previews", () => { + const question = "What is an agent trace?"; + const parts = [{ type: "text", content: question }]; + const input = JSON.stringify([{ role: "user", parts }]); + expect(parseMessages(input)).toEqual([{ role: "user", parts, content: question }]); + expect(previewText(input)).toBe(question); + expect( + parseMessages(JSON.stringify({ role: "assistant", content: [{ type: "text", text: "An execution record" }] })), + ).toEqual([{ role: "assistant", content: "An execution record" }]); + expect(parseMessages('[{"role":"assistant","tool_calls":[]}]')).toBeNull(); + expect(parseMessages('[{"role":"user","content":42}]')).toBeNull(); + }); }); describe("treeGuides", () => { diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts index 00d24d4eaef..4379255c430 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts @@ -9,6 +9,9 @@ import type { Span, TraceMessage, TraceSummary } from "./traceTypes"; /* Formatting */ /* ------------------------------------------------------------------ */ +export const traceAgentNames = (trace: TraceSummary): readonly string[] => + trace.agent_names ?? (trace.service ? [trace.service] : []); + export const fmtMs = (ms: number): string => { if (ms >= 60_000) return `${(ms / 60_000).toFixed(1)}m`; if (ms >= 1000) return `${(ms / 1000).toFixed(2)}s`; @@ -267,14 +270,10 @@ export const parseJson = (value: string): unknown => { } }; -const isMessage = (value: unknown): value is TraceMessage => { - const isObject = typeof value === "object" && value !== null; - return isObject && "role" in value && typeof (value as TraceMessage).role === "string"; -}; - const blockText = (block: unknown): string | null => { if (typeof block !== "object" || block === null) return null; - const text: unknown = Reflect.get(block, "text"); + const text: unknown = + Reflect.get(block, "text") ?? (Reflect.get(block, "type") === "text" ? Reflect.get(block, "content") : undefined); return typeof text === "string" ? text : null; }; @@ -296,13 +295,19 @@ export function messageText(content: string): string { .join("\n\n"); } -const withText = (message: TraceMessage): TraceMessage => ({ ...message, content: messageText(message.content) }); +const parseMessage = (value: unknown): TraceMessage | null => { + if (typeof value !== "object" || value === null) return null; + const role: unknown = Reflect.get(value, "role"); + const content: unknown = Reflect.get(value, "content") ?? Reflect.get(value, "parts"); + if (typeof role !== "string" || (typeof content !== "string" && !Array.isArray(content))) return null; + return { ...value, role, content: messageText(typeof content === "string" ? content : JSON.stringify(content)) }; +}; /** An llm span's input (array of messages) or output (one message); null when it isn't one. */ export function parseMessages(value: string): TraceMessage[] | null { const parsed = parseJson(value); - if (Array.isArray(parsed)) return parsed.every(isMessage) ? parsed.map(withText) : null; - return isMessage(parsed) ? [withText(parsed)] : null; + const messages = (Array.isArray(parsed) ? parsed : [parsed]).map(parseMessage); + return messages.every((message) => message !== null) ? messages : null; } /** Pretty JSON when the payload is JSON, else the raw string. */ diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index fc1a9c67946..a376f8d7910 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -46904,6 +46904,8 @@ export interface components { agent_count: number; /** Agent Invocations */ agent_invocations: number; + /** Agent Names */ + agent_names?: string[]; /** Duration Ms */ duration_ms: number; /** Error Count */ From b21e44cbf97b1b1b5ce71f7edc31059c7836cf17 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 12:11:33 -0700 Subject: [PATCH 019/139] feat(jwt): auto_register_map_existing_key maps JWT to the user's existing virtual key (#42375) * test(e2e): jwt auto_register map-existing-key repro Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(jwt): auto_register_map_existing_key maps JWT to the user's existing virtual key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(jwt): exclude blocked keys from auto_register_map_existing_key reuse Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(jwt): route existing-key lookup through VerificationTokenRepository Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(e2e): stop requiring LITELLM_SALT_KEY for the owned JWT gateway Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(e2e): gate the owned JWT gateway tests behind E2E_OWNED_GATEWAY Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(jwt): only reuse keys that can call LLM routes in auto_register_map_existing_key Skip Admin UI session keys and keys whose allowed_routes restrict them to anything other than llm_api_routes (management, read_only, password-reset sessions). Mapping a JWT to one of those left the user with 401s or 403s on every LLM call, since the mapping persists. * fix(jwt): scope auto_register_map_existing_key reuse to the JWT-resolved team Only reuse a key whose team_id matches the team auth_builder resolved for the JWT (no team matches no team), so a personal key can no longer bypass the resolved team's model and budget limits. With the flag on, the first JWT request now falls through to the same virtual-key checks later mapped requests get, instead of returning early, so a reused key's own limits apply from request one rather than 200 then 403. Flag off keeps the early return unchanged. * fix(jwt): keep the early return when no master key is set Without a master key the generic virtual-key path returns a bare INTERNAL_USER object, so falling through on the first auto-registered request dropped the key's team, models and budgets. Only fall through when a master key is configured. Tests now assert the reused key per team rather than the query shape, and cover the flag-off early return and the no-master-key case. * test(jwt): assert on race-loser's returned key, not only mocks (TQ002) Co-Authored-By: Claude Opus 5.5 * fix(jwt): close the auto_register_map_existing_key race, shared-claim and expiry holes A key auto_register just minted is never adopted by a concurrent request, so the race loser's cleanup can no longer delete a key another request mapped and cascade its mapping away (503, user left with no key) Reuse only happens when the claim value is the JWT-resolved user_id. A shared claim such as azp or client_id falls back to minting, so one user can no longer land on another user's personal key and budget Only keys that never expire are reused, so an expiring key can no longer pin the claim to a permanent 401 Integration tests on a real proxy and Postgres cover all three. The race test holds the first mapping insert in a Postgres relay, so the interleaving is forced rather than timed. The where-clause shape unit tests are replaced by these, since only a real database proves the filter * test(e2e): create the reused key in the team the JWT resolves to The flag only reuses a key in the JWT-resolved team, and this identity's groups claim resolves to its team, so a teamless key was never eligible and the test could not pass * test(integration): match the held statement across TCP reads The relay looked for the trigger inside one read, so an insert split across two reads was never held and the race test would fail waiting for it. It now matches one exact trigger over a window that keeps the end of the previous read * fix(jwt): gate key reuse on the claim field, not on the claim value Requiring the claim value to equal the resolved user_id skipped reuse for users matched through the sso_user_id or case-insensitive email fallback, whose stored user_id differs from the JWT sub. That is the lookup LIT-5378 asks for. Reuse is now allowed when the virtual key claim is the user_id or user_email JWT field, globally or for the token's issuer, which still keeps shared claims such as azp or client_id on the mint path * fix(jwt): let an issuer's own user field replace the global one when gating key reuse An issuer that identifies users by uid no longer treats the global sub field as a user identity claim, so a shared sub under that issuer mints instead of reusing a personal key * test(jwt): make the flag-off test fail when the flag no longer gates key reuse The flag-off test used a config where sub was not a user identity claim, so deleting the flag check still passed. Configure user_id_jwt_field=sub so only the flag keeps the lookup off, and drop test docstrings * chore(lint): drop mutable-ok suppressions that LIT013 flags as no-ops --------- Co-authored-by: yuneng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mrinal Co-authored-by: Mrinal Chanshetty Co-authored-by: Claude Opus 5.5 --- .github/workflows/test-e2e-changed.yml | 1 + litellm/proxy/_types.py | 20 + litellm/proxy/auth/user_api_key_auth.py | 238 ++++++---- .../verification_token_repository.py | 32 ++ tests/e2e/conftest.py | 7 + tests/e2e/coverage_registry/other.yaml | 3 + tests/e2e/e2e_config.py | 12 + tests/e2e/mcp/oauth_gateway.py | 12 +- tests/e2e/models.py | 35 ++ tests/e2e/other/other_client.py | 46 ++ tests/e2e/other/owned_jwt_gateway.py | 108 +++++ tests/e2e/other/test_jwt_auto_register_e2e.py | 182 +++++++ tests/e2e/pytest.ini | 1 + tests/integration/_support/database_relay.py | 80 +++- ...test_jwt_auto_register_map_existing_key.py | 219 +++++++++ .../test_user_api_key_auth_request_flow.py | 448 ++++++++++++++++++ 16 files changed, 1333 insertions(+), 111 deletions(-) create mode 100644 tests/e2e/other/owned_jwt_gateway.py create mode 100644 tests/e2e/other/test_jwt_auto_register_e2e.py create mode 100644 tests/integration/authorization/test_jwt_auto_register_map_existing_key.py diff --git a/.github/workflows/test-e2e-changed.yml b/.github/workflows/test-e2e-changed.yml index 8e03a902383..228e23f60d7 100644 --- a/.github/workflows/test-e2e-changed.yml +++ b/.github/workflows/test-e2e-changed.yml @@ -176,6 +176,7 @@ jobs: TESTS: ${{ needs.detect.outputs.tests }} E2E_FIXTURE_MODE: live E2E_PROVIDER_EDGE_HOST_REACHABLE: '1' + E2E_OWNED_GATEWAY: '1' COLUMNS: '400' run: | umask 077 diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index ef4c545507b..cabb04cc52b 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -5464,6 +5464,17 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): "'auto_register': auto-create a virtual key and mapping on first encounter." ), ) + auto_register_map_existing_key: bool = Field( + default=False, + description=( + "Only used with unregistered_jwt_client_behavior='auto_register'. When True and the virtual key claim " + "field is the user_id_jwt_field or user_email_jwt_field, the JWT claim is mapped to a virtual key the " + "JWT-resolved user already owns instead of minting a new one. If the user owns several, the most recently created key in the " + "JWT-resolved team (or with no team when the JWT resolves none) is chosen among keys that never " + "expire, are not blocked, are not Admin UI session keys, were not minted by auto_register, and " + "have no allowed_routes or include llm_api_routes. Otherwise a new key is minted as usual." + ), + ) routing_overrides: list[JWTRoutingOverride] | None = Field( default=None, description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.", @@ -5564,6 +5575,15 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): return issuer_config.virtual_key_claim_field return self.virtual_key_claim_field + def is_user_identity_claim(self, claim_field: str, issuer: str | None) -> bool: + issuer_config: Final = self.get_issuer_config(issuer) + if issuer_config is None: + return claim_field in (self.user_id_jwt_field, self.user_email_jwt_field) + return claim_field in ( + issuer_config.user_id_jwt_field or self.user_id_jwt_field, + issuer_config.user_email_jwt_field or self.user_email_jwt_field, + ) + def get_unregistered_jwt_client_behavior(self, issuer: str | None) -> UnregisteredJWTClientBehavior: issuer_config: Final = self.get_issuer_config(issuer) if issuer_config is not None and issuer_config.unregistered_jwt_client_behavior is not None: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 3400dccf2a7..e82f3eed7cc 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -140,6 +140,7 @@ from litellm.proxy.utils import ( normalize_route_for_root_path, ) from litellm.repositories.table_repositories import TeamMembershipRepository +from litellm.repositories.verification_token_repository import VerificationTokenRepository from litellm.router_utils.common_utils import resolve_model_group_alias from litellm.secret_managers.main import get_secret_bool from litellm.types.services import ServiceTypes @@ -939,6 +940,24 @@ class _PendingAutoRegister(NamedTuple): jwt_issuer: str | None = None +def _claim_identifies_user(jwt_handler: JWTHandler, claim_field: str, jwt_issuer: str | None) -> bool: + if not jwt_handler.litellm_jwtauth.auto_register_map_existing_key: + return False + if jwt_handler.litellm_jwtauth.is_user_identity_claim(claim_field, jwt_issuer): + return True + verbose_proxy_logger.warning( + "JWT Key Mapping (auto_register_map_existing_key): claim '%s' is not the user_id or user_email JWT field " + "and may be shared by several users, so a new key is minted instead of reusing one the user owns.", + claim_field, + ) + return False + + +async def _reusable_key_hash_for_user(prisma_client: PrismaClient, user_id: str, team_id: str | None) -> str | None: + key: Final = await VerificationTokenRepository(prisma_client).find_newest_reusable_llm_api_key(user_id, team_id) + return None if key is None else key.token + + async def _auto_register_jwt_mapping( virtual_key_claim_field: str, claim_value: str, @@ -957,8 +976,10 @@ async def _auto_register_jwt_mapping( ) -> UserAPIKeyAuth | None: """ Auto-register: create a new virtual key + mapping for an unrecognised JWT - claim value. ``team_id`` and ``user_id`` must come from a successful - ``JWTAuthManager.auth_builder`` run — they encode the JWT identity AFTER + claim value, or point the mapping at a key the resolved user already owns + when ``auto_register_map_existing_key`` is set. ``team_id`` and ``user_id`` + must come from a successful ``JWTAuthManager.auth_builder`` run — they + encode the JWT identity AFTER RBAC/scope/custom_validate/email-domain policy has been enforced. The key is stamped with those values so the cached future-request path inherits the same team/user/org limits the auth_builder path would have applied. @@ -974,29 +995,38 @@ async def _auto_register_jwt_mapping( generate_key_helper_fn, ) - # ``table_name="key"`` is required: without it, generate_key_helper_fn - # falls into the user-upsert branch (`table_name is None or "user"`) and - # attempts to insert into LiteLLM_UserTable with user_id=None, which fails - # the NOT NULL @id constraint. Every successful key-creation caller (e.g. - # /key/generate) passes table_name="key" explicitly. - key_data: Final = await generate_key_helper_fn( - llm_router=None, - request_type="key", - table_name="key", - team_id=team_id, - user_id=user_id, - organization_id=org_id, - agent_id=agent_id, - metadata={ - "auto_registered": True, - "jwt_claim_field": virtual_key_claim_field, - "jwt_claim_value": claim_value, - }, + existing_token_hash: Final = ( + await _reusable_key_hash_for_user(prisma_client, user_id, team_id) + if user_id is not None and _claim_identifies_user(jwt_handler, virtual_key_claim_field, jwt_issuer) + else None ) - # generate_key_helper_fn returns the plaintext key in "token"; the persisted - # row in LiteLLM_VerificationToken uses its hash, so hash here to get the FK - # value referenced by LiteLLM_JWTKeyMapping.token. - token_hash = hash_token(key_data["token"]) + minted: Final = existing_token_hash is None + if existing_token_hash is not None: + token_hash = existing_token_hash + else: + # ``table_name="key"`` is required: without it, generate_key_helper_fn + # falls into the user-upsert branch (`table_name is None or "user"`) and + # attempts to insert into LiteLLM_UserTable with user_id=None, which fails + # the NOT NULL @id constraint. Every successful key-creation caller (e.g. + # /key/generate) passes table_name="key" explicitly. + key_data: Final = await generate_key_helper_fn( + llm_router=None, + request_type="key", + table_name="key", + team_id=team_id, + user_id=user_id, + organization_id=org_id, + agent_id=agent_id, + metadata={ + "auto_registered": True, + "jwt_claim_field": virtual_key_claim_field, + "jwt_claim_value": claim_value, + }, + ) + # generate_key_helper_fn returns the plaintext key in "token"; the persisted + # row in LiteLLM_VerificationToken uses its hash, so hash here to get the FK + # value referenced by LiteLLM_JWTKeyMapping.token. + token_hash = hash_token(key_data["token"]) try: await prisma_client.db.litellm_jwtkeymapping.create( @@ -1023,15 +1053,16 @@ async def _auto_register_jwt_mapping( virtual_key_claim_field, claim_value, ) - try: - await prisma_client.db.litellm_verificationtoken.delete(where={"token": token_hash}) - except Exception as delete_err: - # Don't fail the request if cleanup fails — the orphan is - # unmapped and inert. Log so an operator can prune it later. - verbose_proxy_logger.warning( - "JWT Key Mapping (auto_register): failed to delete orphaned key after race: %s", - delete_err, - ) + if minted: + try: + await prisma_client.db.litellm_verificationtoken.delete(where={"token": token_hash}) + except Exception as delete_err: + # Don't fail the request if cleanup fails — the orphan is + # unmapped and inert. Log so an operator can prune it later. + verbose_proxy_logger.warning( + "JWT Key Mapping (auto_register): failed to delete orphaned key after race: %s", + delete_err, + ) token_hash = await get_jwt_key_mapping_object( jwt_claim_name=virtual_key_claim_field, jwt_claim_value=claim_value, @@ -1061,7 +1092,8 @@ async def _auto_register_jwt_mapping( ) verbose_proxy_logger.info( - "JWT Key Mapping (auto_register): created new virtual key for %s='%s'.", + "JWT Key Mapping (auto_register): %s virtual key for %s='%s'.", + "created new" if minted else "mapped existing", virtual_key_claim_field, claim_value, ) @@ -1075,7 +1107,8 @@ async def _auto_register_jwt_mapping( ).resolve(hashed_token=token_hash) ) if auto_registered_key is not None: - auto_registered_key.org_id = org_id + if minted: + auto_registered_key.org_id = org_id auto_registered_key.end_user_id = end_user_id auto_registered_key.api_key = auto_registered_key.token return auto_registered_key @@ -1771,8 +1804,8 @@ async def _user_api_key_auth_builder( # mapping + virtual key from the *validated* identity, then # replace valid_token with the new key so downstream checks # use the key-scoped path. - if pending_auto_register is not None and prisma_client is not None: - auto_registered: Final = await _auto_register_jwt_mapping( + auto_registered: Final = ( + await _auto_register_jwt_mapping( virtual_key_claim_field=pending_auto_register.claim_field, claim_value=pending_auto_register.claim_value, jwt_handler=jwt_handler, @@ -1788,72 +1821,81 @@ async def _user_api_key_auth_builder( end_user_id=end_user_id, agent_id=agent_id, ) - if auto_registered is not None: - auto_registered.jwt_claims = jwt_claims - auto_registered.user_email = user_email - # The auto-registered token is built from the new key's - # columns, which carry no user budget. Carry over the - # already-loaded user row rather than re-reading it, or - # the budget check below has nothing to enforce. - auto_registered.user_model_max_budget = ( - user_object.model_max_budget if user_object is not None else None - ) - valid_token = auto_registered - api_key = valid_token.token or "" - - # Check if model has zero cost - if so, skip all budget checks - model = _get_model_from_request_context( - request_data=request_data, - route=route, - request=request, - llm_router=llm_router, - team_id=valid_token.team_id, + if pending_auto_register is not None and prisma_client is not None + else None ) - skip_budget_checks = False - if model is not None and llm_router is not None: - from litellm.proxy.auth.auth_checks import _is_model_cost_zero - - skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) - if skip_budget_checks: - verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) - - # Fetch project object for JWT path if project_id is set - _jwt_project_obj = None - if valid_token.project_id is not None: - _jwt_project_obj = await get_project_object( - project_id=valid_token.project_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, + if auto_registered is not None: + auto_registered.jwt_claims = jwt_claims + auto_registered.user_email = user_email + # The auto-registered token is built from the new key's + # columns, which carry no user budget. Carry over the + # already-loaded user row rather than re-reading it, or + # the budget check below has nothing to enforce. + auto_registered.user_model_max_budget = ( + user_object.model_max_budget if user_object is not None else None ) - if _jwt_project_obj is not None: - valid_token.project_metadata = _jwt_project_obj.metadata - valid_token.project_alias = _jwt_project_obj.project_alias + valid_token = auto_registered + api_key = valid_token.token or "" - # JWT auth returns here rather than falling through to the - # virtual-key checks below, so the user's per-model budget - # has to be enforced on this path too. Without it the - # post-call increment still charges the counter and nothing - # ever reads it, which is worse than not tracking at all. - # Guarded by the same flag the virtual-key path uses, or a - # zero-cost model would be refused here and allowed there, - # while the log above claims all budget checks were skipped. - if not skip_budget_checks: - await _check_user_model_budget( - valid_token=cast(UserAPIKeyAuth, valid_token), - model_max_budget_limiter=model_max_budget_limiter, - models=_get_model_names_for_budget_checks( - model=_get_model_from_request_context( - request_data=request_data, - route=route, - request=request, - llm_router=llm_router, - team_id=valid_token.team_id, - ) - ), + falls_through_to_key_checks: Final = ( + auto_registered is not None + and jwt_handler.litellm_jwtauth.auto_register_map_existing_key + and master_key is not None + ) + if not falls_through_to_key_checks: + # Check if model has zero cost - if so, skip all budget checks + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + llm_router=llm_router, + team_id=valid_token.team_id, ) + skip_budget_checks = False + if model is not None and llm_router is not None: + from litellm.proxy.auth.auth_checks import _is_model_cost_zero - return cast(UserAPIKeyAuth, valid_token) + skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) + if skip_budget_checks: + verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) + + # Fetch project object for JWT path if project_id is set + _jwt_project_obj = None + if valid_token.project_id is not None: + _jwt_project_obj = await get_project_object( + project_id=valid_token.project_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if _jwt_project_obj is not None: + valid_token.project_metadata = _jwt_project_obj.metadata + valid_token.project_alias = _jwt_project_obj.project_alias + + # JWT auth returns here rather than falling through to the + # virtual-key checks below, so the user's per-model budget + # has to be enforced on this path too. Without it the + # post-call increment still charges the counter and nothing + # ever reads it, which is worse than not tracking at all. + # Guarded by the same flag the virtual-key path uses, or a + # zero-cost model would be refused here and allowed there, + # while the log above claims all budget checks were skipped. + if not skip_budget_checks: + await _check_user_model_budget( + valid_token=cast(UserAPIKeyAuth, valid_token), + model_max_budget_limiter=model_max_budget_limiter, + models=_get_model_names_for_budget_checks( + model=_get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + llm_router=llm_router, + team_id=valid_token.team_id, + ) + ), + ) + + return cast(UserAPIKeyAuth, valid_token) #### ELSE #### ## CHECK PASS-THROUGH ENDPOINTS ## diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py index d02c2114136..b20fb47306e 100644 --- a/litellm/repositories/verification_token_repository.py +++ b/litellm/repositories/verification_token_repository.py @@ -8,6 +8,7 @@ from datetime import datetime from types import TracebackType from typing import TYPE_CHECKING, Final, Protocol +from litellm.constants import UI_SESSION_TOKEN_TEAM_ID from litellm.models.verification_token import ( LiteLLM_VerificationToken, ) @@ -123,6 +124,37 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"user_id": user_id}) return self._to_model_list(records) + async def find_newest_reusable_llm_api_key( + self, user_id: str, team_id: str | None + ) -> LiteLLM_VerificationToken | None: + records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many( + where={ + "user_id": user_id, + "team_id": team_id, + "expires": None, + "AND": [ + {"OR": [{"blocked": False}, {"blocked": None}]}, + { + "OR": [ + {"team_id": None}, + {"team_id": {"not": UI_SESSION_TOKEN_TEAM_ID}}, + ] + }, + { + "OR": [ + {"allowed_routes": {"is_empty": True}}, + {"allowed_routes": {"has": "llm_api_routes"}}, + ] + }, + ], + }, + order={"created_at": "desc"}, + ) + return next( + (key for key in self._to_model_list(records) if key.metadata.get("auto_registered") is not True), + None, + ) + async def find_by_team_id(self, team_id: str) -> list[LiteLLM_VerificationToken]: """Find all tokens belonging to a team.""" records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"team_id": team_id}) diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 62153e38a83..37f0bf00da6 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -32,6 +32,7 @@ from e2e_config import ( MCP_OAUTH_LIVE_OPT_IN_ENV, OTEL_TLS_OPT_IN_ENV, OTEL_V2_OPT_IN_ENV, + OWNED_GATEWAY_OPT_IN_ENV, PROMPT_CACHING_OPT_IN_ENV, PROVIDER_EDGE_HOST_OPT_IN_ENV, PROXY_BASE_URL, @@ -70,6 +71,7 @@ OPT_IN_MARKERS: Final = MappingProxyType( "cli_determinism": CLI_DETERMINISM_OPT_IN_ENV, "mcp_oauth_live": MCP_OAUTH_LIVE_OPT_IN_ENV, "provider_edge_host": PROVIDER_EDGE_HOST_OPT_IN_ENV, + "owned_gateway": OWNED_GATEWAY_OPT_IN_ENV, "otel_v2": OTEL_V2_OPT_IN_ENV, "otel_tls": OTEL_TLS_OPT_IN_ENV, "secret_manager": SECRET_MANAGER_OPT_IN_ENV, @@ -172,6 +174,11 @@ def pytest_configure(config: pytest.Config) -> None: "provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the " "gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set", ) + config.addinivalue_line( + "markers", + "owned_gateway: boots its own proxy from source against the stack's Postgres, so it needs DATABASE_URL " + "on the pytest host; deselected unless E2E_OWNED_GATEWAY is set", + ) config.addinivalue_line( "markers", "otel_v2: needs a proxy running with LITELLM_OTEL_V2=true; deselected unless E2E_OTEL_V2 is set", diff --git a/tests/e2e/coverage_registry/other.yaml b/tests/e2e/coverage_registry/other.yaml index 0b9249d7420..39747607531 100644 --- a/tests/e2e/coverage_registry/other.yaml +++ b/tests/e2e/coverage_registry/other.yaml @@ -60,6 +60,9 @@ - {id: other.auth.jwt.wrong_issuer_denied, module: other, tier: P0, area: auth, assertions: [wrong_issuer_denied], source: "auth/handle_jwt.py", rationale: "A signed token with the correct audience and an unexpected issuer is rejected"} - {id: other.auth.jwt.wrong_audience_denied, module: other, tier: P0, area: auth, assertions: [wrong_audience_denied], source: "auth/handle_jwt.py", rationale: "A signed token from the trusted issuer intended for another app is rejected"} +- {id: other.auth.jwt.auto_register_maps_existing_key, module: other, tier: P0, area: auth, assertions: [maps_existing_key], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "With auto_register_map_existing_key: true, the first JWT call of a user who already owns a key writes the sub-claim mapping to that existing key hash and mints nothing; the spend row lands on the pre-existing key (LIT-5378)", fail_before_fix: proven} +- {id: other.auth.jwt.auto_register_mints_when_keyless, module: other, tier: P0, area: auth, assertions: [mints_when_keyless], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "With auto_register_map_existing_key: true, a user with no keys still gets exactly one minted key and a sub-claim mapping on their first JWT call (LIT-5378)", fail_before_fix: proven} +- {id: other.auth.jwt.auto_register_default_mints, module: other, tier: P0, area: auth, assertions: [default_mints], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "Without auto_register_map_existing_key, auto_register keeps the current behavior: it mints a second key for a user who already has one and bills the minted key (LIT-5378)", fail_before_fix: proven} - {id: other.auth.session_token.valid_allows, module: other, tier: P0, area: auth, assertions: [valid_allows], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An unexpired LiteLLM-minted session token authenticates with the role it carries"} - {id: other.auth.session_token.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "auth/user_api_key_auth.py expiry check", rationale: "An expired session token is rejected with the expired-key error"} - {id: other.auth.session_token.encrypted_value_denied, module: other, tier: P0, area: auth, assertions: [encrypted_value_denied], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An encrypted value read back from a management route is not accepted as a bearer token"} diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 3fa9f534ffd..e88bfad8388 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -7,6 +7,7 @@ environment so the same tests run against localhost or a deployed proxy. from __future__ import annotations import os +import socket from dataclasses import dataclass import time import uuid @@ -16,6 +17,7 @@ from typing import Final from dotenv import load_dotenv from fixture_mode import deterministic_marker, parse_fixture_mode, registration_owner from provider_edge import provider_edge_api_base +from pydantic import TypeAdapter # Local runs keep provider / DataDog keys in tests/e2e/.env (see CONTRIBUTING.md). # Compose injects them into the proxy container, but pytest on the host does not @@ -206,6 +208,7 @@ REDIS_CHAOS_OPT_IN_ENV = "E2E_REDIS_CHAOS" CLI_DETERMINISM_OPT_IN_ENV = "E2E_CLI_DETERMINISM" MCP_OAUTH_LIVE_OPT_IN_ENV: Final = "E2E_MCP_OAUTH_LIVE" PROVIDER_EDGE_HOST_OPT_IN_ENV: Final = "E2E_PROVIDER_EDGE_HOST_REACHABLE" +OWNED_GATEWAY_OPT_IN_ENV: Final = "E2E_OWNED_GATEWAY" OTEL_V2_OPT_IN_ENV: Final = "E2E_OTEL_V2" OTEL_TLS_OPT_IN_ENV: Final = "E2E_OTEL_EXPORTER_ENDPOINT" SECRET_MANAGER_OPT_IN_ENV: Final = "E2E_SECRET_MANAGER" @@ -296,6 +299,15 @@ def unique_marker() -> str: return uuid.uuid4().hex[:12] +INHERITED_ENV_PREFIXES: Final = ("REDIS_", "MICROSOFT_", "GOOGLE_", "GENERIC_", "PROXY_") + + +def available_port() -> int: + with socket.socket() as listener: + listener.bind(("127.0.0.1", 0)) + return TypeAdapter(tuple[str, int]).validate_python(listener.getsockname())[1] + + def settle_propagation(written_at: float) -> None: """Block until PROPAGATION_TIMEOUT has elapsed since `written_at`, a `time.monotonic()` stamp taken the moment a control-plane write returned. diff --git a/tests/e2e/mcp/oauth_gateway.py b/tests/e2e/mcp/oauth_gateway.py index b328c81687b..029b0135900 100644 --- a/tests/e2e/mcp/oauth_gateway.py +++ b/tests/e2e/mcp/oauth_gateway.py @@ -8,7 +8,6 @@ The optional live edge measures headers without recording credentials or bodies. from __future__ import annotations import os -import socket import subprocess import sys import threading @@ -20,13 +19,12 @@ from pathlib import Path from typing import Final import psycopg +from e2e_config import INHERITED_ENV_PREFIXES, available_port from e2e_http import NoBody from idp import Keycloak, stop_process_group from proxy_client import ProxyClient, build_proxy_client from psycopg.rows import class_row -from pydantic import BaseModel, SecretStr, TypeAdapter, ValidationError - -INHERITED_ENV_PREFIXES: Final = ("REDIS_", "MICROSOFT_", "GOOGLE_", "GENERIC_", "PROXY_") +from pydantic import BaseModel, SecretStr, ValidationError class StoredOAuth(BaseModel): @@ -101,12 +99,6 @@ class OAuthObservation: assert all(not item[2] for item in snapshot), "gateway bearer leaked to the upstream" -def available_port() -> int: - with socket.socket() as listener: - listener.bind(("127.0.0.1", 0)) - return TypeAdapter(tuple[str, int]).validate_python(listener.getsockname())[1] - - @dataclass(slots=True) class OAuthGateway: base_url: str diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 83d8a9884a3..e027c410e44 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -1537,6 +1537,7 @@ class UserNewBody(BaseModel): class UserNewResponse(BaseModel): user_id: str + key: str | None = None class UserUpdateBody(BaseModel): @@ -1580,6 +1581,40 @@ class UserListResponse(BaseModel): total: int +class UserKeyRow(BaseModel): + token: str + key_alias: str | None = None + + +class UserInfoWithKeysResponse(BaseModel): + user_id: str | None = None + keys: list[UserKeyRow] = [] + + +class JwtKeyMappingRow(BaseModel): + id: str + jwt_claim_name: str + jwt_claim_value: str + created_by: str | None = None + + +class JwtKeyMappingListParams(BaseModel): + size: int = 100 + + +class JwtKeyMappingListResponse(BaseModel): + mappings: list[JwtKeyMappingRow] + total_count: int + + +class JwtKeyMappingDeleteBody(BaseModel): + id: str + + +class JwtKeyMappingDeleteResponse(BaseModel): + status: str + + class OrgNewBody(BaseModel): organization_alias: str models: list[str] = [] diff --git a/tests/e2e/other/other_client.py b/tests/e2e/other/other_client.py index 93c198586f6..d7bad4f1ed1 100644 --- a/tests/e2e/other/other_client.py +++ b/tests/e2e/other/other_client.py @@ -21,12 +21,20 @@ from idp import Keycloak, keycloak_from_env from models import ( ChatBody, ChatResponse, + JwtKeyMappingDeleteBody, + JwtKeyMappingDeleteResponse, + JwtKeyMappingListParams, + JwtKeyMappingListResponse, ModelsListParams, ModelsListResponse, ReadinessDetailsResponse, ReadinessResponse, + UserInfoParams, + UserInfoWithKeysResponse, UserListParams, UserListResponse, + UserNewBody, + UserNewResponse, ) from proxy_client import ProxyClient from pydantic import Field @@ -79,6 +87,44 @@ class OtherClient: response_type=ReadinessDetailsResponse, ) + def user_new(self, body: UserNewBody) -> Result[UserNewResponse]: + """POST /user/new under the master key: seed the litellm user a JWT + `sub` claim resolves to, before that token ever reaches the proxy.""" + return self.proxy.transport.post( + "/user/new", + headers=self.proxy.transport.master, + json=body, + response_type=UserNewResponse, + ) + + def user_info(self, user_id: str) -> Result[UserInfoWithKeysResponse]: + """GET /user/info under the master key. Only the user's key rows are + modelled: `token` is the stored key hash, never the plaintext key.""" + return self.proxy.transport.get( + "/user/info", + headers=self.proxy.transport.master, + params=UserInfoParams(user_id=user_id), + response_type=UserInfoWithKeysResponse, + ) + + def jwt_mapping_list(self) -> Result[JwtKeyMappingListResponse]: + """GET /jwt/key/mapping/list under the master key.""" + return self.proxy.transport.get( + "/jwt/key/mapping/list", + headers=self.proxy.transport.master, + params=JwtKeyMappingListParams(size=100), + response_type=JwtKeyMappingListResponse, + ) + + def jwt_mapping_delete(self, mapping_id: str) -> Result[JwtKeyMappingDeleteResponse]: + """POST /jwt/key/mapping/delete under the master key.""" + return self.proxy.transport.post( + "/jwt/key/mapping/delete", + headers=self.proxy.transport.master, + json=JwtKeyMappingDeleteBody(id=mapping_id), + response_type=JwtKeyMappingDeleteResponse, + ) + def chat_as_team(self, token: str, team: str, body: ChatBody) -> Result[ChatResponse]: """POST /chat/completions under `token` with `x-litellm-team-id: team`.""" return self.proxy.transport.post( diff --git a/tests/e2e/other/owned_jwt_gateway.py b/tests/e2e/other/owned_jwt_gateway.py new file mode 100644 index 00000000000..1af348cac60 --- /dev/null +++ b/tests/e2e/other/owned_jwt_gateway.py @@ -0,0 +1,108 @@ +"""An owned, source-built proxy whose `litellm_jwtauth` block a test controls. + +The shared proxy on :4000 runs the CONTRIBUTING.md JWT block, so a test that +needs a different `litellm_jwtauth` config boots its own gateway on a free port +against the same database and the same Keycloak realm. The caller supplies the +`litellm_jwtauth` mapping verbatim, which is exactly what makes a config an +unfixed proxy rejects observable as a boot failure in this gateway's own log. +""" + +from __future__ import annotations + +import os +import subprocess +import sys +import time +from collections.abc import Mapping +from contextlib import ExitStack +from dataclasses import dataclass, field +from pathlib import Path +from typing import Final + +from e2e_config import INHERITED_ENV_PREFIXES, available_port +from e2e_http import NoBody +from idp import Keycloak, stop_process_group +from proxy_client import ProxyClient, build_proxy_client + +MODEL_NAME: Final = "gemini-3.8-flash" + + +@dataclass(slots=True) +class OwnedJwtGateway: + base_url: str + proxy: ProxyClient + _environment: Mapping[str, str] = field(repr=False) + _command: tuple[str, ...] = field(repr=False) + _log_path: Path + _child: subprocess.Popen[bytes] | None = field(default=None, init=False, repr=False) + + def start(self) -> None: + with self._log_path.open("ab") as log: + self._child = subprocess.Popen( + self._command, + env=self._environment, + stdout=log, + stderr=log, + start_new_session=True, + ) + deadline: Final = time.monotonic() + 120 + while time.monotonic() < deadline: + assert self._child.poll() is None, "owned JWT gateway exited; inspect its private log" + result = self.proxy.transport.probe("/health/liveliness", params=NoBody()) + if result.status_code == 200: + return + time.sleep(0.5) + raise AssertionError("owned JWT gateway did not become ready") + + def stop(self) -> None: + if self._child is not None: + stop_process_group(self._child) + assert self._child.poll() is not None, "old gateway process is still alive" + + +def owned_jwt_gateway( + idp: Keycloak, directory: Path, cleanup: ExitStack, *, litellm_jwtauth: str, name: str +) -> OwnedJwtGateway: + for env_name in ("DATABASE_URL", "LITELLM_LICENSE", "LITELLM_MASTER_KEY"): + assert os.environ.get(env_name), f"{env_name} is required for the owned JWT gateway" + port: Final = available_port() + base_url: Final = f"http://127.0.0.1:{port}" + config: Final = directory / f"{name}.yaml" + config.write_text( + "model_list:\n" + f" - model_name: {MODEL_NAME}\n" + " litellm_params:\n" + f" model: gemini/{MODEL_NAME}\n" + " api_key: os.environ/GEMINI_API_KEY\n" + "general_settings:\n" + " master_key: os.environ/LITELLM_MASTER_KEY\n" + " database_url: os.environ/DATABASE_URL\n" + " proxy_batch_write_at: 5\n" + " enable_jwt_auth: true\n" + " litellm_jwtauth:\n" + "".join(f" {line}\n" for line in litellm_jwtauth.strip().splitlines()) + ) + environment: Final = { + **{key: value for key, value in os.environ.items() if not key.startswith(INHERITED_ENV_PREFIXES)}, + "JWT_PUBLIC_KEY_URL": idp.jwks_url, + "JWT_ISSUER": idp.issuer, + "JWT_AUDIENCE": "litellm-e2e", + "LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY": "true", + "DISABLE_SCHEMA_UPDATE": "true", + "STORE_MODEL_IN_DB": "True", + "PYTHONPATH": str(Path(__file__).resolve().parents[3]), + } + gateway: Final = OwnedJwtGateway( + base_url=base_url, + proxy=build_proxy_client( + base_url=base_url, + control_plane_base_url=base_url, + replica_urls=(base_url,), + master_key=os.environ["LITELLM_MASTER_KEY"], + ), + _environment=environment, + _command=(sys.executable, "-m", "litellm.proxy.proxy_cli", "--config", str(config), "--port", str(port)), + _log_path=directory / f"{name}.log", + ) + cleanup.callback(gateway.stop) + gateway.start() + return gateway diff --git a/tests/e2e/other/test_jwt_auto_register_e2e.py b/tests/e2e/other/test_jwt_auto_register_e2e.py new file mode 100644 index 00000000000..8f7c2c6a693 --- /dev/null +++ b/tests/e2e/other/test_jwt_auto_register_e2e.py @@ -0,0 +1,182 @@ +"""auto_register with auto_register_map_existing_key binds the JWT claim to the user's existing key. + +`unregistered_jwt_client_behavior: auto_register` on `virtual_key_claim_field: sub` mints a fresh +virtual key on the user's first JWT call. With `auto_register_map_existing_key: true` the proxy must +instead point the new JWT mapping at a key the resolved user already owns, and mint only when the +user has none. Each behavior gets its own gateway because the flag lives in `litellm_jwtauth`, so +this file boots two owned proxies against the shared database and Keycloak realm. +""" + +from __future__ import annotations + +import hashlib +from collections.abc import Iterator +from contextlib import ExitStack +from typing import Final + +import pytest +from e2e_config import unique_marker +from e2e_http import unwrap +from idp import Identity, Keycloak +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, JwtKeyMappingRow, KeyGenerateBody, TeamNewBody, UserNewBody +from other_client import OtherClient +from owned_jwt_gateway import MODEL_NAME, OwnedJwtGateway, owned_jwt_gateway + +pytestmark = pytest.mark.e2e + +_JWT_COMMON: Final = ( + "user_id_jwt_field: sub\n" + "user_email_jwt_field: email\n" + "team_ids_jwt_field: groups\n" + "user_id_upsert: true\n" + "virtual_key_claim_field: sub\n" + "unregistered_jwt_client_behavior: auto_register" +) + + +def _key_hash(key: str) -> str: + return hashlib.sha256(key.encode()).hexdigest() + + +def _ping() -> ChatBody: + return ChatBody( + model=MODEL_NAME, + messages=[ChatMessage(role="user", content=f"Reply with the single word ok. {unique_marker()}")], + max_tokens=5, + ) + + +def _identity_with_user(idp: Keycloak, client: OtherClient, resources: ResourceManager) -> Identity: + """An IdP identity plus the litellm user and team its claims resolve to, with + teardown that also sweeps the user's keys and JWT mapping rows the proxy + wrote, since those outlive the user row itself.""" + marker: Final = unique_marker() + identity: Final = idp.provision(marker=marker, group=f"e2e-jwt-team-{marker}", defer=resources.defer) + resources.defer(lambda: client.proxy.delete_user(identity.user_id)) + team_id: Final = client.proxy.create_team(TeamNewBody(team_alias=f"e2e-jwt-{marker}", team_id=identity.group)) + resources.defer(lambda: client.proxy.delete_team(team_id)) + unwrap( + client.user_new( + UserNewBody( + user_id=identity.user_id, + user_email=f"{identity.username}@example.com", + user_role="internal_user", + auto_create_key=False, + ) + ) + ) + + def delete_user_keys() -> None: + for row in unwrap(client.user_info(identity.user_id)).keys: + client.proxy.delete_key(row.token) + + def delete_user_mappings() -> None: + for mapping in unwrap(client.jwt_mapping_list()).mappings: + if mapping.jwt_claim_value == identity.user_id: + _ = client.jwt_mapping_delete(mapping.id) + + resources.defer(delete_user_keys) + resources.defer(delete_user_mappings) + return identity + + +def _mapping_for(client: OtherClient, claim_value: str) -> JwtKeyMappingRow | None: + return next( + (row for row in unwrap(client.jwt_mapping_list()).mappings if row.jwt_claim_value == claim_value), + None, + ) + + +@pytest.fixture(scope="module") +def mapping_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedJwtGateway]: + with ExitStack() as cleanup: + yield owned_jwt_gateway( + idp, + tmp_path_factory.mktemp("jwt-mapping"), + cleanup, + litellm_jwtauth=f"{_JWT_COMMON}\nauto_register_map_existing_key: true", + name="jwt-mapping-gateway", + ) + + +@pytest.fixture(scope="module") +def minting_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedJwtGateway]: + with ExitStack() as cleanup: + yield owned_jwt_gateway( + idp, + tmp_path_factory.mktemp("jwt-minting"), + cleanup, + litellm_jwtauth=_JWT_COMMON, + name="jwt-minting-gateway", + ) + + +@pytest.mark.owned_gateway +class TestJwtAutoRegisterMapExistingKey: + @pytest.mark.covers("other.auth.jwt.auto_register_maps_existing_key") + def test_first_jwt_call_maps_to_the_users_existing_key_and_mints_none( + self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway + ) -> None: + identity: Final = _identity_with_user(idp, client, resources) + existing_key: Final = client.proxy.generate_key( + KeyGenerateBody( + user_id=identity.user_id, team_id=identity.group, key_alias=f"e2e-jwt-existing-{unique_marker()}" + ) + ) + resources.defer(lambda: client.proxy.delete_key(existing_key)) + + response: Final = unwrap(mapping_gateway.proxy.chat(idp.access_token(identity), _ping())) + + keys: Final = unwrap(client.user_info(identity.user_id)).keys + assert [row.token for row in keys] == [_key_hash(existing_key)], ( + f"map_existing_key must leave the user with only their pre-existing key, got {keys}" + ) + mapping: Final = _mapping_for(client, identity.user_id) + assert mapping is not None, ( + f"no JWT mapping row for sub={identity.user_id}: {unwrap(client.jwt_mapping_list())}" + ) + assert mapping.jwt_claim_name == "sub", f"mapping must bind the sub claim, got {mapping}" + assert mapping.created_by == "auto_register", f"mapping must be written by auto_register, got {mapping}" + rows: Final = client.proxy.poll_logs_for_key(existing_key) + assert any(row.request_id == response.id for row in rows), ( + f"the JWT chat must be billed to the user's existing key, spend rows for it: {rows}" + ) + + @pytest.mark.covers("other.auth.jwt.auto_register_mints_when_keyless") + def test_first_jwt_call_mints_a_key_when_the_user_has_none( + self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway + ) -> None: + identity: Final = _identity_with_user(idp, client, resources) + + response: Final = unwrap(mapping_gateway.proxy.chat(idp.access_token(identity), _ping())) + assert response.choices, f"JWT chat returned no completion: {response}" + + keys: Final = unwrap(client.user_info(identity.user_id)).keys + assert len(keys) == 1, f"a keyless user must get exactly one minted key, got {keys}" + mapping: Final = _mapping_for(client, identity.user_id) + assert mapping is not None and mapping.jwt_claim_name == "sub", ( + f"the minted key must be recorded as a sub-claim mapping, mappings: {unwrap(client.jwt_mapping_list())}" + ) + + @pytest.mark.covers("other.auth.jwt.auto_register_default_mints") + def test_default_behavior_still_mints_when_the_user_already_has_a_key( + self, client: OtherClient, idp: Keycloak, resources: ResourceManager, minting_gateway: OwnedJwtGateway + ) -> None: + identity: Final = _identity_with_user(idp, client, resources) + existing_key: Final = client.proxy.generate_key( + KeyGenerateBody(user_id=identity.user_id, key_alias=f"e2e-jwt-existing-{unique_marker()}") + ) + resources.defer(lambda: client.proxy.delete_key(existing_key)) + + response: Final = unwrap(minting_gateway.proxy.chat(idp.access_token(identity), _ping())) + assert response.id is not None, f"JWT chat returned no response id: {response}" + + keys: Final = unwrap(client.user_info(identity.user_id)).keys + assert len(keys) == 2, ( + f"default auto_register must mint a second key for a user who already has one, got {keys}" + ) + rows: Final = client.proxy.poll_logs_for_request_id(response.id) + assert rows and all(row.api_key != _key_hash(existing_key) for row in rows), ( + f"the default path must bill the minted key, not the user's existing one: {rows}" + ) diff --git a/tests/e2e/pytest.ini b/tests/e2e/pytest.ini index e795ebe5721..dbd3ff47daa 100644 --- a/tests/e2e/pytest.ini +++ b/tests/e2e/pytest.ini @@ -16,6 +16,7 @@ markers = quiet_stack: measures the proxy itself, so it runs while no other test on this host is hitting the stack; every other test waits for it to finish mcp_oauth_live: real Linear OAuth consent via a captured browser session; deselected unless E2E_MCP_OAUTH_LIVE is set provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set + owned_gateway: boots its own proxy from source against the stack's Postgres, so it needs DATABASE_URL on the pytest host; deselected unless E2E_OWNED_GATEWAY is set otel_v2: needs a proxy running with LITELLM_OTEL_V2=true; deselected unless E2E_OTEL_V2 is set otel_tls: needs a stack whose gateway exports OTLP over TLS signed by the CA in SSL_CERT_FILE; deselected unless E2E_OTEL_EXPORTER_ENDPOINT is set secret_manager: needs a proxy booted from gateway/secret_manager__ci_config.yml against that live secret manager; deselected unless E2E_SECRET_MANAGER names the backend (see secret_manager/secret_backends.py) diff --git a/tests/integration/_support/database_relay.py b/tests/integration/_support/database_relay.py index 1b3bc0183a5..cb8b9098a64 100644 --- a/tests/integration/_support/database_relay.py +++ b/tests/integration/_support/database_relay.py @@ -88,15 +88,89 @@ class DatabaseRelay: ) +class HeldStatementRelay: + def __init__(self, upstream_host: str, upstream_port: int, trigger: bytes) -> None: + self.port: Final = _free_port() + self._upstream_host: Final = upstream_host + self._upstream_port: Final = upstream_port + self._trigger: Final = trigger + self._loop: Final = asyncio.new_event_loop() + self._released: Final = asyncio.Event() + self.held: Final = threading.Event() + self._ready: Final = threading.Event() + self._thread: Final = threading.Thread(target=self._run, daemon=True) + + def release(self) -> None: + self._loop.call_soon_threadsafe(self._released.set) + + def start(self) -> None: + self._thread.start() + assert self._ready.wait(10), "Database relay did not start" + + def stop(self) -> None: + self.release() + self._loop.call_soon_threadsafe(self._loop.stop) + self._thread.join(10) + + def _run(self) -> None: + asyncio.set_event_loop(self._loop) + self._loop.run_until_complete(asyncio.start_server(self._serve, "127.0.0.1", self.port)) + self._ready.set() + self._loop.run_forever() + + def _holds(self, window: bytes) -> bool: + return not self.held.is_set() and self._trigger in window + + async def _serve(self, client_reader: asyncio.StreamReader, client_writer: asyncio.StreamWriter) -> None: + server_reader, server_writer = await asyncio.open_connection(self._upstream_host, self._upstream_port) + + async def forward(reader: asyncio.StreamReader, writer: asyncio.StreamWriter, inspect: bool) -> None: + tail = b"" # rebind-ok: carries the previous read's end so a trigger split across reads still matches + try: + while chunk := await reader.read(65536): + window: Final = tail + chunk + if inspect and self._holds(window): + self.held.set() + await self._released.wait() + tail = window[-(len(self._trigger) - 1) :] + writer.write(chunk) + await writer.drain() + except (ConnectionError, asyncio.IncompleteReadError): + return + finally: + writer.close() + + await asyncio.gather( + forward(client_reader, server_writer, True), + forward(server_reader, client_writer, False), + ) + + +def _relayed_url(database_url: str, port: int) -> str: + parts: Final = urlsplit(database_url) + credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else "" + return urlunsplit(parts._replace(netloc=f"{credentials}127.0.0.1:{port}")) + + @contextmanager def database_relay(database_url: str, trigger: bytes) -> Generator[tuple[DatabaseRelay, str]]: parts: Final = urlsplit(database_url) assert parts.hostname is not None and parts.port is not None, database_url relay: Final = DatabaseRelay(parts.hostname, parts.port, trigger) relay.start() - credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else "" - relayed: Final = urlunsplit(parts._replace(netloc=f"{credentials}127.0.0.1:{relay.port}")) try: - yield relay, relayed + yield relay, _relayed_url(database_url, relay.port) + finally: + relay.stop() + + +@contextmanager +def held_statement_relay(database_url: str, trigger: bytes) -> Generator[tuple[HeldStatementRelay, str]]: + parts: Final = urlsplit(database_url) + assert parts.hostname is not None and parts.port is not None, database_url + relay: Final = HeldStatementRelay(parts.hostname, parts.port, trigger) + relay.start() + try: + yield relay, _relayed_url(database_url, relay.port) finally: relay.stop() diff --git a/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py b/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py new file mode 100644 index 00000000000..5215054a364 --- /dev/null +++ b/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py @@ -0,0 +1,219 @@ +import json +import os +import time +import uuid +from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import httpx +import jwt +import pytest +import yaml +from cryptography.hazmat.primitives.asymmetric import rsa + +from tests.integration._support.client import Gateway, eventually, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.database_relay import held_statement_relay +from tests.integration._support.process import owned_proxy +from tests.integration._support.wire import Reply, Request, wire_server + +KEY_ID: Final = "integration-jwt-map-existing-key" +MAPPING_INSERT: Final = b'INSERT INTO "public"."LiteLLM_JWTKeyMapping"' + +pytestmark = pytest.mark.timeout(240) + + +def _hash(key: str) -> str: + return sha256(key.encode()).hexdigest() + + +def _config(directory: Path, claim_field: str) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"] = { + **config["general_settings"], + "enable_jwt_auth": True, + "litellm_jwtauth": { + "user_id_jwt_field": "sub", + "user_email_jwt_field": "email", + "virtual_key_claim_field": claim_field, + "unregistered_jwt_client_behavior": "auto_register", + "auto_register_map_existing_key": True, + }, + } + path: Final = directory / f"jwt_map_existing_key_{claim_field}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@contextmanager +def _issuer() -> Iterator[tuple[rsa.RSAPrivateKey, str]]: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = json.loads(jwt.algorithms.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.method == "GET", request + return Reply(body=jwks) + + with wire_server(respond) as server: + yield private_key, server.url + + +def _token(private_key: rsa.RSAPrivateKey, subject: str, **claims: str) -> str: + now: Final = int(time.time()) + return jwt.encode( + {"sub": subject, **claims, "iat": now, "exp": now + 300}, + private_key, + algorithm="RS256", + headers={"kid": KEY_ID}, + ) + + +def _chat(candidate: Gateway, model: str, token: str) -> httpx.Response: + return candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "map existing key control"}]}, + key=token, + ) + + +def _mapped_token(claim_name: str, claim_value: str) -> str: + rows: Final = read_rows( + 'SELECT token FROM "LiteLLM_JWTKeyMapping" WHERE jwt_claim_name = %s AND jwt_claim_value = %s', + (claim_name, claim_value), + ) + assert len(rows) == 1, rows + return string_value(rows[0]["token"]) + + +def _user_key_hashes(user: str) -> frozenset[str]: + rows: Final = read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE user_id = %s', (user,)) + return frozenset(string_value(row["token"]) for row in rows) + + +def _billed_key(response: httpx.Response) -> str: + rows: Final = eventually( + lambda: read_rows( + 'SELECT api_key FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (str(response.json()["id"]),) + ), + lambda values: len(values) == 1, + seconds=70, + ) + return string_value(rows[0]["api_key"]) + + +def test_first_jwt_call_reuses_the_newest_durable_llm_key_and_skips_every_ineligible_newer_key( + gateway: Gateway, tmp_path: Path +) -> None: + with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario: + model: Final = scenario.model() + user: Final = scenario.user(user_role="internal_user") + older_durable: Final = scenario.key(user_id=user) + durable: Final = scenario.key(user_id=user) + skipped: Final = { + "older_durable": older_durable, + "expiring": scenario.key(user_id=user, duration="1h"), + "management_only": scenario.key(user_id=user, allowed_routes=["management_routes"]), + "auto_registered_look_alike": scenario.key(user_id=user, metadata={"auto_registered": True}), + "other_team": scenario.key(user_id=user, team_id=scenario.team()), + "blocked": scenario.key(user_id=user), + } + gateway.post("/key/block", {"key": skipped["blocked"]}) + keys_before: Final = _user_key_hashes(user) + + with owned_proxy( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "sub") + ) as candidate: + response: Final = _chat(candidate, model, _token(private_key, user)) + + assert response.status_code == 200, response.text + mapped: Final = _mapped_token("sub", user) + assert mapped == _hash(durable), { + "mapped_to": next((name for name, key in skipped.items() if _hash(key) == mapped), mapped) + } + assert _user_key_hashes(user) == keys_before, "a key was minted although a reusable one existed" + assert _billed_key(response) == _hash(durable) + + +def test_user_matched_by_email_instead_of_sub_still_reuses_their_existing_key(gateway: Gateway, tmp_path: Path) -> None: + with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario: + model: Final = scenario.model() + email: Final = f"integration-{uuid.uuid4().hex}@example.com" + user: Final = scenario.user(user_role="internal_user", user_email=email) + existing: Final = scenario.key(user_id=user) + subject: Final = f"integration-idp-subject-{uuid.uuid4().hex}" + + with owned_proxy( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "sub") + ) as candidate: + response: Final = _chat(candidate, model, _token(private_key, subject, email=email.upper())) + + assert response.status_code == 200, response.text + assert _mapped_token("sub", subject) == _hash(existing) + assert _user_key_hashes(user) == frozenset({_hash(existing)}), "a key was minted for an email-matched user" + assert _billed_key(response) == _hash(existing) + + +def test_shared_client_claim_never_maps_a_second_user_onto_the_first_users_personal_key( + gateway: Gateway, tmp_path: Path +) -> None: + with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario: + model: Final = scenario.model() + first_user: Final = scenario.user(user_role="internal_user") + second_user: Final = scenario.user(user_role="internal_user") + personal: Final = scenario.key(user_id=first_user) + client_id: Final = f"integration-shared-client-{uuid.uuid4().hex}" + + with owned_proxy( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "client_id") + ) as candidate: + first: Final = _chat(candidate, model, _token(private_key, first_user, client_id=client_id)) + second: Final = _chat(candidate, model, _token(private_key, second_user, client_id=client_id)) + + assert first.status_code == 200, first.text + assert second.status_code == 200, second.text + mapped: Final = _mapped_token("client_id", client_id) + assert mapped != _hash(personal), "the shared client claim was mapped to the first user's personal key" + assert (_billed_key(first), _billed_key(second)) == (mapped, mapped) + assert read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE api_key = %s', (_hash(personal),)) == [] + + +def test_concurrent_first_jwt_calls_of_a_keyless_user_both_succeed_on_one_surviving_mapped_key( + gateway: Gateway, tmp_path: Path +) -> None: + writer_url: Final = os.environ.get("INTEGRATION_PROXY_DATABASE_URL") or os.environ["DATABASE_URL"] + with ( + _issuer() as (private_key, jwks_url), + gateway.scenario() as scenario, + held_statement_relay(writer_url, MAPPING_INSERT) as (relay, relayed_url), + ): + model: Final = scenario.model() + user: Final = scenario.user(user_role="internal_user") + token: Final = _token(private_key, user) + overrides: Final = { + "JWT_PUBLIC_KEY_URL": jwks_url, + "DATABASE_URL": relayed_url, + "PRISMA_HEALTH_WATCHDOG_ENABLED": "false", + } + + with ( + owned_proxy(gateway, tmp_path, overrides, config=_config(tmp_path, "sub")) as candidate, + ThreadPoolExecutor(max_workers=1) as pool, + ): + held_call: Final = pool.submit(_chat, candidate, model, token) + assert relay.held.wait(60), "the first call never reached its mapping insert" + racing: Final = _chat(candidate, model, token) + relay.release() + held: Final = held_call.result(timeout=60) + + assert racing.status_code == 200, racing.text + assert held.status_code == 200, held.text + keys: Final = _user_key_hashes(user) + assert len(keys) == 1, keys + assert _mapped_token("sub", user) in keys + assert (_billed_key(held), _billed_key(racing)) == (_mapped_token("sub", user),) * 2 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 781d0a13bfd..a7219ac059b 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 @@ -2020,6 +2020,454 @@ async def test_auto_register_binds_api_key_to_token_hash(): assert result.end_user_id == "validated-end-user" +def _auto_register_patches(*, plaintext_key: str | None = "sk-minted-plaintext"): + from litellm.proxy.auth.auth_method import AuthMethod + from litellm.proxy.auth.resolvers.models import CredentialRef + from litellm.proxy.auth.resolvers.store import IdentityStore + from litellm.proxy.proxy_server import hash_token + + resolved_key = UserAPIKeyAuth( + token="existing-hash" if plaintext_key is None else hash_token(plaintext_key), + user_id="validated-user", + team_id="validated-team", + org_id="key-own-org", + ) + principal = IdentityStore._principal_from_key( + resolved_key, + auth_method=AuthMethod.API_KEY, + credential_ref=CredentialRef(token_id=resolved_key.token), + ) + return ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + new_callable=AsyncMock, + return_value={"token": plaintext_key}, + ), + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore.resolve", + new_callable=AsyncMock, + return_value=principal, + ), + ) + + +def _auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler, **over): + kwargs = { + "virtual_key_claim_field": "sub", + "claim_value": "validated-user", + "jwt_handler": jwt_handler, + "prisma_client": prisma_client, + "user_api_key_cache": user_api_key_cache, + "parent_otel_span": None, + "proxy_logging_obj": MagicMock(), + "cache_key": "jwt_key_mapping:sub:validated-user", + "team_id": "validated-team", + "user_id": "validated-user", + "org_id": "jwt-org", + "end_user_id": "validated-end-user", + } + kwargs.update(over) + return kwargs + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_reuses_users_key_but_never_an_auto_registered_one(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[ + {"token": "auto-registered-hash", "metadata": {"auto_registered": True}}, + {"token": "existing-hash", "metadata": {}}, + ] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None) + with generate_patch as generate_key, resolve_patch: + result = await _auto_register_jwt_mapping( + **_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler) + ) + + generate_key.assert_not_awaited() + + create_data = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"] + assert create_data["token"] == "existing-hash" + assert create_data["created_by"] == "auto_register" + assert user_api_key_cache.async_set_cache.await_args.kwargs["value"] == "existing-hash" + assert result is not None + assert result.token == "existing-hash" + assert result.api_key == "existing-hash" + assert result.org_id == "key-own-org" + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_mints_when_user_has_no_key(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + from litellm.proxy.proxy_server import hash_token + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches() + with generate_patch as generate_key, resolve_patch: + result = await _auto_register_jwt_mapping( + **_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler) + ) + + generate_key.assert_awaited_once() + create_data = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"] + assert create_data["token"] == hash_token("sk-minted-plaintext") + assert result is not None + assert result.token == hash_token("sk-minted-plaintext") + + +@pytest.mark.asyncio +async def test_auto_register_default_never_looks_up_existing_keys(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="sub", virtual_key_mapping_cache_ttl=300) + + generate_patch, resolve_patch = _auto_register_patches() + with generate_patch as generate_key, resolve_patch: + await _auto_register_jwt_mapping(**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler)) + + prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited() + generate_key.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_race_loser_keeps_reused_key(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_verificationtoken.delete = AsyncMock() + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock(side_effect=Exception("Unique constraint failed (P2002)")) + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None) + with ( + generate_patch, + resolve_patch, + patch( + "litellm.proxy.auth.user_api_key_auth.get_jwt_key_mapping_object", + new_callable=AsyncMock, + return_value="winner-hash", + ), + ): + result = await _auto_register_jwt_mapping( + **_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler) + ) + + assert result is not None + assert result.org_id == "key-own-org" + prisma_client.db.litellm_verificationtoken.delete.assert_not_awaited() + assert user_api_key_cache.async_set_cache.await_args.kwargs["value"] == "winner-hash" + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_user_id_none_mints(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches() + with generate_patch as generate_key, resolve_patch: + await _auto_register_jwt_mapping( + **_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler, user_id=None) + ) + + prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited() + generate_key.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_reuses_when_the_user_was_matched_by_a_fallback_lookup(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + user_email_jwt_field="email", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None) + with generate_patch as generate_key, resolve_patch: + await _auto_register_jwt_mapping( + **_auto_register_kwargs( + prisma_client, + user_api_key_cache, + jwt_handler, + claim_value="idp-subject-not-the-db-user-id", + cache_key="jwt_key_mapping:sub:idp-subject-not-the-db-user-id", + ) + ) + + generate_key.assert_not_awaited() + assert prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] == "existing-hash" + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_mints_when_the_claim_is_not_a_user_identity_field(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + from litellm.proxy.proxy_server import hash_token + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches() + with generate_patch as generate_key, resolve_patch: + await _auto_register_jwt_mapping( + **_auto_register_kwargs( + prisma_client, + user_api_key_cache, + jwt_handler, + virtual_key_claim_field="azp", + claim_value="shared-client-app", + cache_key="jwt_key_mapping:azp:shared-client-app", + ) + ) + + prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited() + generate_key.assert_awaited_once() + assert prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] == hash_token( + "sk-minted-plaintext" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("issuer_user_id_field", "expect_reuse"), + [("uid", False), (None, True)], +) +async def test_auto_register_map_existing_key_uses_the_issuers_own_user_field_over_the_global_one( + issuer_user_id_field, expect_reuse +): + from litellm.proxy._types import JWTIssuerConfig + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + from litellm.proxy.proxy_server import hash_token + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + issuers=[ + JWTIssuerConfig( + issuer="https://idp.example.com", audience="litellm", user_id_jwt_field=issuer_user_id_field + ) + ], + ) + + generate_patch, resolve_patch = _auto_register_patches() + with generate_patch as generate_key, resolve_patch: + await _auto_register_jwt_mapping( + **_auto_register_kwargs( + prisma_client, user_api_key_cache, jwt_handler, jwt_issuer="https://idp.example.com" + ) + ) + + mapped_token = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] + assert mapped_token == ("existing-hash" if expect_reuse else hash_token("sk-minted-plaintext")) + assert generate_key.await_count == (0 if expect_reuse else 1) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("map_existing_key", "master_key", "reused_key_models", "expect_denied"), + [ + (True, "sk-master", ["some-other-model"], True), + (True, "sk-master", [], False), + (False, "sk-master", ["some-other-model"], False), + (True, None, ["some-other-model"], False), + ], +) +async def test_auto_register_map_existing_key_first_request_runs_key_checks( + map_existing_key: bool, master_key: str | None, reused_key_models: list[str], expect_denied: bool +) -> None: + jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature" + user_api_key_cache = DualCache() + prisma_client = MagicMock() + jwt_handler = MagicMock() + jwt_handler.is_jwt.return_value = True + jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "user1"}) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + virtual_key_claim_field="sub", + virtual_key_mapping_cache_ttl=300, + auto_register_map_existing_key=map_existing_key, + ) + reused_key = UserAPIKeyAuth( + token="hashed-existing-key", + api_key="hashed-existing-key", + user_id="validated-user", + team_id="validated-team", + models=reused_key_models, + ) + mock_jwt_result = { + "is_proxy_admin": False, + "team_object": None, + "user_object": LiteLLM_UserTable(user_id="validated-user", user_role="internal_user"), + "end_user_object": None, + "org_object": None, + "token": jwt_token, + "team_id": "validated-team", + "user_id": "validated-user", + "user_email": None, + "end_user_id": None, + "org_id": None, + "team_membership": None, + "jwt_claims": {"sub": "user1"}, + } + + mock_request = MagicMock() + mock_request.url.path = "/v1/chat/completions" + mock_request.method = "POST" + mock_request.headers = {"authorization": f"Bearer {jwt_token}"} + mock_request.query_params = {} + mock_request.state = SimpleNamespace() + + with ( + patch("litellm.proxy.proxy_server.general_settings", {"enable_jwt_auth": True}), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", master_key), + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + MagicMock(post_call_failure_hook=AsyncMock(return_value=None)), + ), + patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler), + patch( + "litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key", + new_callable=AsyncMock, + return_value=_PendingAutoRegister( + claim_field="sub", + claim_value="user1", + cache_key="jwt_key_mapping:sub:user1", + ), + ), + patch( + "litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", + new_callable=AsyncMock, + return_value=mock_jwt_result, + ), + patch( + "litellm.proxy.auth.user_api_key_auth._auto_register_jwt_mapping", + new_callable=AsyncMock, + return_value=reused_key, + ), + ): + call = _user_api_key_auth_builder( + request=mock_request, + api_key=jwt_token, + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={"model": "gpt-4o-mini"}, + ) + if expect_denied: + with pytest.raises(ProxyException, match="not available for this API key"): + await call + return + result = await call + + assert result.api_key == "hashed-existing-key" + assert result.user_id == "validated-user" + assert result.team_id == "validated-team" + assert result.models == reused_key_models + + @pytest.mark.asyncio @pytest.mark.parametrize("active", [True, False]) async def test_auto_register_first_request_propagates_user_email(active: bool) -> None: From b1e0e9e84b6710f45d746111b47fd21e04bada06 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Fri, 2 Oct 2026 12:17:55 -0700 Subject: [PATCH 020/139] feat(ui): select Laya for OSS classification (#43768) --- .../AutoRouters/autoRouterRows.test.ts | 3 +- .../components/AutoRouters/autoRouterRows.ts | 3 +- .../add_model/AutoRouterAvailability.tsx | 2 +- ...oRouterClassifierTabs.integration.test.tsx | 2 +- .../add_model/AutoRouterClassifierTabs.tsx | 30 ++++- .../add_model/ClassificationMethodConfig.tsx | 4 +- .../add_model/ClassifierTypeRadios.tsx | 4 +- .../add_model/ComplexityRouterConfig.tsx | 2 +- .../JevClassifierConfig.integration.test.tsx | 111 ++++++++++-------- .../add_model/JevClassifierConfig.tsx | 35 ++++-- .../JevConnectionTest.integration.test.tsx | 8 +- .../add_auto_router_tab.integration.test.tsx | 39 +++--- .../add_model/auto_router_connection_test.tsx | 10 +- ...d_auto_router_routing_test_request.test.ts | 14 +-- .../build_auto_router_routing_test_request.ts | 15 ++- .../build_complexity_router_config.test.ts | 22 ++-- .../build_complexity_router_config.ts | 16 ++- .../add_model/jev_classifier_config.ts | 34 +++++- ...d_updated_complexity_router_config.test.ts | 41 ++++--- .../edit_auto_router_modal.tsx | 1 + .../hydrate_complexity_router_config.ts | 12 +- .../LogDetailsDrawer/RoutingDecisionCard.tsx | 2 +- .../src/lib/autorouter_presets.test.ts | 13 +- .../src/lib/autorouter_presets.ts | 7 +- 24 files changed, 280 insertions(+), 150 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts index 79c4243271e..153b77666a7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts @@ -85,7 +85,8 @@ describe("autoRouterRows", () => { it.each([ ["llm", "LLM Classifier"], - ["jev", "JEV Classifier"], + ["jev", "OSS Classifier"], + ["oss_classifier", "OSS Classifier"], ])("labels a router using the %s classifier", (classifierType, label) => { const row = toAutoRouterRow( { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts index 1faf3408c23..3c3366a767d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts @@ -57,7 +57,8 @@ const dedupe = (models: string[]): string[] => Array.from(new Set(models)); const COMPLEXITY_TYPE_LABELS: Record = { llm: "LLM Classifier", - jev: "JEV Classifier", + jev: "OSS Classifier", + oss_classifier: "OSS Classifier", capability: "Capability", llm_v2: "Fuse v2", heuristic_first: "Heuristic first", diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx index b1233ef03b8..135ebb958db 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx @@ -136,7 +136,7 @@ export const AutoRouterLimits = () => { Routing and customization limits

- Rule-based, Complexity, and Jev are unlimited with built-in settings. Choose or change tier models freely. + Rule-based, Complexity, and OSS are unlimited with built-in settings. Choose or change tier models freely. Customization allowances are shared across this proxy.

diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx index f843a472d15..08f329e4ad7 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx @@ -68,7 +68,7 @@ describe("Auto-router classifier selection", () => { llm: "LLM", heuristic_first: "LLM", hybrid: "LLM", - jev: "Jev", + jev: "OSS Classifier", }[classifier_type]; expect(screen.getByRole("radio", { name: new RegExp(`^${family}$`) })).toBeChecked(); fireEvent.click(screen.getByRole("radio", { name: new RegExp(`^${family}$`) })); diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx index 92fc8d2a335..3a1e0065530 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx @@ -16,6 +16,7 @@ import { type ClassifierType, type ComplexityRouterConfigValue, } from "./ComplexityRouterConfig"; +import { defaultJevClassifierConfig, normalizeJevClassifierConfig } from "./jev_classifier_config"; import { transitionClassifierType } from "./classifier_type_transition"; import { isForecastClassifier } from "./forecast_classifier_config"; import { @@ -148,6 +149,14 @@ const AutoRouterClassifierTabs: React.FC = ({ val if (next === "llm") changeType("llm"); if (next === "jev") changeType("jev"); }; + const changeProvider = (provider: unknown) => { + if (provider !== "jev" && provider !== "laya") return; + const defaults = defaultJevClassifierConfig(provider); + onChange({ + ...value, + jev_classifier_config: { ...defaults, ...value.jev_classifier_config, provider, model: defaults.model }, + }); + }; const approachLabels: Partial> = { capability: "Capability", llm_v2: "Fuse v2" }; const approachDescription: Partial> = { capability: "Use the efficient model when it is likely to succeed", @@ -164,7 +173,7 @@ const AutoRouterClassifierTabs: React.FC = ({ val {[ { value: "heuristics", label: "Heuristics", description: "Classify locally, with no API call" }, { value: "llm", label: "LLM", description: "Use a judge model to choose a solver" }, - { value: "jev", label: "Jev", description: "Use TypeSafe System One Choice to choose a tier" }, + { value: "jev", label: "OSS Classifier", description: "Use Jev or Laya to choose a tier" }, ].map((option) => (
diff --git a/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx b/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx index 1fd6dfa6a20..2eda2c63f64 100644 --- a/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx @@ -54,8 +54,8 @@ const ClassifierTypeRadios: React.FC = ({ value, clas diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index 32f9ebf97ad..e324886549b 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -237,7 +237,7 @@ const TierSetToolbar: React.FC<{ {editing && ( Add or remove tiers to define your own set. Every custom tier needs a definition the classifier routes on, and - an edited set requires the LLM or Jev classification method + an edited set requires the LLM or OSS classification method )} {editing && keywordRulesError && ( diff --git a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx index 7da8b12c2d7..e175fc934b9 100644 --- a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx @@ -1,4 +1,5 @@ import React, { useState } from "react"; +import userEvent from "@testing-library/user-event"; import { afterEach, describe, expect, it, vi } from "vitest"; import { fireEvent, renderWithProviders, screen } from "../../../tests/test-utils"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; @@ -96,49 +97,65 @@ function Form() { describe("JEV classifier editor", () => { afterEach(() => vi.mocked(useAuthorized).mockReset()); - it("uses built-in JEV without a license and preserves custom tiers and context through reload", () => { - renderWithProviders(
); - expect(screen.getByLabelText("Judge model")).toBeInTheDocument(); - expect(screen.getByText("Reasoning Effort")).toBeInTheDocument(); - expect(screen.getByText("Classifier Prompt")).toBeInTheDocument(); - expect(screen.getByRole("switch", { name: "Use images for classification" })).toBeInTheDocument(); - fireEvent.click(screen.getByRole("radio", { name: /Jev Classifier/ })); - expect(screen.getByRole("radio", { name: /^Jev Classifier/ })).toBeChecked(); - expect(screen.getByLabelText("Jev Model")).toHaveValue("jev-latest"); - expect(screen.getByLabelText("Jev Instructions")).toBeEnabled(); - expect(screen.queryByLabelText("Judge model")).not.toBeInTheDocument(); - expect(screen.queryByText("Reasoning Effort")).not.toBeInTheDocument(); - expect(screen.queryByText("Classifier Prompt")).not.toBeInTheDocument(); - expect(screen.queryByRole("switch", { name: "Use images for classification" })).not.toBeInTheDocument(); - fireEvent.change(screen.getByLabelText("Jev Model"), { target: { value: "jev-test" } }); - fireEvent.change(screen.getByLabelText("Jev Timeout (ms)"), { target: { value: "4200" } }); - fireEvent.change(screen.getByLabelText("Context Window Size"), { target: { value: "6" } }); - fireEvent.change(screen.getByLabelText("Circuit breaker cooldown (seconds)"), { target: { value: "50" } }); - fireEvent.click(screen.getByRole("switch", { name: "Classifier circuit breaker" })); - fireEvent.click(screen.getByRole("button", { name: "Customize tiers" })); - fireEvent.click(screen.getByRole("button", { name: "Save and reload" })); - expect(screen.getByRole("radio", { name: /Jev Classifier/ })).toBeChecked(); - expect(screen.getByLabelText("Jev Model")).toHaveValue("jev-test"); - expect(screen.getByLabelText("Jev Timeout (ms)")).toHaveValue(4200); - expect(screen.getByLabelText("Context Window Size")).toHaveValue("6"); - expect(screen.getByRole("switch", { name: "Classifier circuit breaker" })).not.toBeChecked(); - fireEvent.click(screen.getByRole("button", { name: "Probe current config" })); - expect(testAutoRouterRouting).toHaveBeenCalledWith( - "token", - expect.objectContaining({ - complexity_router_config: expect.objectContaining({ - classifier_type: "jev", - jev_classifier_config: { - model: "jev-test", - timeout_ms: 4200, - circuit_breaker_enabled: false, - circuit_breaker_cooldown_seconds: 50, - }, - tiers: expect.objectContaining({ QUICK: ["fast"] }), + it.each(["jev", "laya"] as const)( + "preserves %s, custom tiers and context through save, reload and probe", + async (provider) => { + renderWithProviders(); + expect(screen.getByLabelText("Judge model")).toBeInTheDocument(); + expect(screen.getByText("Reasoning Effort")).toBeInTheDocument(); + expect(screen.getByText("Classifier Prompt")).toBeInTheDocument(); + expect(screen.getByRole("switch", { name: "Use images for classification" })).toBeInTheDocument(); + fireEvent.click(screen.getByRole("radio", { name: /^OSS Classifier$/ })); + expect(screen.getByRole("radio", { name: /^OSS Classifier$/ })).toBeChecked(); + expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-latest"); + expect(screen.getByLabelText("Classifier Instructions")).toBeEnabled(); + expect(screen.queryByLabelText("Judge model")).not.toBeInTheDocument(); + expect(screen.queryByText("Reasoning Effort")).not.toBeInTheDocument(); + expect(screen.queryByText("Classifier Prompt")).not.toBeInTheDocument(); + expect(screen.queryByRole("switch", { name: "Use images for classification" })).not.toBeInTheDocument(); + fireEvent.click(screen.getByRole("radio", { name: "Laya" })); + expect(screen.getByLabelText("Classifier Model")).toHaveTextContent("english"); + fireEvent.click(screen.getByRole("radio", { name: "Jev" })); + expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-latest"); + if (provider === "laya") { + fireEvent.click(screen.getByRole("radio", { name: "Laya" })); + await userEvent.click(screen.getByLabelText("Classifier Model")); + await userEvent.click(screen.getByRole("option", { name: "multilingual" })); + } else { + fireEvent.change(screen.getByLabelText("Classifier Model"), { target: { value: "jev-test" } }); + } + fireEvent.change(screen.getByLabelText("Classifier Timeout (ms)"), { target: { value: "4200" } }); + fireEvent.change(screen.getByLabelText("Context Window Size"), { target: { value: "6" } }); + fireEvent.change(screen.getByLabelText("Circuit breaker cooldown (seconds)"), { target: { value: "50" } }); + fireEvent.click(screen.getByRole("switch", { name: "Classifier circuit breaker" })); + fireEvent.click(screen.getByRole("button", { name: "Customize tiers" })); + fireEvent.click(screen.getByRole("button", { name: "Save and reload" })); + expect(screen.getByRole("radio", { name: /^OSS Classifier$/ })).toBeChecked(); + expect(screen.getByRole("radio", { name: provider === "laya" ? "Laya" : "Jev" })).toBeChecked(); + if (provider === "laya") expect(screen.getByLabelText("Classifier Model")).toHaveTextContent("multilingual"); + else expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-test"); + expect(screen.getByLabelText("Classifier Timeout (ms)")).toHaveValue(4200); + expect(screen.getByLabelText("Context Window Size")).toHaveValue("6"); + expect(screen.getByRole("switch", { name: "Classifier circuit breaker" })).not.toBeChecked(); + fireEvent.click(screen.getByRole("button", { name: "Probe current config" })); + expect(testAutoRouterRouting).toHaveBeenCalledWith( + "token", + expect.objectContaining({ + complexity_router_config: expect.objectContaining({ + classifier_type: "oss_classifier", + opensource_classifier_config: { + provider, + model: provider === "laya" ? "multilingual" : "jev-test", + timeout_ms: 4200, + circuit_breaker_enabled: false, + circuit_breaker_cooldown_seconds: 50, + }, + tiers: expect.objectContaining({ QUICK: ["fast"] }), + }), }), - }), - ); - }); + ); + }, + ); it("allows licensed instructions and can restore built-in instructions", () => { const authorized = useAuthorized(); @@ -152,10 +169,10 @@ describe("JEV classifier editor", () => { return ; }; renderWithProviders(); - expect(screen.getByLabelText("Jev Instructions")).toBeEnabled(); - fireEvent.change(screen.getByLabelText("Jev Instructions"), { target: { value: "New instructions" } }); - expect(screen.getByLabelText("Jev Instructions")).toHaveValue("New instructions"); - fireEvent.click(screen.getByRole("button", { name: "Restore built-in Jev instructions" })); - expect(screen.getByLabelText("Jev Instructions")).toHaveValue(""); + expect(screen.getByLabelText("Classifier Instructions")).toBeEnabled(); + fireEvent.change(screen.getByLabelText("Classifier Instructions"), { target: { value: "New instructions" } }); + expect(screen.getByLabelText("Classifier Instructions")).toHaveValue("New instructions"); + fireEvent.click(screen.getByRole("button", { name: "Restore built-in instructions" })); + expect(screen.getByLabelText("Classifier Instructions")).toHaveValue(""); }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx index 97609bcd8c3..22e8708acc7 100644 --- a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx @@ -6,7 +6,8 @@ import { Label } from "@/components/ui/label"; import { Textarea } from "@/components/ui/textarea"; import ClassifierCircuitBreakerConfig from "./ClassifierCircuitBreakerConfig"; import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; -import { defaultJevClassifierConfig } from "./jev_classifier_config"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { defaultJevClassifierConfig, LAYA_MODELS } from "./jev_classifier_config"; export default function JevClassifierConfig({ value, @@ -17,20 +18,38 @@ export default function JevClassifierConfig({ }) { const id = useId(); const config = value.jev_classifier_config ?? defaultJevClassifierConfig(); + const isLaya = config.provider === "laya"; const update = (patch: Partial) => onChange({ ...value, jev_classifier_config: { ...config, ...patch } }); return (

- Uses TypeSafe System One Choice evaluation with your configured tiers + {isLaya + ? "Uses Laya with your configured tiers. Set LAYA_API_BASE on the gateway to connect your Laya server." + : "Uses TypeSafe System One Choice evaluation with your configured tiers"}

- - update({ model: event.target.value })} /> + + {isLaya ? ( + + ) : ( + update({ model: event.target.value })} /> + )}
- +
- + {config.instructions && ( )}

- Built-in Jev is available without a license and uses the shipped tier criteria + Built-in OSS classification is available without a license and uses the shipped tier criteria

diff --git a/ui/litellm-dashboard/src/components/add_model/JevConnectionTest.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/JevConnectionTest.integration.test.tsx index ff147e20d1a..02b065083a0 100644 --- a/ui/litellm-dashboard/src/components/add_model/JevConnectionTest.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/JevConnectionTest.integration.test.tsx @@ -107,10 +107,10 @@ describe("JEV network probes", () => { expect(JSON.parse(String(routingCall?.[1]?.body))).toEqual(expectedRequest); expect(fetchMock).toHaveBeenCalledTimes(5); expect(screen.getAllByTestId("test-status-success")).toHaveLength(4); - expect(screen.getByRole("status", { name: "Jev connection" })).toHaveTextContent( + expect(screen.getByRole("status", { name: "OSS classifier connection" })).toHaveTextContent( cause === "jev_classifier" - ? "Jev classification succeeded" - : `Jev was not reached successfully (routing cause: ${cause})`, + ? "OSS classification succeeded" + : `OSS classifier was not reached successfully (routing cause: ${cause})`, ); }, ); @@ -131,7 +131,7 @@ describe("JEV network probes", () => { ); fireEvent.change(screen.getByTestId("auto-router-routing-test-prompt"), { target: { value: "Hello" } }); fireEvent.click(screen.getByTestId("auto-router-routing-test-send")); - expect(await screen.findByText("JEV classifier")).toBeInTheDocument(); + expect(await screen.findByText("OSS classifier")).toBeInTheDocument(); expect(screen.getByText("jev-latest")).toBeInTheDocument(); expect(screen.getByText("80.0%")).toBeInTheDocument(); expect(screen.getByText("SIMPLE: 80.0%")).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx index 0337d35d698..4a871d5b103 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx @@ -274,13 +274,13 @@ describe("AddAutoRouterTab", () => { mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS); renderWithProviders(); await user.click(await screen.findByRole("button", { name: "Choose models for me" })); - await user.click(screen.getByRole("radio", { name: "Jev" })); + await user.click(screen.getByRole("radio", { name: "OSS Classifier" })); await waitFor(() => expect(apiClient.post).toHaveBeenLastCalledWith( "/auto_router/availability", expect.objectContaining({ body: expect.objectContaining({ - complexity_router_config: expect.objectContaining({ classifier_type: "jev" }), + complexity_router_config: expect.objectContaining({ classifier_type: "oss_classifier" }), }), }), ), @@ -301,13 +301,13 @@ describe("AddAutoRouterTab", () => { expect(within(screen.getByRole("alert")).getByRole("link", { name: "Talk to our team" })).toBeVisible(); await user.click(screen.getByRole("button", { name: "Restore defaults" })); await waitFor(() => expect(screen.queryByRole("alert")).not.toBeInTheDocument()); - expect(screen.getByRole("radio", { name: "Jev" })).toBeChecked(); + expect(screen.getByRole("radio", { name: "OSS Classifier" })).toBeChecked(); await waitFor(() => expect(screen.getByRole("button", { name: "Add Auto Router" })).toBeEnabled()); await user.click(screen.getByRole("button", { name: "Add Auto Router" })); await waitFor(() => expect(handleAddAutoRouterSubmit).toHaveBeenCalled()); const saved = vi.mocked(handleAddAutoRouterSubmit).mock.calls.at(-1)?.[0].complexity_router_config; expect(saved).not.toHaveProperty("tier_definitions"); - expect(saved?.classifier_type).toBe("jev"); + expect(saved?.classifier_type).toBe("oss_classifier"); expect(Object.keys(saved?.tiers ?? {})).toEqual(["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]); expect(saved?.tiers).toEqual(initialRequest.complexity_router_config.tiers); }); @@ -317,7 +317,7 @@ describe("AddAutoRouterTab", () => { mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS); renderWithProviders(); await user.click(await screen.findByRole("button", { name: "Choose models for me" })); - await user.click(screen.getByRole("radio", { name: "Jev" })); + await user.click(screen.getByRole("radio", { name: "OSS Classifier" })); fireEvent.change(screen.getByLabelText("Auto Router Name"), { target: { value: "checked-router" } }); await waitFor(() => expect(screen.getByRole("button", { name: "Add Auto Router" })).toBeEnabled()); let complete: ((result: unknown) => void) | undefined; @@ -357,17 +357,22 @@ describe("AddAutoRouterTab", () => { expect(screen.getByRole("button", { name: "Routing approach" })).toHaveTextContent("Complexity"); }); - it.each(["LLM", "Jev"])("keeps %s and the frequency when choosing models automatically", async (family) => { - mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS); - renderWithProviders(); - const automatic = await screen.findByRole("button", { name: "Choose models for me" }); - await userEvent.click(screen.getByRole("radio", { name: family })); - await selectAutoRouterOption("How often to classify", "Every new user message"); - await userEvent.click(automatic); - expect(screen.getByRole("radio", { name: family })).toBeChecked(); - expect(screen.getByRole("combobox", { name: "How often to classify" })).toHaveTextContent("Every new user message"); - expect(screen.getByRole("button", { name: "Advanced settings" })).toHaveAttribute("aria-expanded", "false"); - }); + it.each(["LLM", "OSS Classifier"])( + "keeps %s and the frequency when choosing models automatically", + async (family) => { + mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS); + renderWithProviders(); + const automatic = await screen.findByRole("button", { name: "Choose models for me" }); + await userEvent.click(screen.getByRole("radio", { name: family })); + await selectAutoRouterOption("How often to classify", "Every new user message"); + await userEvent.click(automatic); + expect(screen.getByRole("radio", { name: family })).toBeChecked(); + expect(screen.getByRole("combobox", { name: "How often to classify" })).toHaveTextContent( + "Every new user message", + ); + expect(screen.getByRole("button", { name: "Advanced settings" })).toHaveAttribute("aria-expanded", "false"); + }, + ); it.each(["Capability", "Fuse v2"])( "creates %s from its dedicated tab without complexity templates", @@ -1902,7 +1907,7 @@ describe("preset catalog fetch states", () => { await waitFor(() => expect(handleAddAutoRouterSubmit).toHaveBeenCalledOnce()); expect(vi.mocked(handleAddAutoRouterSubmit).mock.calls[0][0].complexity_router_config).toMatchObject({ - classifier_type: "jev", + classifier_type: "oss_classifier", classifier_context_per_turn_chars: 450, }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx index aa2ce91766a..eafa5122b59 100644 --- a/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx @@ -50,7 +50,7 @@ const AutoRouterConnectionTest: React.FC = ({ ? { status: "success" } : { status: "error", - error: `Jev was not reached successfully (routing cause: ${decision.cause ?? "unknown"})`, + error: `OSS classifier was not reached successfully (routing cause: ${decision.cause ?? "unknown"})`, }, ); }; @@ -91,11 +91,11 @@ const AutoRouterConnectionTest: React.FC = ({ classifier probe includes its reasoning effort override.

{jevRequest && ( -
- Jev Classifier +
+ OSS Classifier

- {jevResult.status === "pending" && "Testing Jev classification"} - {jevResult.status === "success" && "Jev classification succeeded"} + {jevResult.status === "pending" && "Testing OSS classification"} + {jevResult.status === "success" && "OSS classification succeeded"} {jevResult.status === "error" && jevResult.error}

diff --git a/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts b/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts index de0fb6fe6e1..3686103f234 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts @@ -33,20 +33,20 @@ describe("buildAutoRouterRoutingTestRequest", () => { const expectedRequest = { prompt: JEV_CONNECTION_TEST_PROMPT, complexity_router_config: { - classifier_type: "jev", + classifier_type: "oss_classifier", tiers: CONFIG.tiers, - jev_classifier_config: defaultJevClassifierConfig(), + opensource_classifier_config: defaultJevClassifierConfig(), }, saved_model_id: "saved-id", }; expect(request).toEqual(expectedRequest); - expect(request?.complexity_router_config.jev_classifier_config).not.toHaveProperty("api_key"); - expect(request?.complexity_router_config.jev_classifier_config).not.toHaveProperty("api_base"); + expect(request?.complexity_router_config.opensource_classifier_config).not.toHaveProperty("api_key"); + expect(request?.complexity_router_config.opensource_classifier_config).not.toHaveProperty("api_base"); }); - it.each(["object", "json"])("probes saved JEV %s configuration with custom tiers and team context", (format) => { + it.each(["object", "json"])("probes saved Laya %s configuration with custom tiers and team context", (format) => { const config = { - classifier_type: "jev", - jev_classifier_config: { model: "jev-test", timeout_ms: 900 }, + classifier_type: "oss_classifier", + opensource_classifier_config: { provider: "laya", model: "english", timeout_ms: 900 }, tiers: { QUICK: ["fast"], DEEP: ["strong"] }, tier_definitions: { QUICK: "Simple questions", DEEP: "Complex questions" }, fallback_tier: "DEEP", diff --git a/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.ts b/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.ts index 6a9d1ce7d92..9a5be0837a0 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.ts @@ -1,7 +1,7 @@ import { AutoRouterRoutingTestRequest } from "../networking"; import { ComplexityRouterConfigPayload } from "./build_complexity_router_config"; import { z } from "zod"; -import { jevClassifierConfigSchema } from "./jev_classifier_config"; +import { hydrateOssClassifier, jevClassifierConfigSchema, normalizeJevClassifierConfig } from "./jev_classifier_config"; export const JEV_CONNECTION_TEST_PROMPT = "What is 2 plus 2?"; @@ -23,16 +23,23 @@ export const buildSavedJevConnectionTestRequest = ( : rawConfig; const result = z .object({ - classifier_type: z.literal("jev"), + classifier_type: z.enum(["jev", "oss_classifier"]), tiers: z.record(z.unknown()), - jev_classifier_config: jevClassifierConfigSchema.default({}), + jev_classifier_config: jevClassifierConfigSchema.optional(), + opensource_classifier_config: jevClassifierConfigSchema.optional(), }) .passthrough() .safeParse(parsed); if (!result.success) return undefined; + const { jev_classifier_config, opensource_classifier_config, ...config } = result.data; + const classifier = hydrateOssClassifier({ ...config, jev_classifier_config, opensource_classifier_config }); return { prompt: JEV_CONNECTION_TEST_PROMPT, - complexity_router_config: result.data, + complexity_router_config: { + ...config, + classifier_type: "oss_classifier", + opensource_classifier_config: normalizeJevClassifierConfig(classifier.jev_classifier_config), + }, saved_model_id: savedModelId, ...(teamId && { team_id: teamId }), }; diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts index 6d458d61797..c054f6bbe9d 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts @@ -64,6 +64,7 @@ describe("buildComplexityRouterConfig", () => { it.each([ { model: "" }, { model: " " }, + { provider: "laya" as const, model: "unsupported" }, { timeout_ms: 0 }, { timeout_ms: 1.5 }, { timeout_ms: Number.NaN }, @@ -75,15 +76,16 @@ describe("buildComplexityRouterConfig", () => { classifier_type: "jev", jev_classifier_config: { model: "jev-latest", timeout_ms: 3000, ...patch }, }), - ).toBe("Enter a JEV model, a positive whole-number timeout and a positive cooldown"); + ).toBe("Enter a valid classifier model, a positive whole-number timeout and a positive cooldown"); }); - it.each([false, true])("serializes JEV with shared context and no LLM config, custom tiers: %s", (custom) => { + it.each([false, true])("serializes Laya with shared context and no LLM config, custom tiers: %s", (custom) => { const params: BuildComplexityRouterConfigParams = { ...baseParams, classifierType: "jev", jevClassifierConfig: { - model: "jev-test", + provider: "laya", + model: "english", timeout_ms: 4500, instructions: " Choose the configured tier ", circuit_breaker_enabled: false, @@ -108,15 +110,17 @@ describe("buildComplexityRouterConfig", () => { }), }; const config = buildComplexityRouterConfig(params); - expect(config.classifier_type).toBe("jev"); + expect(config.classifier_type).toBe("oss_classifier"); const expectedJevConfig = { - model: "jev-test", + provider: "laya", + model: "english", timeout_ms: 4500, instructions: "Choose the configured tier", circuit_breaker_enabled: false, circuit_breaker_cooldown_seconds: 12.5, }; - expect(config.jev_classifier_config).toEqual(expectedJevConfig); + expect(config.opensource_classifier_config).toEqual(expectedJevConfig); + expect(config).not.toHaveProperty("jev_classifier_config"); expect(config.classifier_context_window_size).toBe(4); expect(config.classifier_context_budget_chars).toBe(2000); expect(config.classifier_context_per_turn_chars).toBe(450); @@ -139,15 +143,15 @@ describe("buildComplexityRouterConfig", () => { classifierType: "jev", jevClassifierConfig: { model: "jev-latest", timeout_ms: 3000, instructions: " " }, }); - expect(jev.jev_classifier_config).toEqual({ model: "jev-latest", timeout_ms: 3000 }); + expect(jev.opensource_classifier_config).toEqual({ provider: "jev", model: "jev-latest", timeout_ms: 3000 }); const llmParams: BuildComplexityRouterConfigParams = { ...baseParams, classifierType: "llm", classifierLlmConfig: { model: "judge", timeout_ms: 1000 }, - jevClassifierConfig: jev.jev_classifier_config, + jevClassifierConfig: jev.opensource_classifier_config, }; const llm = buildComplexityRouterConfig(llmParams); - expect(llm).not.toHaveProperty("jev_classifier_config"); + expect(llm).not.toHaveProperty("opensource_classifier_config"); }); it("forwards preset references and explicit overrides without materializing absent text on create", () => { diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts index 13df40464f0..7756d0ee8fa 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts @@ -153,12 +153,13 @@ export interface StoredComplexityRouterConfig { heuristic_first_max_tier?: unknown; hybrid_boundary_margin?: unknown; tier_labels?: unknown; - classifier_type?: ClassifierType; + classifier_type?: ClassifierType | "oss_classifier"; heuristic_v2_success_threshold?: unknown; capability_classifier_config?: unknown; llm_v2_config?: unknown; classifier_llm_config?: ClassifierLLMConfig; jev_classifier_config?: unknown; + opensource_classifier_config?: unknown; classifier_context_window_size?: unknown; classifier_context_budget_chars?: unknown; classifier_context_per_turn_chars?: unknown; @@ -283,12 +284,13 @@ export interface ComplexityRouterConfigPayload { default_model?: string; plan_mode_min_tier?: string; tier_labels?: ComplexityTierLabels; - classifier_type: ClassifierType; + classifier_type: ClassifierType | "oss_classifier"; heuristic_v2_success_threshold?: number; capability_classifier_config?: CapabilitySettings; llm_v2_config?: FuseSettings; classifier_llm_config?: ClassifierLLMConfig; jev_classifier_config?: JevClassifierConfig; + opensource_classifier_config?: JevClassifierConfig; classifier_context_window_size?: number; classifier_context_budget_chars?: number; classifier_context_per_turn_chars?: number; @@ -438,7 +440,9 @@ export const getClassifierModelError = ( ): string | null => { if (effectiveClassifierType(config) === "jev") { const parsed = jevClassifierConfigSchema.safeParse(config.jev_classifier_config ?? {}); - return parsed.success ? null : "Enter a JEV model, a positive whole-number timeout and a positive cooldown"; + return parsed.success + ? null + : "Enter a valid classifier model, a positive whole-number timeout and a positive cooldown"; } if (!usesLlmClassifier(effectiveClassifierType(config)) || config.classifier_llm_config?.model) return null; return config.custom_tier_set @@ -498,7 +502,7 @@ export const customTierWireFields = ( tiers: Object.fromEntries(rows.map((row) => [activeTierName(row), row.models])), tier_definitions: tierDefinitionsFromRows(rows), ...(fallback && { fallback_tier: activeTierName(fallback) }), - classifier_type: classifierType === "jev" ? "jev" : "llm", + classifier_type: classifierType === "jev" ? "oss_classifier" : "llm", // Rebuilt from the fields an edited tier set allows. The backend rejects system_prompt and // classification_rubric beside tier_definitions, and both live inside this object rather than at // the top level the omit list covers. The opening instructions ride classification_prompt below. @@ -769,8 +773,8 @@ export const buildComplexityRouterConfig = ({ ...(defaultModel?.trim() && { default_model: defaultModel }), ...(planModeMinTier?.trim() && { plan_mode_min_tier: planModeMinTier }), ...(cleanedTierLabels && { tier_labels: cleanedTierLabels }), - classifier_type: classifierType, - ...(effectiveType === "jev" && { jev_classifier_config: normalizeJevClassifierConfig(jevClassifierConfig) }), + classifier_type: classifierType === "jev" ? "oss_classifier" : classifierType, + ...(effectiveType === "jev" && { opensource_classifier_config: normalizeJevClassifierConfig(jevClassifierConfig) }), ...(heuristicV2SuccessThreshold !== undefined && { heuristic_v2_success_threshold: heuristicV2SuccessThreshold, }), diff --git a/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts b/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts index 478c763351c..d1e1c96863c 100644 --- a/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts @@ -1,7 +1,11 @@ import { z } from "zod"; +import type { ClassifierType } from "./classifier_types"; + +export const LAYA_MODELS = ["english", "multilingual", "typed-decisions"] as const; const jevClassifierConfigFields = { - model: z.string().trim().min(1).default("jev-latest"), + provider: z.preprocess((value) => (value === "typesafe" ? "jev" : value), z.enum(["jev", "laya"]).optional()), + model: z.string().trim().min(1).optional(), timeout_ms: z.number().int().positive().default(3000), instructions: z .string() @@ -11,15 +15,39 @@ const jevClassifierConfigFields = { circuit_breaker_cooldown_seconds: z.number().finite().positive().optional(), }; -export const jevClassifierConfigSchema = z.object(jevClassifierConfigFields); +export const jevClassifierConfigSchema = z + .object(jevClassifierConfigFields) + .transform((config) => ({ + ...config, + model: config.model ?? (config.provider === "laya" ? "english" : "jev-latest"), + })) + .refine((config) => config.provider !== "laya" || LAYA_MODELS.some((model) => model === config.model), { + message: "Select a supported Laya model", + path: ["model"], + }); export type JevClassifierConfig = z.infer; -export const defaultJevClassifierConfig = (): JevClassifierConfig => jevClassifierConfigSchema.parse({}); +export const defaultJevClassifierConfig = (provider: "jev" | "laya" = "jev"): JevClassifierConfig => + jevClassifierConfigSchema.parse({ provider }); + +export const hydrateOssClassifier = (config: { + classifier_type?: ClassifierType | "oss_classifier"; + opensource_classifier_config?: unknown; + jev_classifier_config?: unknown; +}): { classifier_type: ClassifierType; jev_classifier_config?: JevClassifierConfig } => ({ + classifier_type: config.classifier_type === "oss_classifier" ? "jev" : config.classifier_type ?? "heuristic", + jev_classifier_config: + config.classifier_type === "oss_classifier" || config.classifier_type === "jev" + ? jevClassifierConfigSchema.safeParse(config.opensource_classifier_config ?? config.jev_classifier_config ?? {}) + .data ?? defaultJevClassifierConfig() + : undefined, +}); export const normalizeJevClassifierConfig = ( config: JevClassifierConfig = defaultJevClassifierConfig(), ): JevClassifierConfig => ({ + provider: config.provider ?? "jev", model: config.model.trim(), timeout_ms: config.timeout_ms, ...(config.instructions?.trim() && { instructions: config.instructions.trim() }), diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts index 4b8e298191d..91514c3a7d6 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts @@ -48,7 +48,7 @@ const hydratedState: KeywordMatchingState = { }; describe("buildUpdatedComplexityRouterConfig keyword matching", () => { - it.each([false, true])("omits masked JEV credentials from dashboard saves, edited: %s", (edited) => { + it.each([false, true])("omits masked Jev credentials from legacy/canonical saves, edited: %s", (edited) => { const stored = { classifier_type: "jev" as const, tiers: FORM_VALUE.tiers, @@ -60,7 +60,15 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { api_base: "https://jev.example.com", }, }; - const hydrated = hydrateComplexityRouterConfig(stored, undefined); + const source = edited + ? { + ...stored, + classifier_type: "oss_classifier" as const, + jev_classifier_config: undefined, + opensource_classifier_config: { ...stored.jev_classifier_config, provider: "typesafe" }, + } + : stored; + const hydrated = hydrateComplexityRouterConfig(source, undefined); expect(hydrated.jev_classifier_config).not.toHaveProperty("api_key"); expect(hydrated.jev_classifier_config).not.toHaveProperty("api_base"); const value = edited @@ -69,8 +77,9 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { jev_classifier_config: { model: "jev-updated", timeout_ms: 8100, instructions: "" }, } : hydrated; - const saved = buildUpdatedComplexityRouterConfig(stored, value); - expect(saved.jev_classifier_config).toEqual({ + const saved = buildUpdatedComplexityRouterConfig(source, value); + expect(saved.opensource_classifier_config).toEqual({ + provider: "jev", ...(edited ? { model: "jev-updated", timeout_ms: 8100 } : { model: "jev-configured", timeout_ms: 6100, instructions: "Existing instructions" }), @@ -78,7 +87,7 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { for (const classifierType of ["llm", "heuristic"] as const) { expect( buildUpdatedComplexityRouterConfig(saved, transitionClassifierType(value, classifierType)), - ).not.toHaveProperty("jev_classifier_config"); + ).not.toHaveProperty("opensource_classifier_config"); } }); @@ -94,19 +103,21 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { tiers: FORM_VALUE.tiers, }; const saved = buildUpdatedComplexityRouterConfig(stored, hydrateComplexityRouterConfig(stored, undefined)); - expect(saved.jev_classifier_config).toEqual({ + expect(saved.opensource_classifier_config).toEqual({ + provider: "jev", model: "jev-configured", timeout_ms: 6100, circuit_breaker_enabled: false, }); }); - it.each([false, true])("round trips JEV settings and preserves unmanaged fields, custom: %s", (custom) => { + it.each([false, true])("round trips Laya settings and preserves unmanaged fields, custom: %s", (custom) => { const stored = { ...(custom ? storedCustomConfig() : STORED), classifier_llm_config: { model: "stale-judge", timeout_ms: 3000 }, - classifier_type: "jev" as const, - jev_classifier_config: { - model: "jev-test", + classifier_type: "oss_classifier" as const, + opensource_classifier_config: { + provider: "laya" as const, + model: "english", timeout_ms: 4100, instructions: "Judge the request", circuit_breaker_enabled: false, @@ -121,12 +132,12 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { const hydrated = hydrateComplexityRouterConfig(stored, undefined); expect(effectiveClassifierType(hydrated)).toBe("jev"); expect(hydrated.classifier_llm_config).toBeUndefined(); - expect(hydrated.jev_classifier_config).toEqual(stored.jev_classifier_config); + expect(hydrated.jev_classifier_config).toEqual(stored.opensource_classifier_config); expect(hydrated.classifier_context_per_turn_chars).toBe(450); const saved = buildUpdatedComplexityRouterConfig(stored, hydrated); const expectedSavedConfig = { - classifier_type: "jev", - jev_classifier_config: stored.jev_classifier_config, + classifier_type: "oss_classifier", + opensource_classifier_config: stored.opensource_classifier_config, classifier_context_window_size: 7, classifier_context_budget_chars: 9000, classifier_context_per_turn_chars: 450, @@ -135,12 +146,13 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { }; expect(saved).toMatchObject(expectedSavedConfig); expect(saved).not.toHaveProperty("classifier_llm_config"); + expect(saved).not.toHaveProperty("jev_classifier_config"); const reloaded = hydrateComplexityRouterConfig(saved, undefined); expect(reloaded.jev_classifier_config).toEqual(hydrated.jev_classifier_config); expect(reloaded.classifier_context_per_turn_chars).toBe(450); expect(effectiveClassifierType(reloaded)).toBe("jev"); const llm = buildUpdatedComplexityRouterConfig(saved, transitionClassifierType(reloaded, "llm")); - expect(llm).not.toHaveProperty("jev_classifier_config"); + expect(llm).not.toHaveProperty("opensource_classifier_config"); }); it.each([0, 0.92, 1])("hydrates and saves a success threshold of %s without changing the artifact", (threshold) => { @@ -883,6 +895,7 @@ describe("managed keys survive an untouched open-and-save", () => { "fallback_tier", "hybrid_boundary_margin", "jev_classifier_config", + "opensource_classifier_config", "classifier_plugin_timeout_ms", ]); diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx index 2fd4fba2b31..e91936ebd36 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx @@ -101,6 +101,7 @@ export const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([ "llm_v2_config", "classifier_llm_config", "jev_classifier_config", + "opensource_classifier_config", "classifier_context_window_size", "classifier_context_budget_chars", "classifier_context_include_assistant_turns", diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts b/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts index 6dbd2b19b52..f357fa77cea 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts @@ -1,4 +1,4 @@ -import { defaultJevClassifierConfig, jevClassifierConfigSchema } from "../add_model/jev_classifier_config"; +import { hydrateOssClassifier } from "../add_model/jev_classifier_config"; import { capabilitySettingsSchema, fuseSettingsSchema } from "../add_model/forecast_classifier_config"; import type { StoredComplexityRouterConfig } from "../add_model/build_complexity_router_config"; import { @@ -54,6 +54,7 @@ export const hydrateComplexityRouterConfig = ( parsedConfig: StoredComplexityRouterConfig, complexityRouterDefaultModel: string | null | undefined, ): ComplexityRouterConfigValue => { + const classifier = hydrateOssClassifier(parsedConfig); const builtIn = hydrateBuiltInTiers(parsedConfig.tiers, parsedConfig.enable_non_reasoning_tier); const { tiers: hydratedTiers, enable_non_reasoning_tier } = builtIn; const custom_tier_set = hydrateCustomTierSet(parsedConfig); @@ -70,19 +71,14 @@ export const hydrateComplexityRouterConfig = ( default_model: hydratePinnedDefaultModel(parsedConfig.default_model, complexityRouterDefaultModel, activeTiers), plan_mode_min_tier: hydratePlanModeMinTier(parsedConfig.plan_mode_min_tier, custom_tier_set), tier_labels: hydrateTierLabels(parsedConfig.tier_labels), - classifier_type: parsedConfig.classifier_type || "heuristic", + ...classifier, heuristic_v2_success_threshold: typeof parsedConfig.heuristic_v2_success_threshold === "number" ? parsedConfig.heuristic_v2_success_threshold : undefined, capability_classifier_config: capabilitySettingsSchema.safeParse(parsedConfig.capability_classifier_config).data, llm_v2_config: fuseSettingsSchema.safeParse(parsedConfig.llm_v2_config).data, - classifier_llm_config: parsedConfig.classifier_type === "jev" ? undefined : parsedConfig.classifier_llm_config, - jev_classifier_config: - parsedConfig.classifier_type === "jev" - ? jevClassifierConfigSchema.safeParse(parsedConfig.jev_classifier_config ?? {}).data ?? - defaultJevClassifierConfig() - : undefined, + classifier_llm_config: classifier.classifier_type === "jev" ? undefined : parsedConfig.classifier_llm_config, classifier_context_window_size: typeof parsedConfig.classifier_context_window_size === "number" ? parsedConfig.classifier_context_window_size diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.tsx index 6a8cedd3746..b10aaff7cd7 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.tsx @@ -143,7 +143,7 @@ function describeCause(decision: RoutingDecision): string { case "llm_classifier": return classifierModel ? `LLM classifier (${classifierModel})` : "LLM classifier"; case "jev_classifier": - return "JEV classifier"; + return "OSS classifier"; case "literal_keyword_match": case "keyword": return matchedKeyword ? `Keyword match: "${matchedKeyword}"` : "Keyword match"; diff --git a/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts b/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts index c67befe8fa5..85dcd7df716 100644 --- a/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts +++ b/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts @@ -859,23 +859,28 @@ describe("autorouter_presets", () => { }); describe("buildPresetPrefill", () => { - it("preserves JEV settings and drops inactive classifier settings when prefilling", () => { + it("preserves Laya settings and drops inactive classifier settings when prefilling", () => { const config = { tiers: { SIMPLE: ["fast"], MEDIUM: [], COMPLEX: [], REASONING: [] }, - classifier_type: "jev" as const, + classifier_type: "oss_classifier" as const, classification_mode: "every_request" as const, session_affinity: false, deployment_affinity: true, modality_routing: false, modality_pin_override: false, - jev_classifier_config: { model: "jev-test", timeout_ms: 4000, circuit_breaker_enabled: false }, + opensource_classifier_config: { + provider: "laya" as const, + model: "english", + timeout_ms: 4000, + circuit_breaker_enabled: false, + }, classifier_llm_config: { model: "stale-judge", timeout_ms: 6000 }, classifier_context_window_size: 6, }; const prefill = buildPresetPrefill(config, groupsOnly(["fast"])); const expectedJevConfig = { classifier_type: "jev", - jev_classifier_config: config.jev_classifier_config, + jev_classifier_config: config.opensource_classifier_config, classifier_context_window_size: 6, classifier_llm_config: undefined, }; diff --git a/ui/litellm-dashboard/src/lib/autorouter_presets.ts b/ui/litellm-dashboard/src/lib/autorouter_presets.ts index bbbf151e49b..ac3894f9d8d 100644 --- a/ui/litellm-dashboard/src/lib/autorouter_presets.ts +++ b/ui/litellm-dashboard/src/lib/autorouter_presets.ts @@ -1,3 +1,4 @@ +import { hydrateOssClassifier } from "@/components/add_model/jev_classifier_config"; import { ComplexityRouterConfigPayload, hydrateTierLabels, @@ -302,6 +303,7 @@ export const buildPresetPrefill = ( config: ComplexityRouterConfigPayload, availability: ModelAvailability, ): PresetPrefill => { + const classifier = hydrateOssClassifier(config); const resolve = (model: string): string => resolveAvailableModel(model, availability) ?? model; const resolveTier = (models: string[]): string[] => models.map(resolve); // Params key on the model name the preset spells while every tier entry is rewritten to the @@ -334,11 +336,10 @@ export const buildPresetPrefill = ( }, tier_model_params: resolveParamKeys(hydrateTierModelParams(config.tiers, config.tier_model_configs)), tier_labels: hydrateTierLabels(config.tier_labels), - classifier_type: config.classifier_type, + ...classifier, heuristic_v2_success_threshold: config.heuristic_v2_success_threshold, - jev_classifier_config: config.classifier_type === "jev" ? config.jev_classifier_config : undefined, classifier_llm_config: - config.classifier_type !== "jev" && config.classifier_llm_config + classifier.classifier_type !== "jev" && config.classifier_llm_config ? { ...config.classifier_llm_config, model: resolve(config.classifier_llm_config.model) } : undefined, classifier_context_window_size: config.classifier_context_window_size, From f95446ea3933acb50ae2a35ecdaac9b6ef0ac9fc Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 12:36:46 -0700 Subject: [PATCH 021/139] fix(otel): tolerate non-dict callback_settings.otel and ignore bare EXCLUDED_SERVICES env (#44086) * test(otel): cover non-dict callback_settings.otel and bare EXCLUDED_SERVICES env Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(otel): tolerate non-dict callback_settings.otel and ignore bare EXCLUDED_SERVICES env Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(otel): simplify settings_customise_sources signature Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(otel): poll for present spans instead of waiting the full window Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(otel): audit null otel block and bare EXCLUDED_SERVICES across the excluded-services matrix Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(otel): assert per-trace datastore spans in the unconfigured burst cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(otel): keep pydantic-settings runtime options on the OTel v2 env source Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(otel): require post-auth datastore spans at the tenant on the cache-hit twin Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(otel): cover pydantic-settings runtime options in callback_settings.otel through the proxy Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/integrations/otel/logger.py | 3 +- litellm/integrations/otel/model/config.py | 33 +- .../test_otel_excluded_services.py | 204 ++++++++++++ .../test_otel_excluded_services_matrix.py | 299 +++++++++++++++++- ..._v2_config_baggage_parenting_guardrails.py | 56 ++++ .../otel/test_otel_v2_destinations.py | 12 + 6 files changed, 602 insertions(+), 5 deletions(-) diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index 55eb8e8fb71..e1983a44451 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -915,7 +915,8 @@ def _excluded_db_systems(logger: "OpenTelemetryV2") -> frozenset[str]: logger got published: with ``callbacks: [langfuse_otel, otel]`` the ``otel`` callback folds into the preset, whose config is env-only. """ - configured: Final = litellm.callback_settings.get("otel", {}).get("excluded_services") + otel_settings: Final = (litellm.callback_settings or {}).get("otel") + configured: Final = otel_settings.get("excluded_services") if isinstance(otel_settings, dict) else None if configured is None: return logger.config.excluded_services return excluded_db_systems_from(configured) diff --git a/litellm/integrations/otel/model/config.py b/litellm/integrations/otel/model/config.py index 9eb29157d6f..ddb8e127408 100644 --- a/litellm/integrations/otel/model/config.py +++ b/litellm/integrations/otel/model/config.py @@ -5,7 +5,8 @@ from functools import lru_cache from typing import Annotated, Any, Final from pydantic import AliasChoices, BaseModel, Field, TypeAdapter, ValidationError, field_validator, model_validator -from pydantic_settings import BaseSettings, NoDecode, SettingsConfigDict +from pydantic.fields import FieldInfo +from pydantic_settings import BaseSettings, NoDecode, PydanticBaseSettingsSource, SettingsConfigDict from litellm._logging import verbose_logger from litellm.integrations.otel.model.baggage import ( @@ -121,9 +122,37 @@ class ExporterSpec(BaseModel): ) +class _EnvWithoutBareExcludedServices(PydanticBaseSettingsSource): + def __init__(self, settings_cls: type[BaseSettings], env_settings: PydanticBaseSettingsSource) -> None: + super().__init__(settings_cls) + self._env_settings: Final = env_settings + + def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[object, str, bool]: + return self._env_settings.get_field_value(field, field_name) + + def __call__(self) -> dict[str, object]: + return {key: value for key, value in self._env_settings().items() if key != "excluded_services"} + + class OpenTelemetryV2Config(BaseSettings): model_config = SettingsConfigDict(populate_by_name=True, extra="ignore") + @classmethod + def settings_customise_sources( + cls, + settings_cls: type[BaseSettings], + init_settings: PydanticBaseSettingsSource, + env_settings: PydanticBaseSettingsSource, + dotenv_settings: PydanticBaseSettingsSource, + file_secret_settings: PydanticBaseSettingsSource, + ) -> tuple[PydanticBaseSettingsSource, ...]: + return ( + init_settings, + _EnvWithoutBareExcludedServices(settings_cls, env_settings), + dotenv_settings, + file_secret_settings, + ) + # ----- single-destination shorthand, read from standard OTEL_* envs ----- # exporter: str = Field( default="console", @@ -178,7 +207,7 @@ class OpenTelemetryV2Config(BaseSettings): ) excluded_services: Annotated[frozenset[str], NoDecode] = Field( default_factory=frozenset, - validation_alias=AliasChoices("excluded_services", "LITELLM_OTEL_EXCLUDED_SERVICES"), + validation_alias=AliasChoices("LITELLM_OTEL_EXCLUDED_SERVICES"), description=( "Datastore services whose spans are withheld from key/team ``callback_vars`` " "OTel destinations (the operator's own exporters still receive them). Accepted " diff --git a/tests/integration/observability/test_otel_excluded_services.py b/tests/integration/observability/test_otel_excluded_services.py index 48b20651b97..8768052b488 100644 --- a/tests/integration/observability/test_otel_excluded_services.py +++ b/tests/integration/observability/test_otel_excluded_services.py @@ -124,6 +124,20 @@ def _trace_spans(sink_url: str, trace_id: str, seconds: float = 30) -> tuple[Spa return group +def _trace_spans_when( + sink_url: str, + trace_id: str, + ready: Callable[[tuple[Span, ...]], bool], + seconds: float = 30, +) -> tuple[Span, ...]: + spans: Final = eventually( + lambda: spans_for_trace(recorded_spans(sink_url)[1], trace_id), + ready, + seconds=seconds, + ) + return spans + + def _await_db_span(sink_url: str, trace_id: str | None, needle: str, seconds: float = 40, since: int = 0) -> None: def seen() -> bool: _, spans = recorded_spans(sink_url, since) @@ -242,6 +256,196 @@ def test_without_excluded_services_the_tenant_still_gets_redis_and_postgres_span assert {"redis", "postgresql"} <= systems, f"datastore spans missing at tenant: {systems}" +@pytest.mark.parametrize("otel", [None, True, "on", "", []], ids=["null", "true", "on", "empty_string", "empty_list"]) +@pytest.mark.timeout(180) +def test_a_non_mapping_otel_block_still_publishes_the_tenant_fan_out( + gateway: Gateway, + audit_sinks: SpanSinks, + otel_audit_config: AuditConfigWriter, + langfuse_vars: dict[str, JsonValue], + tmp_path: Path, + otel: JsonValue, +) -> None: + def with_callback_settings(config: dict) -> None: + config["litellm_settings"]["callbacks"] = ["langfuse_otel"] + config["callback_settings"]["otel"] = otel + + config: Final = _config_with(tmp_path, otel_audit_config, extra=with_callback_settings) + overrides: Final = {"LITELLM_OTEL_V2": "1", **_operator_langfuse(audit_sinks)} + with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate: + traffic: Final = _drive(candidate, langfuse_vars) + tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic) + _await_db_span(audit_sinks.tenant, tenant_trace, "redis") + tenant_spans: Final = _trace_spans_when( + audit_sinks.tenant, + tenant_trace, + lambda spans: any(span["kind"] == 2 for span in spans) and "redis" in _db_systems(spans), + seconds=15, + ) + assert any(span["kind"] == 2 for span in tenant_spans), "tenant SERVER root span missing" + assert "redis" in _db_systems(tenant_spans), f"tenant redis span missing: {_db_systems(tenant_spans)}" + operator_trace: Final = _trace_id(audit_sinks.operator, traffic) + operator_spans: Final = _trace_spans_when( + audit_sinks.operator, + operator_trace, + lambda spans: any(span["kind"] == 2 for span in spans), + seconds=15, + ) + assert any(span["kind"] == 2 for span in operator_spans), "operator SERVER root span missing" + + +@pytest.mark.parametrize("name", ["EXCLUDED_SERVICES", "excluded_services"]) +@pytest.mark.timeout(180) +def test_a_bare_excluded_services_env_var_is_ignored( + gateway: Gateway, + audit_sinks: SpanSinks, + otel_audit_config: AuditConfigWriter, + langfuse_vars: dict[str, JsonValue], + tmp_path: Path, + name: str, +) -> None: + config: Final = _config_with(tmp_path, otel_audit_config) + overrides: Final = {"LITELLM_OTEL_V2": "1", name: "redis,postgres"} + with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate: + tenant_start, _ = recorded_spans(audit_sinks.tenant) + traffic: Final = _drive(candidate, langfuse_vars) + tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic) + _await_db_span(audit_sinks.tenant, tenant_trace, "redis") + tenant_spans: Final = _trace_spans_when( + audit_sinks.tenant, + tenant_trace, + lambda spans: "redis" in _db_systems(spans), + seconds=15, + ) + assert "redis" in _db_systems(tenant_spans), f"redis span missing at tenant: {_db_systems(tenant_spans)}" + _await_db_span(audit_sinks.tenant, None, "postgresql", since=tenant_start) + _, all_tenant = recorded_spans(audit_sinks.tenant, tenant_start) + systems: Final = _db_systems(all_tenant) + assert {"redis", "postgresql"} <= systems, f"datastore spans missing at tenant: {systems}" + + +@pytest.mark.timeout(180) +def test_the_documented_env_var_wins_over_a_bare_excluded_services( + gateway: Gateway, + audit_sinks: SpanSinks, + otel_audit_config: AuditConfigWriter, + langfuse_vars: dict[str, JsonValue], + tmp_path: Path, +) -> None: + config: Final = _config_with(tmp_path, otel_audit_config) + overrides: Final = { + "LITELLM_OTEL_V2": "1", + "LITELLM_OTEL_EXCLUDED_SERVICES": "redis", + "EXCLUDED_SERVICES": "postgres", + } + with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate: + tenant_start, _ = recorded_spans(audit_sinks.tenant) + traffic: Final = _drive(candidate, langfuse_vars) + tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic) + _trace_spans(audit_sinks.tenant, tenant_trace, seconds=15) + _await_db_span(audit_sinks.tenant, None, "postgresql", since=tenant_start) + _, tenant_spans = recorded_spans(audit_sinks.tenant, tenant_start) + systems: Final = _db_systems(tenant_spans) + assert "postgresql" in systems, f"postgresql spans missing at tenant: {systems}" + assert "redis" not in systems, f"redis spans reached tenant: {systems}" + + +@pytest.mark.parametrize( + ("env_name", "redis_reaches_tenant"), + [ + pytest.param("LITELLM_OTEL_EXCLUDED_SERVICES", False, id="exact-case"), + pytest.param("litellm_otel_excluded_services", True, id="wrong-case"), + ], +) +@pytest.mark.timeout(180) +def test_case_sensitive_otel_settings_read_only_the_exact_env_name( + gateway: Gateway, + audit_sinks: SpanSinks, + otel_audit_config: AuditConfigWriter, + langfuse_vars: Mapping[str, JsonValue], + tmp_path: Path, + env_name: str, + redis_reaches_tenant: bool, +) -> None: + config: Final = _config_with(tmp_path, otel_audit_config, otel={"_case_sensitive": True}) + overrides: Final = {"LITELLM_OTEL_V2": "1", env_name: "redis"} + with owned_proxy( + gateway, + tmp_path, + overrides, + config=config, + remove_environment=("LITELLM_OTEL_EXCLUDED_SERVICES", "litellm_otel_excluded_services"), + workers=2, + ) as candidate: + tenant_start, _ = recorded_spans(audit_sinks.tenant) + traffic: Final = _drive(candidate, langfuse_vars) + tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic) + if redis_reaches_tenant: + _await_db_span(audit_sinks.tenant, tenant_trace, "redis") + _trace_spans_when( + audit_sinks.tenant, + tenant_trace, + lambda spans: "redis" in _db_systems(spans), + seconds=15, + ) + else: + _trace_spans(audit_sinks.tenant, tenant_trace, seconds=15) + _await_db_span(audit_sinks.tenant, None, "postgresql", seconds=60, since=tenant_start) + _, tenant_spans = recorded_spans(audit_sinks.tenant, tenant_start) + systems: Final = _db_systems(tenant_spans) + assert "postgresql" in systems, f"postgresql spans missing at tenant: {systems}" + assert ("redis" in systems) is redis_reaches_tenant, ( + f"tenant redis presence={('redis' in systems)}; expected={redis_reaches_tenant}; systems={systems}" + ) + + +@pytest.mark.timeout(180) +def test_env_ignore_empty_keeps_the_default_service_name( + gateway: Gateway, + audit_sinks: SpanSinks, + otel_audit_config: AuditConfigWriter, + langfuse_vars: Mapping[str, JsonValue], + tmp_path: Path, +) -> None: + config: Final = _config_with(tmp_path, otel_audit_config, otel={"_env_ignore_empty": True}) + overrides: Final = {"LITELLM_OTEL_V2": "1", "OTEL_SERVICE_NAME": ""} + with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate: + traffic: Final = _drive(candidate, langfuse_vars) + operator_trace: Final = _trace_id(audit_sinks.operator, traffic) + operator_spans: Final = _trace_spans_when( + audit_sinks.operator, + operator_trace, + lambda spans: any(span["kind"] == 2 for span in spans), + seconds=15, + ) + service_names: Final = tuple(span["resource"].get("service.name") for span in operator_spans) + assert service_names and all(name == "litellm" for name in service_names), ( + f"operator service.name values={service_names}" + ) + + +@pytest.mark.timeout(180) +def test_env_parse_none_str_reads_a_null_traces_endpoint_as_unset( + gateway: Gateway, + audit_sinks: SpanSinks, + otel_audit_config: AuditConfigWriter, + langfuse_vars: Mapping[str, JsonValue], + tmp_path: Path, +) -> None: + config: Final = _config_with(tmp_path, otel_audit_config, otel={"_env_parse_none_str": "null"}) + overrides: Final = {"LITELLM_OTEL_V2": "1", "OTEL_TRACES_ENDPOINT": "null"} + with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate: + traffic: Final = _drive(candidate, langfuse_vars) + operator_trace: Final = _trace_id(audit_sinks.operator, traffic) + operator_spans: Final = _trace_spans_when( + audit_sinks.operator, + operator_trace, + lambda spans: any(span["kind"] == 2 for span in spans), + seconds=15, + ) + assert any(span["kind"] == 2 for span in operator_spans), "operator SERVER root span missing" + + def test_env_excluded_services_drops_only_redis( gateway: Gateway, audit_sinks: SpanSinks, diff --git a/tests/integration/observability/test_otel_excluded_services_matrix.py b/tests/integration/observability/test_otel_excluded_services_matrix.py index 0d4b5d087c6..7bca7ffc0dc 100644 --- a/tests/integration/observability/test_otel_excluded_services_matrix.py +++ b/tests/integration/observability/test_otel_excluded_services_matrix.py @@ -9,7 +9,8 @@ from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager from dataclasses import dataclass from pathlib import Path -from typing import Final, Literal +from types import MappingProxyType +from typing import Final, Literal, Protocol, cast import anthropic import httpx @@ -26,6 +27,9 @@ from pydantic import JsonValue, TypeAdapter MARKER: Final = re.compile(rb"excl-[0-9a-f]{32}") FAILING: Final = re.compile(rb"excl-fail-[0-9a-f]{32}") JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +UNCONFIGURED_VARIANT: Final[TypeAdapter[Literal["null_block", "bare_env"]]] = TypeAdapter( + Literal["null_block", "bare_env"] +) REPLY_TEXT: Final = "excluded ok" SERVER: Final = 2 INVALID_NAME_LOG: Final = "is not a datastore service" @@ -37,6 +41,11 @@ CLIENTS: Final[tuple[Client, ...]] = ("raw", "sdk", "async_sdk") AuditConfigWriter = Callable[[Path, Mapping[str, JsonValue]], Path] +class _FixtureRequestParam(Protocol): + @property + def param(self) -> object: ... + + def _marker() -> str: return "excl-" + uuid.uuid4().hex @@ -339,6 +348,11 @@ def _db_systems(spans: tuple[Span, ...]) -> set[str]: } +def _post_auth_datastore_spans(spans: tuple[Span, ...]) -> tuple[Span, ...]: + auth_ids: Final = frozenset(span["span_id"] for span in spans if span["name"].startswith("auth ")) + return tuple(span for span in spans if _db_systems((span,)) and span["parent_span_id"] not in auth_ids) + + def _names(spans: tuple[Span, ...]) -> list[str]: return sorted(span["name"] for span in spans) @@ -408,6 +422,29 @@ def _config(directory: Path, otel_audit_config: AuditConfigWriter, otel: Mapping return path +def _null_otel_config(directory: Path, otel_audit_config: AuditConfigWriter, name: str) -> Path: + written: Final = otel_audit_config(directory, {}) + loaded: Final = object_value(JSON.validate_python(yaml.safe_load(written.read_text()))) + config: Final = { + **loaded, + "litellm_settings": {**object_value(loaded["litellm_settings"]), "callbacks": ["langfuse_otel"]}, + "callback_settings": {**object_value(loaded["callback_settings"]), "otel": None}, + } + path: Final = directory / f"{name}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _operator_langfuse(sinks: SpanSinks) -> dict[str, str]: + return { + "LANGFUSE_HOST": sinks.operator, + "LANGFUSE_PUBLIC_KEY": "pk-lf-operator", + "LANGFUSE_SECRET_KEY": "sk-lf-operator", + "OTEL_EXPORTER": "http/json", + "OTEL_ENDPOINT": sinks.operator, + } + + @contextmanager def _started( provider: Wire, @@ -416,13 +453,14 @@ def _started( directory: Path, langfuse_vars: Mapping[str, JsonValue], workers: int, + environment: Mapping[str, str] = MappingProxyType({}), ) -> Generator[Rig]: with ( gateway_from_environment() as gateway, owned_proxy_process( gateway, directory, - {"LITELLM_OTEL_V2": "1", "OTEL_BSP_SCHEDULE_DELAY": "300"}, + {"LITELLM_OTEL_V2": "1", "OTEL_BSP_SCHEDULE_DELAY": "300", **environment}, config=config, remove_environment=("LITELLM_OTEL_EXCLUDED_SERVICES",), workers=workers, @@ -458,6 +496,32 @@ def rig( yield started +@pytest.fixture(scope="module", params=["null_block", "bare_env"], ids=["null_block", "bare_env"]) +def unconfigured_rig( + request: pytest.FixtureRequest, + provider: Wire, + audit_sinks: SpanSinks, + otel_audit_config: AuditConfigWriter, + langfuse_vars: dict[str, JsonValue], + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[Rig]: + parameter: Final = cast(_FixtureRequestParam, request).param + variant: Final = UNCONFIGURED_VARIANT.validate_python(parameter) + directory: Final = tmp_path_factory.mktemp(f"excluded-{variant}") + config: Final = ( + _null_otel_config(directory, otel_audit_config, variant) + if variant == "null_block" + else _config(directory, otel_audit_config, {}, variant) + ) + environment: Final = ( + _operator_langfuse(audit_sinks) if variant == "null_block" else {"EXCLUDED_SERVICES": "redis,postgres"} + ) + with _started( + provider, audit_sinks, config, directory, langfuse_vars, workers=2, environment=environment + ) as started: + yield started + + @pytest.mark.timeout(120) @pytest.mark.parametrize("stream", [False, True], ids=["unary", "stream"]) @pytest.mark.parametrize("client", CLIENTS) @@ -654,6 +718,237 @@ def test_killing_one_of_two_workers_mid_burst_keeps_the_filter_on_the_survivor(r _assert_withheld(rig, rig.raw("chat", _marker(), stream=False), after) +def _assert_tenant_kept( + rig: Rig, trace_id: str, cursors: Cursors, *, needs_model_span: bool = False +) -> tuple[Span, ...]: + def ready(spans: tuple[Span, ...]) -> bool: + return ( + sum(1 for span in spans if span["kind"] == SERVER) == 1 + and "redis" in _db_systems(spans) + and (not needs_model_span or any("gen_ai.operation.name" in span["attributes"] for span in spans)) + ) + + tenant: Final = eventually( + lambda: spans_for_trace(recorded_spans(rig.sinks.tenant, cursors.tenant)[1], trace_id), + ready, + seconds=40, + return_last_on_timeout=True, + ) + assert sum(1 for span in tenant if span["kind"] == SERVER) == 1, _names(tenant) + assert "redis" in _db_systems(tenant), f"redis spans missing at the tenant: {_names(tenant)}" + assert not needs_model_span or any("gen_ai.operation.name" in span["attributes"] for span in tenant), _names(tenant) + return tenant + + +def _assert_kept(rig: Rig, sent: Sent, cursors: Cursors) -> tuple[Span, ...]: + operator: Final = _operator_trace(rig, sent, cursors) + return _assert_tenant_kept(rig, operator[0]["trace_id"], cursors, needs_model_span=True) + + +@pytest.mark.timeout(120) +@pytest.mark.parametrize("stream", [False, True], ids=["unary", "stream"]) +@pytest.mark.parametrize("client", CLIENTS) +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_unconfigured_tenant_trace_keeps_datastore_spans( + unconfigured_rig: Rig, endpoint: Endpoint, client: Client, stream: bool +) -> None: + cursors: Final = unconfigured_rig.cursors() + marker: Final = _marker() + sent: Final = unconfigured_rig.send(endpoint, client, marker, stream) + assert sent.text == REPLY_TEXT, sent + assert unconfigured_rig.upstream_hits(marker) == 1 + _assert_kept(unconfigured_rig, sent, cursors) + + +@pytest.mark.timeout(120) +@pytest.mark.parametrize("endpoint", ["chat", "messages"]) +def test_unconfigured_cache_hit_twin_keeps_datastore_spans(unconfigured_rig: Rig, endpoint: Endpoint) -> None: + cursors: Final = unconfigured_rig.cursors() + marker: Final = _marker() + first_result: Final = _traced_raw(unconfigured_rig, endpoint, marker) + first: Final = first_result[1] + assert first.text == REPLY_TEXT, first + assert unconfigured_rig.upstream_hits(marker) == 1 + _assert_kept(unconfigured_rig, first, cursors) + hit_cursors: Final = unconfigured_rig.cursors() + + def read_hit() -> tuple[str, Sent, tuple[Span, ...]]: + trace_id, sent = _traced_raw(unconfigured_rig, endpoint, marker) + return trace_id, sent, _operator_trace_by_id(unconfigured_rig, trace_id, hit_cursors) + + trace_id, hit, operator = eventually( + read_hit, + lambda result: ( + unconfigured_rig.upstream_hits(marker) == 0 + and "redis" in _db_systems(_post_auth_datastore_spans(result[2])) + ), + seconds=60, + ) + assert hit.text == REPLY_TEXT, hit + post_auth_datastore: Final = _post_auth_datastore_spans(operator) + post_auth_span_ids: Final = frozenset(span["span_id"] for span in post_auth_datastore) + non_datastore_names: Final = frozenset(span["name"] for span in operator if not _db_systems((span,))) + tenant: Final = eventually( + lambda: spans_for_trace(recorded_spans(unconfigured_rig.sinks.tenant, hit_cursors.tenant)[1], trace_id), + lambda spans: ( + non_datastore_names <= frozenset(span["name"] for span in spans) + and post_auth_span_ids <= frozenset(span["span_id"] for span in spans) + ), + seconds=40, + return_last_on_timeout=True, + ) + tenant_span_ids: Final = frozenset(span["span_id"] for span in tenant) + missing_post_auth_names: Final = tuple( + span["name"] for span in post_auth_datastore if span["span_id"] not in tenant_span_ids + ) + assert sum(1 for span in tenant if span["kind"] == SERVER) == 1, _names(tenant) + assert "redis" in _db_systems(tenant), ( + f"operator datastore systems={sorted(_db_systems(post_auth_datastore))}; " + f"tenant datastore systems={sorted(_db_systems(tenant))}; tenant spans={_names(tenant)}" + ) + assert not missing_post_auth_names, ( + f"missing post-auth datastore span names={missing_post_auth_names}; " + f"operator={_names(post_auth_datastore)}; tenant={_names(tenant)}" + ) + + +@pytest.mark.timeout(120) +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_unconfigured_failed_upstream_keeps_datastore_spans(unconfigured_rig: Rig, endpoint: Endpoint) -> None: + cursors: Final = unconfigured_rig.cursors() + marker: Final = "excl-fail-" + uuid.uuid4().hex + trace_id: Final = uuid.uuid4().hex + path, body = _body(unconfigured_rig.model, endpoint, marker, stream=False) + failed: Final = unconfigured_rig.proxy.client.post( + path, + json=body, + headers={ + "Authorization": f"Bearer {unconfigured_rig.key}", + "traceparent": f"00-{trace_id}-{uuid.uuid4().hex[:16]}-01", + }, + ) + assert failed.status_code == 500, failed.text + assert unconfigured_rig.upstream_hits(marker) >= 1 + operator: Final = eventually( + lambda: spans_for_trace(recorded_spans(unconfigured_rig.sinks.operator, cursors.operator)[1], trace_id), + lambda spans: _has_root(spans) and "redis" in _db_systems(spans), + seconds=40, + ) + assert "redis" in _db_systems(operator), _names(operator) + _assert_tenant_kept(unconfigured_rig, trace_id, cursors) + + +@pytest.mark.timeout(120) +@pytest.mark.parametrize("status", [403, 404]) +def test_unconfigured_rejecting_tenant_destination_recovers(unconfigured_rig: Rig, status: int) -> None: + configure_sink(unconfigured_rig.sinks.tenant, status=status) + try: + cursors: Final = unconfigured_rig.cursors() + marker: Final = _marker() + sent: Final = unconfigured_rig.raw("chat", marker, stream=True) + assert sent.text == REPLY_TEXT, sent + assert unconfigured_rig.upstream_hits(marker) == 1 + _operator_trace(unconfigured_rig, sent, cursors) + finally: + configure_sink(unconfigured_rig.sinks.tenant, status=200) + after: Final = unconfigured_rig.cursors() + recovered: Final = unconfigured_rig.raw("responses", _marker(), stream=False) + assert recovered.text == REPLY_TEXT, recovered + _assert_kept(unconfigured_rig, recovered, after) + + +@pytest.mark.timeout(120) +def test_unconfigured_key_level_destination_keeps_datastore_spans( + unconfigured_rig: Rig, langfuse_vars: dict[str, JsonValue] +) -> None: + key: Final = unconfigured_rig.scenario.key( + metadata={ + "logging": [ + {"callback_name": "langfuse_otel", "callback_type": "success", "callback_vars": dict(langfuse_vars)} + ] + } + ) + cursors: Final = unconfigured_rig.cursors() + marker: Final = _marker() + sent: Final = unconfigured_rig.raw("chat", marker, stream=False, key=key) + assert sent.text == REPLY_TEXT, sent + assert unconfigured_rig.upstream_hits(marker) == 1 + _assert_kept(unconfigured_rig, sent, cursors) + + +def _assert_tenant_kept_the_burst(rig: Rig, cursors: Cursors, traces: set[str]) -> None: + def ready(spans: tuple[Span, ...]) -> bool: + def trace_kept(trace: str) -> bool: + trace_spans: Final = spans_for_trace(spans, trace) + return any(span["kind"] == SERVER for span in trace_spans) and "redis" in _db_systems(trace_spans) + + return all(trace_kept(trace) for trace in traces) + + tenant: Final = eventually( + lambda: recorded_spans(rig.sinks.tenant, cursors.tenant)[1], + ready, + seconds=90, + return_last_on_timeout=True, + ) + burst: Final = tuple(span for span in tenant if span["trace_id"] in traces) + missing_roots: Final = tuple( + trace for trace in traces if not any(span["kind"] == SERVER for span in spans_for_trace(burst, trace)) + ) + missing_redis: Final = tuple(trace for trace in traces if "redis" not in _db_systems(spans_for_trace(burst, trace))) + assert not missing_roots, f"SERVER root missing from tenant burst traces: {missing_roots}, {_names(burst)}" + assert not missing_redis, f"redis spans missing from tenant burst traces: {missing_redis}, {_names(burst)}" + + +@pytest.mark.timeout(300) +def test_unconfigured_tenant_outage_during_a_mixed_burst(unconfigured_rig: Rig) -> None: + cursors: Final = unconfigured_rig.cursors() + configure_sink(unconfigured_rig.sinks.tenant, status=503) + try: + results: Final = _burst(unconfigured_rig, 30) + finally: + configure_sink(unconfigured_rig.sinks.tenant, status=200) + served: Final = _served(results) + assert len(served) == 30, [result for result in results if isinstance(result, str)] + assert all(sent.text == REPLY_TEXT for sent in served), served + traces: Final = _assert_operator_exactly_once(unconfigured_rig, served, cursors) + _assert_tenant_kept_the_burst(unconfigured_rig, cursors, traces) + after: Final = unconfigured_rig.cursors() + _assert_kept(unconfigured_rig, unconfigured_rig.raw("messages", _marker(), stream=True), after) + + +@pytest.mark.timeout(300) +def test_unconfigured_killing_one_of_two_workers_keeps_the_fan_out(unconfigured_rig: Rig) -> None: + root: Final = psutil.Process(unconfigured_rig.owned.process.pid) + workers: Final = eventually( + lambda: tuple(child for child in root.children() if "resource_tracker" not in " ".join(child.cmdline())), + lambda found: len(found) == 2, + seconds=30, + ) + cursors: Final = unconfigured_rig.cursors() + + def one(index: int) -> Sent | str: + if index == 6: + os.kill(workers[0].pid, signal.SIGKILL) + try: + return unconfigured_rig.raw("chat", _marker(), stream=index % 2 == 0) + except (httpx.HTTPError, AssertionError) as error: + return repr(error) + + with ThreadPoolExecutor(max_workers=6) as pool: + results: Final = tuple(pool.map(one, range(18))) + assert unconfigured_rig.owned.process.poll() is None, "Proxy root exited after a worker was killed" + failures: Final = tuple(result for result in results if isinstance(result, str)) + assert all(failure.startswith(("ReadError(", "RemoteProtocolError(", "ConnectError(")) for failure in failures), ( + failures + ) + assert len(failures) <= 6, failures + settled: Final = tuple(result for index, result in enumerate(results) if index > 12 and isinstance(result, Sent)) + traces: Final = _assert_operator_exactly_once(unconfigured_rig, settled, cursors) + _assert_tenant_kept_the_burst(unconfigured_rig, cursors, traces) + after: Final = unconfigured_rig.cursors() + _assert_kept(unconfigured_rig, unconfigured_rig.raw("chat", _marker(), stream=False), after) + + @dataclass(frozen=True, slots=True) class Setting: otel: Mapping[str, JsonValue] diff --git a/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py b/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py index 86837f7f46c..20f8c89b5ff 100644 --- a/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py +++ b/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py @@ -12,6 +12,7 @@ import asyncio import logging +from typing import Final import pytest @@ -110,6 +111,61 @@ def test_excluded_services_from_env_csv(monkeypatch): assert OpenTelemetryV2Config().excluded_services == frozenset({"redis", "postgresql"}) +@pytest.mark.parametrize("name", ["EXCLUDED_SERVICES", "excluded_services", "Excluded_Services"]) +def test_a_bare_excluded_services_env_var_is_ignored(monkeypatch, name): + for env_name in ("LITELLM_OTEL_EXCLUDED_SERVICES", "EXCLUDED_SERVICES", "excluded_services", "Excluded_Services"): + monkeypatch.delenv(env_name, raising=False) + monkeypatch.setenv(name, "redis,postgres") + assert OpenTelemetryV2Config().excluded_services == frozenset() + + +@pytest.mark.parametrize( + ("set_env_name", "env_value", "case_sensitive", "env_ignore_empty", "env_parse_none_str"), + [ + pytest.param("otel_service_name", "lower", True, False, None, id="case-sensitive"), + pytest.param("OTEL_SERVICE_NAME", "", False, True, None, id="ignore-empty"), + pytest.param("OTEL_ENDPOINT", "null", False, False, "null", id="parse-none"), + pytest.param("excluded_services", "redis", True, False, None, id="bare-exclusion"), + ], +) +def test_env_source_preserves_runtime_options( + monkeypatch: pytest.MonkeyPatch, + set_env_name: str, + env_value: str, + case_sensitive: bool, + env_ignore_empty: bool, + env_parse_none_str: str | None, +) -> None: + for env_name in ( + "OTEL_SERVICE_NAME", + "otel_service_name", + "OTEL_ENDPOINT", + "OTEL_EXPORTER_OTLP_ENDPOINT", + "LITELLM_OTEL_EXCLUDED_SERVICES", + "EXCLUDED_SERVICES", + "excluded_services", + "Excluded_Services", + ): + monkeypatch.delenv(env_name, raising=False) + monkeypatch.setenv(set_env_name, env_value) + config: Final = OpenTelemetryV2Config( + _case_sensitive=case_sensitive, + _env_ignore_empty=env_ignore_empty, + _env_parse_none_str=env_parse_none_str, + ) + assert config.service_name == "litellm" + assert config.endpoint is None + assert config.excluded_services == frozenset() + + +def test_the_documented_env_var_wins_over_a_bare_excluded_services(monkeypatch): + for env_name in ("LITELLM_OTEL_EXCLUDED_SERVICES", "EXCLUDED_SERVICES", "excluded_services", "Excluded_Services"): + monkeypatch.delenv(env_name, raising=False) + monkeypatch.setenv("EXCLUDED_SERVICES", "postgres") + monkeypatch.setenv("LITELLM_OTEL_EXCLUDED_SERVICES", "redis") + assert OpenTelemetryV2Config().excluded_services == frozenset({"redis"}) + + def test_excluded_services_config_wins_over_env(monkeypatch): monkeypatch.setenv("LITELLM_OTEL_EXCLUDED_SERVICES", "redis") assert OpenTelemetryV2Config(excluded_services=["postgres"]).excluded_services == frozenset({"postgresql"}) diff --git a/tests/unit/integrations/otel/test_otel_v2_destinations.py b/tests/unit/integrations/otel/test_otel_v2_destinations.py index 5a7057203e4..cb986e61229 100644 --- a/tests/unit/integrations/otel/test_otel_v2_destinations.py +++ b/tests/unit/integrations/otel/test_otel_v2_destinations.py @@ -1116,6 +1116,18 @@ class TestProviderWiring: assert self._fan_out_of(preset)._excluded_db_systems == frozenset({"redis"}) + @pytest.mark.parametrize("otel", [None, True, "on", "", []], ids=["null", "true", "on", "empty_string", "empty_list"]) + def test_a_non_mapping_otel_block_falls_back_to_the_published_logger_config(self, monkeypatch, otel): + monkeypatch.setattr(litellm, "callback_settings", {"otel": otel}, raising=False) + preset = OpenTelemetryV2( + config=OpenTelemetryV2Config(exporters=[ExporterSpec(kind="in_memory")], excluded_services=["redis"]), + callback_name="langfuse_otel", + ) + + publish_global_otel_v2_provider([], lambda _p: None, registered=preset) + + assert self._fan_out_of(preset)._excluded_db_systems == frozenset({"redis"}) + def test_otel_after_a_preset_reuses_it_and_still_takes_callback_settings_exclusions(self, monkeypatch): """``callbacks: [langfuse_otel, otel]`` keeps one v2 logger, exactly as before ``excluded_services`` existed, and the exclusion still comes from From f932e292c523f2d238c4b7361c5168776dc0d042 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 19:38:44 +0000 Subject: [PATCH 022/139] fix(proxy): carry key, team and project tags into pass-through spend logs (#42662) * fix(proxy): carry key, team and project tags into pass-through spend logs Pass-through endpoints built their request metadata without the key, team and project controls that native routes apply, so spend rows for configured routes and provider pass-throughs like /anthropic dropped the key, team and project tags and the key and team spend_logs_metadata. The native team and project controls now live in a shared helper that both paths call, key spend_logs_metadata is copied instead of aliased from the cached key, client metadata cannot overwrite user_api_key_ fields, and header tags dedupe with the same merge used on native routes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): drop covers markers from pass-through tag tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): type the shared team and project control helper Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): validate pass-through endpoint list instead of suppressing pyright Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): cover body, streaming, hostile, forged and native cells for pass-through tags Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test: mock pass-through request url as httpx.URL after rebase on main * refactor(proxy): merge spend_logs_metadata sources without a stacked comprehension * refactor(proxy): name the team and request spend_logs_metadata merge --------- Co-authored-by: ryan Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/litellm_pre_call_utils.py | 83 ++-- .../batch_attribution.py | 4 +- .../pass_through_endpoints.py | 29 +- .../spend/test_passthrough_request_tags.py | 421 ++++++++++++++++++ .../test_pass_through_endpoints.py | 55 +++ 5 files changed, 551 insertions(+), 41 deletions(-) create mode 100644 tests/integration/spend/test_passthrough_request_tags.py diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index d6daf6ebe59..a809f53aa85 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1760,7 +1760,7 @@ class LiteLLMProxyRequestSetup: ): # don't override k-v pair sent by request (user request) data[_metadata_variable_name]["spend_logs_metadata"][key] = value else: - data[_metadata_variable_name]["spend_logs_metadata"] = key_metadata["spend_logs_metadata"] + data[_metadata_variable_name]["spend_logs_metadata"] = dict(key_metadata["spend_logs_metadata"]) ## KEY-LEVEL DISABLE FALLBACKS if "disable_fallbacks" in key_metadata and isinstance(key_metadata["disable_fallbacks"], bool): @@ -1777,6 +1777,53 @@ class LiteLLMProxyRequestSetup: ) return data + @staticmethod + def add_team_and_project_level_controls( + user_api_key_dict: UserAPIKeyAuth, metadata: dict[str, object] + ) -> dict[str, object]: + team_metadata: Final = user_api_key_dict.team_metadata or MappingProxyType({}) + project_metadata: Final = user_api_key_dict.project_metadata or MappingProxyType({}) + request_tags: Final = metadata.get("tags") + team_tags: Final = team_metadata.get("tags") + project_tags: Final = project_metadata.get("tags") + disable_global_guardrails: Final = team_metadata.get("disable_global_guardrails") + opted_out_global_guardrails: Final = team_metadata.get("opted_out_global_guardrails") + spend_logs_metadata: Final = LiteLLMProxyRequestSetup._merge_spend_logs_metadata( + team_spend_logs_metadata=team_metadata.get("spend_logs_metadata"), + request_spend_logs_metadata=metadata.get("spend_logs_metadata"), + ) + tags: Final = LiteLLMProxyRequestSetup._merge_tags( + request_tags=LiteLLMProxyRequestSetup._merge_tags( + request_tags=request_tags if isinstance(request_tags, list) else None, + tags_to_add=team_tags if isinstance(team_tags, list) else None, + ), + tags_to_add=project_tags if isinstance(project_tags, list) else None, + ) + controls: Final = ( + ("tags", tags or None), + ("spend_logs_metadata", spend_logs_metadata), + ( + "disable_global_guardrails", + disable_global_guardrails if isinstance(disable_global_guardrails, bool) else None, + ), + ( + "opted_out_global_guardrails", + opted_out_global_guardrails if isinstance(opted_out_global_guardrails, list) else None, + ), + ) + return {**metadata, **{key: value for key, value in controls if value is not None}} + + @staticmethod + def _merge_spend_logs_metadata( + team_spend_logs_metadata: object, request_spend_logs_metadata: object + ) -> dict[str, object] | None: + """Team values as defaults, the request's own values win on the same key. None when neither is a dict""" + team_values: Final = team_spend_logs_metadata if isinstance(team_spend_logs_metadata, dict) else None + request_values: Final = request_spend_logs_metadata if isinstance(request_spend_logs_metadata, dict) else None + if team_values is None and request_values is None: + return None + return {**(team_values or {}), **(request_values or {})} + @staticmethod def _merge_tags(request_tags: list | None, tags_to_add: list | None) -> list: """ @@ -2312,38 +2359,12 @@ async def add_litellm_data_to_request( data=data, _metadata_variable_name=_metadata_variable_name, ) - ## TEAM-LEVEL SPEND LOGS/TAGS + data[_metadata_variable_name] = LiteLLMProxyRequestSetup.add_team_and_project_level_controls( + user_api_key_dict=user_api_key_dict, + metadata=data[_metadata_variable_name], + ) team_metadata: Final = user_api_key_dict.team_metadata or {} - if "tags" in team_metadata and team_metadata["tags"] is not None: - data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags( - request_tags=data[_metadata_variable_name].get("tags"), - tags_to_add=team_metadata["tags"], - ) - if "disable_global_guardrails" in team_metadata and isinstance(team_metadata["disable_global_guardrails"], bool): - data[_metadata_variable_name]["disable_global_guardrails"] = team_metadata["disable_global_guardrails"] - if "opted_out_global_guardrails" in team_metadata and isinstance( - team_metadata["opted_out_global_guardrails"], list - ): - data[_metadata_variable_name]["opted_out_global_guardrails"] = team_metadata["opted_out_global_guardrails"] - if "spend_logs_metadata" in team_metadata and isinstance(team_metadata["spend_logs_metadata"], dict): - if "spend_logs_metadata" in data[_metadata_variable_name] and isinstance( - data[_metadata_variable_name]["spend_logs_metadata"], dict - ): - for key, value in team_metadata["spend_logs_metadata"].items(): - if ( - key not in data[_metadata_variable_name]["spend_logs_metadata"] - ): # don't override k-v pair sent by request (user request) - data[_metadata_variable_name]["spend_logs_metadata"][key] = value - else: - data[_metadata_variable_name]["spend_logs_metadata"] = team_metadata["spend_logs_metadata"] - - ## PROJECT-LEVEL TAGS project_metadata: Final = user_api_key_dict.project_metadata or {} - if "tags" in project_metadata and project_metadata["tags"] is not None: - data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags( - request_tags=data[_metadata_variable_name].get("tags"), - tags_to_add=project_metadata["tags"], - ) # inherited_tags: every tag key/team/project policy contributed, read # directly from those three sources rather than snapshotted off the shared diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/batch_attribution.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/batch_attribution.py index e7b608e162e..880fdad92bf 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/batch_attribution.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/batch_attribution.py @@ -34,9 +34,7 @@ def is_collection_route(url_route: str, collection_suffix: str) -> bool: def request_tags_from_metadata(request_metadata: Mapping[str, object]) -> tuple[str, ...] | None: """Tags for the batch-cost spend row: the request's own tags when it sent any, - otherwise the key's tags, which auth exposes as user_api_key_auth_metadata (a - tagged key does not put its tags in the top-level metadata "tags" on the - passthrough path) + otherwise the key's tags, which auth exposes as user_api_key_auth_metadata """ tags: Final = _sanitized_str_tuple(request_metadata.get("tags")) if tags: diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index fa592629933..865374a0430 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -619,11 +619,11 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): litellm_metadata: Final = litellm_keys_in_body.get("litellm_metadata") metadata: Final = litellm_keys_in_body.get("metadata") - if litellm_metadata: - _metadata.update(litellm_metadata) - if metadata: - _metadata.update(metadata) + for client_metadata in (litellm_metadata, metadata): + if isinstance(client_metadata, dict): + _metadata.update({k: v for k, v in client_metadata.items() if not k.startswith("user_api_key_")}) + _metadata = _apply_key_team_project_controls(user_api_key_dict=user_api_key_dict, metadata=_metadata) _metadata = _update_metadata_with_tags_in_header( request=request, metadata=_metadata, @@ -1934,6 +1934,20 @@ async def pass_through_request( ) +def _apply_key_team_project_controls( + user_api_key_dict: UserAPIKeyAuth, metadata: dict[str, object] +) -> dict[str, object]: + data: Final = LiteLLMProxyRequestSetup.add_key_level_controls( + key_metadata=user_api_key_dict.metadata, + data={"metadata": metadata}, + _metadata_variable_name="metadata", + ) + return LiteLLMProxyRequestSetup.add_team_and_project_level_controls( + user_api_key_dict=user_api_key_dict, + metadata=data["metadata"], + ) + + def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> dict: """ If tags are in the request headers, add them to the metadata @@ -1954,9 +1968,10 @@ def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> di # Only add tags key if there are tags to add if tags_to_add: - if "tags" not in metadata: - metadata["tags"] = [] - metadata["tags"].extend(tags_to_add) + metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags( + request_tags=metadata.get("tags"), + tags_to_add=tags_to_add, + ) return metadata diff --git a/tests/integration/spend/test_passthrough_request_tags.py b/tests/integration/spend/test_passthrough_request_tags.py new file mode 100644 index 00000000000..e6f5e15e161 --- /dev/null +++ b/tests/integration/spend/test_passthrough_request_tags.py @@ -0,0 +1,421 @@ +import json +import uuid +from collections.abc import Callable, Mapping +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import pytest +from integration._support.client import Gateway, JsonValue, Scenario, eventually, object_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server +from pydantic import TypeAdapter + + +def _chat_reply(marker: str) -> dict[str, JsonValue]: + return { + "id": f"chatcmpl-{marker}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": marker}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}, + } + + +def _anthropic_reply(marker: str) -> dict[str, JsonValue]: + return { + "id": f"msg_{marker}", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "content": [{"type": "text", "text": marker}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 5, "output_tokens": 3}, + } + + +def _chat_stream_frames(marker: str) -> tuple[bytes, ...]: + chunk: Final = {"id": f"chatcmpl-{marker}", "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"} + return ( + f"data: {json.dumps({**chunk, 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': marker}}]})}\n\n".encode(), + f"data: {json.dumps({**chunk, 'choices': [{'index': 0, 'delta': {}, 'finish_reason': 'stop'}], 'usage': {'prompt_tokens': 5, 'completion_tokens': 3, 'total_tokens': 8}})}\n\n".encode(), + b"data: [DONE]\n\n", + ) + + +def _spend_row(digest: str, call_type: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_tags, metadata, team_id FROM "LiteLLM_SpendLogs" WHERE api_key=%s AND call_type=%s', + (digest, call_type), + ), + lambda values: len(values) == 1, + seconds=70, + ) + return rows[0] + + +def _spend_row_tagged(tag: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_tags, metadata, team_id, api_key FROM "LiteLLM_SpendLogs" WHERE request_tags::text LIKE %s', + (f'%"{tag}"%',), + ), + lambda values: len(values) == 1, + seconds=70, + ) + return rows[0] + + +def _policy_tags(row: Mapping[str, JsonValue]) -> list[JsonValue]: + raw: Final = row["request_tags"] + tags: Final = json.loads(raw) if isinstance(raw, str) else raw + assert isinstance(tags, list), row + return [tag for tag in tags if not (isinstance(tag, str) and tag.startswith("User-Agent: "))] + + +def _spend_logs_metadata(row: Mapping[str, JsonValue]) -> JsonValue: + metadata: Final = row["metadata"] + return object_value(json.loads(metadata) if isinstance(metadata, str) else metadata).get("spend_logs_metadata") + + +def _tagged_key(scenario: Scenario, marker: str, **fields: JsonValue) -> tuple[str, str]: + team: Final = scenario.team(metadata={"tags": [f"team-{marker}"], "spend_logs_metadata": {"team_field": marker}}) + project: Final = scenario.project(team, metadata={"tags": [f"project-{marker}"]}) + key: Final = scenario.key( + team_id=team, + project_id=project, + metadata={"tags": [f"key-{marker}"], "spend_logs_metadata": {"cost_center": marker}}, + **fields, + ) + return key, sha256(key.encode()).hexdigest() + + +def _digest(key: str) -> str: + return sha256(key.encode()).hexdigest() + + +def _configured_passthrough(gateway: Gateway, scenario: Scenario, marker: str, target: str, *, auth: bool) -> str: + path: Final = f"/integration-passthrough-{marker}" + created: Final = gateway.post("/config/pass_through_endpoint", {"path": path, "target": target, "auth": auth}) + endpoints: Final = TypeAdapter(list[JsonValue]).validate_python(created["endpoints"]) + endpoint_id: Final = object_value(endpoints[0])["id"] + scenario.cleanups.callback( + lambda: gateway.request("DELETE", "/config/pass_through_endpoint", params={"endpoint_id": str(endpoint_id)}) + ) + return path + + +def _responses_reply(marker: str, stream: bool) -> Reply: + response: Final[dict[str, JsonValue]] = { + "id": f"resp_{marker}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "id": f"msg_{marker}", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": marker, "annotations": []}], + } + ], + "usage": {"input_tokens": 5, "output_tokens": 3, "total_tokens": 8}, + } + if not stream: + return Reply(body=json.dumps(response).encode()) + events: Final[tuple[dict[str, JsonValue], ...]] = ( + { + "type": "response.created", + "sequence_number": 0, + "response": {**response, "status": "in_progress", "output": []}, + }, + { + "type": "response.output_text.delta", + "sequence_number": 1, + "item_id": f"msg_{marker}", + "output_index": 0, + "content_index": 0, + "delta": marker, + }, + {"type": "response.completed", "sequence_number": 2, "response": response}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _echo_upstream(marker: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + if request.method == "GET" and request.target == "/v1/models": + return Reply(body=json.dumps({"object": "list", "data": []}).encode()) + assert request.method == "POST", request + body: Final = object_value(json.loads(request.body)) + assert marker in json.dumps(body), request + if request.target == "/v1/responses": + return _responses_reply(marker, body.get("stream") is True) + assert body["messages"] == [{"role": "user", "content": marker}], request + if body.get("stream") is True: + return Reply(chunks=_chat_stream_frames(marker), content_type="text/event-stream") + return Reply(body=json.dumps(_chat_reply(marker)).encode()) + + return respond + + +def test_configured_passthrough_spend_row_matches_native_route_tags_and_spend_logs_metadata(gateway: Gateway) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(api_base=wire.url + "/v1") + path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True) + key, digest = _tagged_key(scenario, marker, models=[model], allowed_passthrough_routes=[path]) + headers: Final = {"x-litellm-tags": f"caller-{marker},key-{marker}", "User-Agent": "integration-tags/1"} + body: Final[dict[str, JsonValue]] = {"model": model, "messages": [{"role": "user", "content": marker}]} + + native: Final = gateway.request("POST", "/v1/chat/completions", body, key=key, headers=headers) + assert native.status_code == 200, native.text + passthrough: Final = gateway.request("POST", path, body, key=key, headers=headers) + assert passthrough.status_code == 200, passthrough.text + assert json.loads(passthrough.content) == _chat_reply(marker) + + native_row: Final = _spend_row(digest, "acompletion") + passthrough_row: Final = _spend_row(digest, "pass_through_endpoint") + expected: Final = [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"] + assert _policy_tags(native_row) == expected, native_row + assert _policy_tags(passthrough_row) == expected, passthrough_row + assert _spend_logs_metadata(native_row) == {"cost_center": marker, "team_field": marker}, native_row + assert _spend_logs_metadata(passthrough_row) == {"cost_center": marker, "team_field": marker}, passthrough_row + + +@pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata"]) +def test_configured_passthrough_body_tags_lead_and_body_spend_logs_metadata_wins_over_key_and_team( + gateway: Gateway, bucket: str +) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario: + path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True) + key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path]) + body: Final[dict[str, JsonValue]] = { + "messages": [{"role": "user", "content": marker}], + bucket: { + "tags": [f"body-{marker}", f"team-{marker}"], + "spend_logs_metadata": {"cost_center": f"body-{marker}"}, + }, + } + response: Final = gateway.request("POST", path, body, key=key) + assert response.status_code == 200, response.text + row: Final = _spend_row(digest, "pass_through_endpoint") + assert _policy_tags(row) == [f"body-{marker}", f"team-{marker}", f"key-{marker}", f"project-{marker}"], row + assert _spend_logs_metadata(row) == {"cost_center": f"body-{marker}", "team_field": marker}, row + + +def test_configured_passthrough_streaming_upstream_row_carries_key_team_project_and_caller_tags( + gateway: Gateway, +) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario: + path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True) + key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path]) + response: Final = gateway.request( + "POST", + path, + {"stream": True, "messages": [{"role": "user", "content": marker}]}, + key=key, + headers={"x-litellm-tags": f"caller-{marker}"}, + ) + assert response.status_code == 200, response.text + assert response.content == b"".join(_chat_stream_frames(marker)), response.text + row: Final = _spend_row(digest, "pass_through_endpoint") + assert _policy_tags(row) == [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"], row + assert _spend_logs_metadata(row) == {"cost_center": marker, "team_field": marker}, row + + +def test_configured_passthrough_key_outside_any_team_carries_its_own_tags_and_spend_logs_metadata( + gateway: Gateway, +) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario: + path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True) + key: Final = scenario.key( + allowed_passthrough_routes=[path], + metadata={"tags": [f"key-{marker}"], "spend_logs_metadata": {"cost_center": marker}}, + ) + response: Final = gateway.request( + "POST", + path, + {"messages": [{"role": "user", "content": marker}]}, + key=key, + headers={"x-litellm-tags": f"caller-{marker}"}, + ) + assert response.status_code == 200, response.text + row: Final = _spend_row(_digest(key), "pass_through_endpoint") + assert _policy_tags(row) == [f"key-{marker}", f"caller-{marker}"], row + assert _spend_logs_metadata(row) == {"cost_center": marker}, row + + +def test_configured_passthrough_untagged_key_row_keeps_only_caller_tag_and_no_spend_logs_metadata( + gateway: Gateway, +) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario: + path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True) + team: Final = scenario.team() + key: Final = scenario.key(team_id=team, allowed_passthrough_routes=[path]) + response: Final = gateway.request( + "POST", + path, + {"messages": [{"role": "user", "content": marker}]}, + key=key, + headers={"x-litellm-tags": f"caller-{marker}"}, + ) + assert response.status_code == 200, response.text + row: Final = _spend_row(_digest(key), "pass_through_endpoint") + assert _policy_tags(row) == [f"caller-{marker}"], row + assert _spend_logs_metadata(row) is None, row + assert row["team_id"] == team, row + + +def test_open_passthrough_without_auth_row_carries_only_caller_tag(gateway: Gateway) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario: + path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=False) + response: Final = gateway.client.post( + path, + json={"messages": [{"role": "user", "content": marker}]}, + headers={"x-litellm-tags": f"caller-{marker}"}, + ) + assert response.status_code == 200, response.text + row: Final = _spend_row_tagged(f"caller-{marker}") + assert _policy_tags(row) == [f"caller-{marker}"], row + assert _spend_logs_metadata(row) is None, row + assert row["api_key"] == "", row + + +@pytest.mark.parametrize( + ("metadata", "leading_tags"), + [ + ({"tags": "string-not-list"}, []), + ({"tags": [1, None, "z"]}, [1, None, "z"]), + ({"spend_logs_metadata": "string-not-object"}, []), + ], +) +def test_configured_passthrough_hostile_body_metadata_shapes_still_carry_key_team_project_tags( + gateway: Gateway, metadata: JsonValue, leading_tags: list[JsonValue] +) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario: + path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True) + key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path]) + response: Final = gateway.request( + "POST", path, {"messages": [{"role": "user", "content": marker}], "metadata": metadata}, key=key + ) + assert response.status_code == 200, response.text + row: Final = _spend_row(digest, "pass_through_endpoint") + assert _policy_tags(row) == [*leading_tags, f"key-{marker}", f"team-{marker}", f"project-{marker}"], row + assert _spend_logs_metadata(row) == {"cost_center": marker, "team_field": marker}, row + + +def test_configured_passthrough_body_cannot_forge_user_api_key_attribution_fields(gateway: Gateway) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario: + path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True) + forged_team: Final = scenario.team() + key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path]) + forged: Final[dict[str, JsonValue]] = { + "user_api_key": "forged-" + marker, + "user_api_key_team_id": forged_team, + "user_api_key_user_id": "forged-" + marker, + "user_api_key_alias": "forged-" + marker, + } + body: Final[dict[str, JsonValue]] = {"messages": [{"role": "user", "content": marker}], "metadata": forged} + response: Final = gateway.request("POST", path, body, key=key) + assert response.status_code == 200, response.text + row: Final = _spend_row(digest, "pass_through_endpoint") + assert row["team_id"] != forged_team, row + assert _policy_tags(row) == [f"key-{marker}", f"team-{marker}", f"project-{marker}"], row + assert read_rows('SELECT api_key FROM "LiteLLM_SpendLogs" WHERE team_id=%s', (forged_team,)) == [], forged_team + + +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("route", ["/v1/chat/completions", "/v1/messages"]) +def test_native_routes_carry_key_team_project_and_caller_tags_and_key_over_team_spend_logs_metadata( + gateway: Gateway, route: str, stream: bool +) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(api_base=wire.url + "/v1") + key, digest = _tagged_key(scenario, marker, models=[model]) + body: Final[dict[str, JsonValue]] = { + "model": model, + "max_tokens": 16, + "stream": stream, + "messages": [{"role": "user", "content": marker}], + } + response: Final = gateway.request( + "POST", + route, + body, + key=key, + headers={"x-litellm-tags": f"caller-{marker}", "User-Agent": "integration-tags/1"}, + ) + assert response.status_code == 200, response.text + rows: Final = eventually( + lambda: read_rows('SELECT request_tags, metadata FROM "LiteLLM_SpendLogs" WHERE api_key=%s', (digest,)), + lambda values: len(values) == 1, + seconds=70, + ) + assert _policy_tags(rows[0]) == [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"], ( + rows + ) + assert _spend_logs_metadata(rows[0]) == {"cost_center": marker, "team_field": marker}, rows + + +def test_anthropic_passthrough_spend_row_carries_key_team_project_tags_and_spend_logs_metadata( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages", request + assert request.headers["x-api-key"] == "synthetic-anthropic-key" + return Reply(body=json.dumps(_anthropic_reply(marker)).encode()) + + config: Final = tmp_path / "proxy_config.yaml" + config.write_text( + "model_list: []\n" + "general_settings:\n" + " master_key: os.environ/LITELLM_MASTER_KEY\n" + " database_url: os.environ/DATABASE_URL\n" + " store_model_in_db: true\n" + " disable_spend_logs: false\n" + " proxy_batch_write_at: 1\n" + "router_settings:\n" + " disable_cooldowns: true\n" + ) + with wire_server(respond) as wire: + overrides: Final = {"ANTHROPIC_API_BASE": wire.url, "ANTHROPIC_API_KEY": "synthetic-anthropic-key"} + with owned_proxy(gateway, tmp_path, overrides, config=config) as candidate, candidate.scenario() as scenario: + key, digest = _tagged_key(scenario, marker) + response: Final = candidate.request( + "POST", + "/anthropic/v1/messages", + { + "model": "claude-sonnet-4-5", + "max_tokens": 16, + "messages": [{"role": "user", "content": marker}], + }, + key=key, + headers={"x-litellm-tags": f"caller-{marker}", "User-Agent": "integration-tags/1"}, + ) + assert response.status_code == 200, response.text + assert json.loads(response.content) == _anthropic_reply(marker) + row: Final = _spend_row(digest, "pass_through_endpoint") + assert _policy_tags(row) == [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"], ( + row + ) + assert _spend_logs_metadata(row) == {"cost_center": marker, "team_field": marker}, row diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 72d76e54a5d..c6c81c14b16 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -7592,6 +7592,61 @@ def test_passthrough_sees_the_public_list_rebound_after_import(monkeypatch: pyte assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]} +def test_passthrough_metadata_carries_key_team_project_tags_and_key_spend_logs_metadata(): + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = httpx.URL("http://0.0.0.0:4000/anthropic/v1/messages") + mock_request.headers = Headers({"x-litellm-tags": "caller-tag,key-tag"}) + mock_request.scope = {} + + cached_key = UserAPIKeyAuth( + api_key="hashed-key", + metadata={"tags": ["key-tag", "shared-tag"], "spend_logs_metadata": {"cost_center": "key"}}, + team_metadata={ + "tags": ["team-tag", "shared-tag"], + "spend_logs_metadata": {"cost_center": "team", "team_field": "team"}, + }, + project_metadata={"tags": ["project-tag"]}, + ) + + kwargs = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=mock_request, + user_api_key_dict=cached_key, + passthrough_logging_payload=MagicMock(), + logging_obj=MagicMock(), + _parsed_body={ + "metadata": { + "tags": ["body-tag"], + "spend_logs_metadata": {"request_id": "body"}, + "user_api_key_auth_metadata": "forged", + } + }, + litellm_call_id="lit-5359-call-id", + ) + second = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=mock_request, + user_api_key_dict=cached_key, + passthrough_logging_payload=MagicMock(), + logging_obj=MagicMock(), + _parsed_body={}, + litellm_call_id="lit-5359-second-call-id", + ) + + metadata = kwargs["litellm_params"]["metadata"] + assert metadata["tags"] == ["body-tag", "key-tag", "shared-tag", "team-tag", "project-tag", "caller-tag"] + assert metadata["spend_logs_metadata"] == {"request_id": "body", "cost_center": "key", "team_field": "team"} + assert metadata["user_api_key_auth_metadata"] == { + "tags": ["key-tag", "shared-tag"], + "spend_logs_metadata": {"cost_center": "key"}, + } + assert second["litellm_params"]["metadata"]["spend_logs_metadata"] == {"cost_center": "key", "team_field": "team"} + assert cached_key.metadata == {"tags": ["key-tag", "shared-tag"], "spend_logs_metadata": {"cost_center": "key"}} + assert cached_key.team_metadata == { + "tags": ["team-tag", "shared-tag"], + "spend_logs_metadata": {"cost_center": "team", "team_field": "team"}, + } + + @pytest.mark.asyncio async def test_chat_completion_pass_through_endpoint_answers_an_openai_typed_error_for_an_unknown_model( monkeypatch: pytest.MonkeyPatch, From c8cd8852513313e3f08eb713d3a9ebddc55314b3 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 20:10:50 +0000 Subject: [PATCH 023/139] feat(ui): add System One (Jev) tab to the playground (#44043) * feat(ui): add System One (Jev) playground tab Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): validate System One inputs and refresh request context Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): validate System One response payloads Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): fall back to requested model for System One Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(ui): use useMutation and zod schemas for the System One tab, mark it Beta Replace the hand-written request and response guards with zod schemas, which also provide the types. Send requests through useMutation instead of manual loading, error and race-guard state. Unknown spec fields now pass through to the upstream model, and validation errors point at the exact offending key * feat(ui): flag the System One tab as a TypeSafe-only beta * feat(ui): highlighted JSON editor for the System One tab Line numbers, JSON syntax highlighting through the existing react-syntax-highlighter dependency, a valid or issue-count status badge, and a compact path plus message issue list replace the bare textarea and stacked alerts. Answer card type badges now sit on the header row * fix(ui): let the System One results scroll to the bottom The tab panel was viewport height but sat below the tab bar, so its bottom was cut off. On wide screens the editor and results now scroll independently, answers render above the question breakdown, and the raw response no longer nests its own scroll area * feat(ui): color the model, state and questions blocks in the System One editor Tints each top-level request block in the JSON editor and marks the matching breakdown sections with the same color, so it is clear which part of the payload feeds which panel * feat(ui): wrap long lines in the System One JSON editor Long state strings no longer need horizontal scrolling. Each line renders as its own row with its number and block color, so wrapped lines keep their line number and the caret stays aligned * fix(ui): remove horizontal scrolling from the System One tab Long unbroken text in the state, question ids, choice labels and the raw response now wraps instead of widening its box * refactor(ui): replace System One presets with one example and a reset button The tab now starts with a single product review example that uses all three question types, and Reset example restores it after editing * refactor(ui): use an issue triage request as the System One example * fix(ui): type the System One line renderer from exported props rendererProps is not exported by the react-syntax-highlighter types, which broke the dashboard build * fix(ui): address System One review feedback Highlights the score level nearest a fractional calibrated score, keeps extra noul criteria fields in the sent payload, and clears an answer when the request key changes * refactor(ui): parse System One root blocks without mutation and preview all noul criteria The root-block finder is now a tokenizer plus a pure reduce, and the question preview lists every noul criterion that will be sent * refactor(ui): group System One playground files into components and lib Drops the repeated SystemOne prefix from the inner component files and moves the pure logic (schemas, example, payload validation, root block parsing) into lib/. Tests stay colocated with their files, matching the rest of the dashboard. No behavior change * feat(ui): link the decision models discussion from the System One beta notice * feat(ui): ask for decision model feedback in the System One beta notice * feat(ui): make the decision model feedback text the discussion link --------- Co-authored-by: ryan Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../systemOneUI/JsonEditor.test.tsx | 70 +++++ .../components/systemOneUI/JsonEditor.tsx | 167 ++++++++++++ .../systemOneUI/QuestionBreakdown.test.tsx | 33 +++ .../systemOneUI/QuestionBreakdown.tsx | 108 ++++++++ .../systemOneUI/ResponseView.test.tsx | 67 +++++ .../components/systemOneUI/ResponseView.tsx | 196 ++++++++++++++ .../SystemOneUI.integration.test.tsx | 248 +++++++++++++++++ .../components/systemOneUI/SystemOneUI.tsx | 181 +++++++++++++ .../components/systemOneUI/lib/example.ts | 38 +++ .../systemOneUI/lib/rootBlocks.test.ts | 53 ++++ .../components/systemOneUI/lib/rootBlocks.ts | 65 +++++ .../components/systemOneUI/lib/schemas.ts | 120 +++++++++ .../systemOneUI/lib/validatePayload.test.ts | 250 ++++++++++++++++++ .../systemOneUI/lib/validatePayload.ts | 62 +++++ .../playground/llm_calls/system_one.test.ts | 107 ++++++++ .../playground/llm_calls/system_one.ts | 38 +++ .../app/(dashboard)/playground/page.test.tsx | 14 + .../src/app/(dashboard)/playground/page.tsx | 10 +- 18 files changed, 1826 insertions(+), 1 deletion(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/JsonEditor.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/JsonEditor.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/QuestionBreakdown.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/QuestionBreakdown.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/ResponseView.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/ResponseView.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/SystemOneUI.integration.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/SystemOneUI.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/example.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/rootBlocks.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/rootBlocks.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/schemas.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/validatePayload.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/validatePayload.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/system_one.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/system_one.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/JsonEditor.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/JsonEditor.test.tsx new file mode 100644 index 00000000000..c861c1d466d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/JsonEditor.test.tsx @@ -0,0 +1,70 @@ +import { fireEvent, render, screen } from "@testing-library/react"; +import { describe, expect, it, vi } from "vitest"; +import JsonEditor from "./JsonEditor"; +import { validateSystemOnePayload } from "./lib/validatePayload"; + +const validPayload = JSON.stringify( + { state: "Hi", questions: { escalate: { type: "noul", instructions: "Escalate?" } } }, + null, + 2, +); + +describe("JsonEditor", () => { + it("marks a valid payload as ready to send and counts its lines", () => { + render(); + + expect(screen.getByText("Valid payload")).toBeInTheDocument(); + expect(screen.getByRole("status")).toHaveTextContent("Ready to send"); + expect(screen.getByText(`${validPayload.split("\n").length} lines`)).toBeInTheDocument(); + expect(screen.getByRole("textbox", { name: "System One JSON payload" })).toHaveAttribute("aria-invalid", "false"); + }); + + it("lists each issue with its path and counts only errors in the status badge", () => { + const payload = JSON.stringify({ + state: "Hi", + questions: { + category: { type: "choice", instructions: 1, criteria: { support: "Help" } }, + urgency: { type: "score", instructions: "Rate", criteria: ["Low"] }, + }, + }); + render(); + + expect(screen.getByText("2 issues")).toBeInTheDocument(); + const issues = screen.getByRole("list", { name: "Payload validation issues" }); + expect(issues).toHaveTextContent("questions.category.instructionsInstructions must be a string."); + expect(issues).toHaveTextContent("questions.urgency.criteriaScore criteria must contain at least 2 levels."); + expect(screen.getByRole("textbox", { name: "System One JSON payload" })).toHaveAttribute("aria-invalid", "true"); + }); + + it("shows warnings without counting them as issues", () => { + const payload = JSON.stringify({ + state: "Hi", + questions: { urgency: { type: "score", instructions: "Rate", criteria: Array.from({ length: 11 }, () => "L") } }, + }); + render(); + + expect(screen.getByText("Valid payload")).toBeInTheDocument(); + expect(screen.getByRole("list", { name: "Payload validation issues" })).toHaveTextContent( + "More than 10 score levels may reduce result quality.", + ); + }); + + it("numbers every line, including a trailing empty one, so wrapped lines keep their number", () => { + const value = `${validPayload}\n`; + render(); + + const lineCount = value.split("\n").length; + expect(screen.getByText(`${lineCount} lines`)).toBeInTheDocument(); + expect(screen.getByText(String(lineCount))).toBeInTheDocument(); + expect(screen.queryByText(String(lineCount + 1))).not.toBeInTheDocument(); + }); + + it("reports edits to the caller", () => { + const onChange = vi.fn(); + render(); + + fireEvent.change(screen.getByRole("textbox", { name: "System One JSON payload" }), { target: { value: "{}" } }); + + expect(onChange).toHaveBeenCalledWith("{}"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/JsonEditor.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/JsonEditor.tsx new file mode 100644 index 00000000000..1effe55ef41 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/JsonEditor.tsx @@ -0,0 +1,167 @@ +import { Badge } from "@/components/ui/badge"; +import { cn } from "@/lib/cva.config"; +import { CircleAlert, CircleCheck, TriangleAlert } from "lucide-react"; +import { useId, useMemo, useRef } from "react"; +import { createElement, PrismLight as SyntaxHighlighter } from "react-syntax-highlighter"; +import type { SyntaxHighlighterProps } from "react-syntax-highlighter"; +import json from "react-syntax-highlighter/dist/esm/languages/prism/json"; +import { findRootBlocks, ROOT_BLOCK_STYLES, type RootBlock } from "./lib/rootBlocks"; +import type { SystemOnePayloadValidation } from "./lib/validatePayload"; + +SyntaxHighlighter.registerLanguage("json", json); + +const EDITOR_TEXT = "m-0 whitespace-pre-wrap wrap-anywhere py-3 font-mono text-xs leading-5 [scrollbar-gutter:stable]"; +const GUTTER_WIDTH = "w-11"; +type LineRendererProps = Parameters>[0]; +const CONTENT_INSET = "pl-14 pr-3"; +const CODE_TAG_PROPS = { className: "language-json", style: { whiteSpace: "pre-wrap" } } as const; + +const TOKEN_COLORS = [ + "[&_.token.property]:text-sky-700 dark:[&_.token.property]:text-sky-300", + "[&_.token.string]:text-emerald-700 dark:[&_.token.string]:text-emerald-300", + "[&_.token.number]:text-amber-700 dark:[&_.token.number]:text-amber-300", + "[&_.token.boolean]:text-violet-700 dark:[&_.token.boolean]:text-violet-300", + "[&_.token.null]:text-violet-700 dark:[&_.token.null]:text-violet-300", + "[&_.token.punctuation]:text-muted-foreground [&_.token.operator]:text-muted-foreground", +].join(" "); + +interface JsonEditorProps { + value: string; + onChange: (value: string) => void; + validation: SystemOnePayloadValidation; +} + +function ValidationStatus({ validation }: { validation: SystemOnePayloadValidation }) { + const errorCount = validation.issues.filter((issue) => issue.severity === "error").length; + if (errorCount > 0) { + return ( + + {errorCount} {errorCount === 1 ? "issue" : "issues"} + + ); + } + return Valid payload; +} + +function IssueList({ id, validation }: { id: string; validation: SystemOnePayloadValidation }) { + if (validation.issues.length === 0) { + return ( +

+ + Ready to send +

+ ); + } + return ( +
    + {validation.issues.map((issue, index) => ( +
  • + {issue.severity === "error" ? ( + + ) : ( + + )} + {issue.path} + {issue.message} +
  • + ))} +
+ ); +} + +function renderLines(rootBlocks: readonly RootBlock[]) { + return function LineRows({ rows, stylesheet, useInlineStyles }: LineRendererProps) { + return rows.map((row, line) => { + const lineElement = { node: row, stylesheet, useInlineStyles, key: line }; + const block = rootBlocks.find(({ startLine, endLine }) => line >= startLine && line <= endLine); + return ( +
+ + {line + 1} + + + {createElement(lineElement)} + +
+ ); + }); + }; +} + +export default function JsonEditor({ value, onChange, validation }: JsonEditorProps) { + const issuesId = useId(); + const highlightRef = useRef(null); + const renderer = useMemo(() => renderLines(findRootBlocks(value)), [value]); + const lineCount = value.split("\n").length; + const hasErrors = !validation.isValid; + + return ( +
+
+
+ Request JSON + +
+ + {lineCount} {lineCount === 1 ? "line" : "lines"} + +
+
+