From eef117bb790ec3ac347585b6ec90a6bdfcbd5626 Mon Sep 17 00:00:00 2001 From: fzowl Date: Fri, 27 Dec 2024 22:04:02 +0100 Subject: [PATCH 01/60] Refresh VoyageAI models and prices and context --- model_prices_and_context_window.json | 94 ++++++++++++++++++---------- 1 file changed, 62 insertions(+), 32 deletions(-) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 5f98c0e68b0..a12d114f952 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -7566,30 +7566,6 @@ "litellm_provider": "cloudflare", "mode": "chat" }, - "voyage/voyage-01": { - "max_tokens": 4096, - "max_input_tokens": 4096, - "input_cost_per_token": 0.0000001, - "output_cost_per_token": 0.000000, - "litellm_provider": "voyage", - "mode": "embedding" - }, - "voyage/voyage-lite-01": { - "max_tokens": 4096, - "max_input_tokens": 4096, - "input_cost_per_token": 0.0000001, - "output_cost_per_token": 0.000000, - "litellm_provider": "voyage", - "mode": "embedding" - }, - "voyage/voyage-large-2": { - "max_tokens": 16000, - "max_input_tokens": 16000, - "input_cost_per_token": 0.00000012, - "output_cost_per_token": 0.000000, - "litellm_provider": "voyage", - "mode": "embedding" - }, "voyage/voyage-law-2": { "max_tokens": 16000, "max_input_tokens": 16000, @@ -7614,14 +7590,6 @@ "litellm_provider": "voyage", "mode": "embedding" }, - "voyage/voyage-lite-02-instruct": { - "max_tokens": 4000, - "max_input_tokens": 4000, - "input_cost_per_token": 0.0000001, - "output_cost_per_token": 0.000000, - "litellm_provider": "voyage", - "mode": "embedding" - }, "voyage/voyage-finance-2": { "max_tokens": 4000, "max_input_tokens": 4000, @@ -7630,6 +7598,68 @@ "litellm_provider": "voyage", "mode": "embedding" }, + "voyage/voyage-3-large": { + "max_tokens": 32000, + "max_input_tokens": 32000, + "input_cost_per_token": 0.00000018, + "output_cost_per_token": 0.000000, + "litellm_provider": "voyage", + "mode": "embedding" + }, + "voyage/voyage-3": { + "max_tokens": 32000, + "max_input_tokens": 32000, + "input_cost_per_token": 0.00000006, + "output_cost_per_token": 0.000000, + "litellm_provider": "voyage", + "mode": "embedding" + }, + "voyage/voyage-3-lite": { + "max_tokens": 32000, + "max_input_tokens": 32000, + "input_cost_per_token": 0.00000002, + "output_cost_per_token": 0.000000, + "litellm_provider": "voyage", + "mode": "embedding" + }, + "voyage/voyage-code-3": { + "max_tokens": 32000, + "max_input_tokens": 32000, + "input_cost_per_token": 0.00000018, + "output_cost_per_token": 0.000000, + "litellm_provider": "voyage", + "mode": "embedding" + }, + "voyage/voyage-multimodal-3": { + "max_tokens": 32000, + "max_input_tokens": 32000, + "input_cost_per_token": 0.00000012, + "output_cost_per_token": 0.000000, + "litellm_provider": "voyage", + "mode": "embedding" + }, + "voyage/rerank-2": { + "max_tokens": 16000, + "max_input_tokens": 16000, + "max_output_tokens": 16000, + "max_query_tokens": 16000, + "input_cost_per_token": 0.00000005, + "input_cost_per_query": 0.00000005, + "output_cost_per_token": 0.0, + "litellm_provider": "voyage", + "mode": "rerank" + }, + "voyage/rerank-2-lite": { + "max_tokens": 8000, + "max_input_tokens": 8000, + "max_output_tokens": 8000, + "max_query_tokens": 8000, + "input_cost_per_token": 0.00000002, + "input_cost_per_query": 0.00000002, + "output_cost_per_token": 0.0, + "litellm_provider": "voyage", + "mode": "rerank" + }, "databricks/databricks-meta-llama-3-1-405b-instruct": { "max_tokens": 128000, "max_input_tokens": 128000, From 6876a14838605865276ac17ea9991e1874f8d64a Mon Sep 17 00:00:00 2001 From: fzowl Date: Mon, 30 Dec 2024 11:50:37 +0100 Subject: [PATCH 02/60] Refresh VoyageAI models and prices and context --- model_prices_and_context_window.json | 32 ++++++++++++++++++++++++++++ 1 file changed, 32 insertions(+) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 2ff5e40247d..0701b20af49 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -7613,6 +7613,38 @@ "litellm_provider": "cloudflare", "mode": "chat" }, + "voyage/voyage-01": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "input_cost_per_token": 0.0000001, + "output_cost_per_token": 0.000000, + "litellm_provider": "voyage", + "mode": "embedding" + }, + "voyage/voyage-lite-01": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "input_cost_per_token": 0.0000001, + "output_cost_per_token": 0.000000, + "litellm_provider": "voyage", + "mode": "embedding" + }, + "voyage/voyage-large-2": { + "max_tokens": 16000, + "max_input_tokens": 16000, + "input_cost_per_token": 0.00000012, + "output_cost_per_token": 0.000000, + "litellm_provider": "voyage", + "mode": "embedding" + }, + "voyage/voyage-lite-02-instruct": { + "max_tokens": 4000, + "max_input_tokens": 4000, + "input_cost_per_token": 0.0000001, + "output_cost_per_token": 0.000000, + "litellm_provider": "voyage", + "mode": "embedding" + }, "voyage/voyage-law-2": { "max_tokens": 16000, "max_input_tokens": 16000, From bb0f74c48f9a85874f2f104a6de0f043d4658ea6 Mon Sep 17 00:00:00 2001 From: fzowl Date: Mon, 30 Dec 2024 11:54:37 +0100 Subject: [PATCH 03/60] Refresh VoyageAI models and prices and context --- model_prices_and_context_window.json | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 0701b20af49..1a9ed533e38 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -7637,6 +7637,14 @@ "litellm_provider": "voyage", "mode": "embedding" }, + "voyage/voyage-finance-2": { + "max_tokens": 32000, + "max_input_tokens": 32000, + "input_cost_per_token": 0.00000012, + "output_cost_per_token": 0.000000, + "litellm_provider": "voyage", + "mode": "embedding" + }, "voyage/voyage-lite-02-instruct": { "max_tokens": 4000, "max_input_tokens": 4000, @@ -7669,14 +7677,6 @@ "litellm_provider": "voyage", "mode": "embedding" }, - "voyage/voyage-finance-2": { - "max_tokens": 4000, - "max_input_tokens": 4000, - "input_cost_per_token": 0.00000012, - "output_cost_per_token": 0.000000, - "litellm_provider": "voyage", - "mode": "embedding" - }, "voyage/voyage-3-large": { "max_tokens": 32000, "max_input_tokens": 32000, From 9690cb66c6e51f81db586152d822965e1db94a60 Mon Sep 17 00:00:00 2001 From: fzowl Date: Mon, 3 Feb 2025 15:21:20 +0100 Subject: [PATCH 04/60] Updating the available VoyageAI models in the docs --- docs/my-website/docs/providers/voyage.md | 27 +++++++++++++++--------- 1 file changed, 17 insertions(+), 10 deletions(-) diff --git a/docs/my-website/docs/providers/voyage.md b/docs/my-website/docs/providers/voyage.md index a56a1408ea9..6ab6b1846f5 100644 --- a/docs/my-website/docs/providers/voyage.md +++ b/docs/my-website/docs/providers/voyage.md @@ -14,7 +14,7 @@ import os os.environ['VOYAGE_API_KEY'] = "" response = embedding( - model="voyage/voyage-01", + model="voyage/voyage-3-large", input=["good morning from litellm"], ) print(response) @@ -23,13 +23,20 @@ print(response) ## Supported Models All models listed here https://docs.voyageai.com/embeddings/#models-and-specifics are supported -| Model Name | Function Call | -|--------------------------|------------------------------------------------------------------------------------------------------------------------------------------------------------------| -| voyage-2 | `embedding(model="voyage/voyage-2", input)` | -| voyage-large-2 | `embedding(model="voyage/voyage-large-2", input)` | -| voyage-law-2 | `embedding(model="voyage/voyage-law-2", input)` | -| voyage-code-2 | `embedding(model="voyage/voyage-code-2", input)` | +| Model Name | Function Call | +|-------------------------|------------------------------------------------------------| +| voyage-3-large | `embedding(model="voyage/voyage-3-large", input)` | +| voyage-3 | `embedding(model="voyage/voyage-3", input)` | +| voyage-3-lite | `embedding(model="voyage/voyage-3-lite", input)` | +| voyage-code-3 | `embedding(model="voyage/voyage-code-3", input)` | +| voyage-finance-2 | `embedding(model="voyage/voyage-finance-2", input)` | +| voyage-law-2 | `embedding(model="voyage/voyage-law-2", input)` | +| voyage-code-2 | `embedding(model="voyage/voyage-code-2", input)` | +| voyage-multilingual-2 | `embedding(model="voyage/voyage-multilingual-2 ", input)` | +| voyage-large-2-instruct | `embedding(model="voyage/voyage-large-2-instruct", input)` | +| voyage-large-2 | `embedding(model="voyage/voyage-large-2", input)` | +| voyage-2 | `embedding(model="voyage/voyage-2", input)` | | voyage-lite-02-instruct | `embedding(model="voyage/voyage-lite-02-instruct", input)` | -| voyage-01 | `embedding(model="voyage/voyage-01", input)` | -| voyage-lite-01 | `embedding(model="voyage/voyage-lite-01", input)` | -| voyage-lite-01-instruct | `embedding(model="voyage/voyage-lite-01-instruct", input)` | \ No newline at end of file +| voyage-01 | `embedding(model="voyage/voyage-01", input)` | +| voyage-lite-01 | `embedding(model="voyage/voyage-lite-01", input)` | +| voyage-lite-01-instruct | `embedding(model="voyage/voyage-lite-01-instruct", input)` | From 818db5e59032f37bd8f8440fd064607ae3f4c924 Mon Sep 17 00:00:00 2001 From: fzowl Date: Wed, 21 May 2025 13:09:42 +0200 Subject: [PATCH 05/60] Updating the available VoyageAI models in the docs --- docs/my-website/docs/providers/voyage.md | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/docs/my-website/docs/providers/voyage.md b/docs/my-website/docs/providers/voyage.md index 6ab6b1846f5..4b729bc9f58 100644 --- a/docs/my-website/docs/providers/voyage.md +++ b/docs/my-website/docs/providers/voyage.md @@ -25,6 +25,8 @@ All models listed here https://docs.voyageai.com/embeddings/#models-and-specific | Model Name | Function Call | |-------------------------|------------------------------------------------------------| +| voyage-3.5 | `embedding(model="voyage/voyage-3.5", input)` | +| voyage-3.5-lite | `embedding(model="voyage/voyage-3.5-lite", input)` | | voyage-3-large | `embedding(model="voyage/voyage-3-large", input)` | | voyage-3 | `embedding(model="voyage/voyage-3", input)` | | voyage-3-lite | `embedding(model="voyage/voyage-3-lite", input)` | @@ -35,8 +37,8 @@ All models listed here https://docs.voyageai.com/embeddings/#models-and-specific | voyage-multilingual-2 | `embedding(model="voyage/voyage-multilingual-2 ", input)` | | voyage-large-2-instruct | `embedding(model="voyage/voyage-large-2-instruct", input)` | | voyage-large-2 | `embedding(model="voyage/voyage-large-2", input)` | -| voyage-2 | `embedding(model="voyage/voyage-2", input)` | +| voyage-2 | `embedding(model="voyage/voyage-2", input)` | | voyage-lite-02-instruct | `embedding(model="voyage/voyage-lite-02-instruct", input)` | -| voyage-01 | `embedding(model="voyage/voyage-01", input)` | -| voyage-lite-01 | `embedding(model="voyage/voyage-lite-01", input)` | +| voyage-01 | `embedding(model="voyage/voyage-01", input)` | +| voyage-lite-01 | `embedding(model="voyage/voyage-lite-01", input)` | | voyage-lite-01-instruct | `embedding(model="voyage/voyage-lite-01-instruct", input)` | From f88dd65590460df01c49a6a58ab95ee7f28d8b58 Mon Sep 17 00:00:00 2001 From: fzowl Date: Wed, 29 Jul 2026 14:31:25 +0200 Subject: [PATCH 06/60] feat(voyage): add voyage-4 family + voyage-context-4, fix contextual list[str] input - Add voyage-4, voyage-4-large, voyage-4-lite, voyage-4-nano (/bin/bash) and voyage-context-4 to the cost/context map (both root and backup json). - Contextual embeddings: normalize flat list[str] input. Voyage rejects a flat list[str] unless input_type='query', so wrap each string as its own single-chunk document otherwise; single str -> [[str]]; list[list[str]] passes through. Keeps list[str] when input_type='query'. - Add tests for input normalization and voyage-4 family pricing. --- .../embedding/transformation_contextual.py | 33 ++++++- ...odel_prices_and_context_window_backup.json | 40 +++++++++ model_prices_and_context_window.json | 40 +++++++++ tests/llm_translation/test_voyage_ai.py | 85 +++++++++++++++++++ 4 files changed, 197 insertions(+), 1 deletion(-) diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index 4df2fa4ba31..2d661610eaa 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -106,11 +106,42 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): headers: dict, ) -> dict: return { - "inputs": input, + "inputs": self._transform_contextual_inputs(input, optional_params), "model": model, **optional_params, } + @staticmethod + def _transform_contextual_inputs( + input: Union[AllEmbeddingInputValues, List[List[str]]], + optional_params: dict, + ) -> List[List[str]]: + """ + Voyage's contextualized embeddings API expects ``inputs`` to be a + ``list[list[str]]`` (each inner list is a document made of chunks that + share context). + + It also accepts a flat ``list[str]`` - but only when + ``input_type == "query"``. In every other case a flat list must be + wrapped so each string becomes its own single-chunk document, otherwise + the API rejects the request with a 400. + + Reference: https://docs.voyageai.com/reference/contextualized-embeddings-api + """ + # Single string -> one document with a single chunk. + if isinstance(input, str): + return [[input]] + + # Flat list[str]: keep as list[str] when the API allows it + # (input_type="query"), otherwise wrap each string as its own document. + if isinstance(input, list) and all(isinstance(i, str) for i in input): + if optional_params.get("input_type") == "query": + return input # type: ignore[return-value] + return [[i] for i in input] + + # Already list[list[str]] (or another shape) -> pass through unchanged. + return input # type: ignore[return-value] + def transform_embedding_response( self, model: str, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index c584deb683a..fc710497592 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -27345,6 +27345,38 @@ "mode": "embedding", "output_cost_per_token": 0.0 }, + "voyage/voyage-4": { + "input_cost_per_token": 6e-08, + "litellm_provider": "voyage", + "max_input_tokens": 32000, + "max_tokens": 32000, + "mode": "embedding", + "output_cost_per_token": 0.0 + }, + "voyage/voyage-4-large": { + "input_cost_per_token": 1.2e-07, + "litellm_provider": "voyage", + "max_input_tokens": 32000, + "max_tokens": 32000, + "mode": "embedding", + "output_cost_per_token": 0.0 + }, + "voyage/voyage-4-lite": { + "input_cost_per_token": 2e-08, + "litellm_provider": "voyage", + "max_input_tokens": 32000, + "max_tokens": 32000, + "mode": "embedding", + "output_cost_per_token": 0.0 + }, + "voyage/voyage-4-nano": { + "input_cost_per_token": 0.0, + "litellm_provider": "voyage", + "max_input_tokens": 32000, + "max_tokens": 32000, + "mode": "embedding", + "output_cost_per_token": 0.0 + }, "voyage/voyage-code-2": { "input_cost_per_token": 1.2e-07, "litellm_provider": "voyage", @@ -27369,6 +27401,14 @@ "mode": "embedding", "output_cost_per_token": 0.0 }, + "voyage/voyage-context-4": { + "input_cost_per_token": 1.2e-07, + "litellm_provider": "voyage", + "max_input_tokens": 120000, + "max_tokens": 120000, + "mode": "embedding", + "output_cost_per_token": 0.0 + }, "voyage/voyage-finance-2": { "input_cost_per_token": 1.2e-07, "litellm_provider": "voyage", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index c584deb683a..fc710497592 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -27345,6 +27345,38 @@ "mode": "embedding", "output_cost_per_token": 0.0 }, + "voyage/voyage-4": { + "input_cost_per_token": 6e-08, + "litellm_provider": "voyage", + "max_input_tokens": 32000, + "max_tokens": 32000, + "mode": "embedding", + "output_cost_per_token": 0.0 + }, + "voyage/voyage-4-large": { + "input_cost_per_token": 1.2e-07, + "litellm_provider": "voyage", + "max_input_tokens": 32000, + "max_tokens": 32000, + "mode": "embedding", + "output_cost_per_token": 0.0 + }, + "voyage/voyage-4-lite": { + "input_cost_per_token": 2e-08, + "litellm_provider": "voyage", + "max_input_tokens": 32000, + "max_tokens": 32000, + "mode": "embedding", + "output_cost_per_token": 0.0 + }, + "voyage/voyage-4-nano": { + "input_cost_per_token": 0.0, + "litellm_provider": "voyage", + "max_input_tokens": 32000, + "max_tokens": 32000, + "mode": "embedding", + "output_cost_per_token": 0.0 + }, "voyage/voyage-code-2": { "input_cost_per_token": 1.2e-07, "litellm_provider": "voyage", @@ -27369,6 +27401,14 @@ "mode": "embedding", "output_cost_per_token": 0.0 }, + "voyage/voyage-context-4": { + "input_cost_per_token": 1.2e-07, + "litellm_provider": "voyage", + "max_input_tokens": 120000, + "max_tokens": 120000, + "mode": "embedding", + "output_cost_per_token": 0.0 + }, "voyage/voyage-finance-2": { "input_cost_per_token": 1.2e-07, "litellm_provider": "voyage", diff --git a/tests/llm_translation/test_voyage_ai.py b/tests/llm_translation/test_voyage_ai.py index a0b9ee0a44b..577ee808035 100644 --- a/tests/llm_translation/test_voyage_ai.py +++ b/tests/llm_translation/test_voyage_ai.py @@ -71,6 +71,26 @@ class TestVoyageAI(BaseLLMEmbeddingTest): assert response.usage.total_tokens > 0 +@pytest.mark.parametrize( + "model, expected_input_cost", + [ + ("voyage/voyage-4", 6e-08), + ("voyage/voyage-4-large", 1.2e-07), + ("voyage/voyage-4-lite", 2e-08), + ("voyage/voyage-4-nano", 0.0), + ("voyage/voyage-context-4", 1.2e-07), + ], +) +def test_voyage_4_family_pricing_registered(model, expected_input_cost): + """The voyage-4 family and voyage-context-4 must be in the cost map.""" + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + info = litellm.get_model_info(model=model) + assert info["litellm_provider"] == "voyage" + assert info["mode"] == "embedding" + assert info["input_cost_per_token"] == expected_input_cost + + def test_voyage_ai_embedding_extra_params(): """Test Voyage AI embedding with extra parameters""" try: @@ -142,6 +162,7 @@ class TestVoyageContextualEmbeddings: # Test contextual model detection assert config.is_contextualized_embeddings("voyage-context-3") is True + assert config.is_contextualized_embeddings("voyage-context-4") is True assert config.is_contextualized_embeddings("voyage-context-2") is True assert config.is_contextualized_embeddings("context-model") is True @@ -198,6 +219,70 @@ class TestVoyageContextualEmbeddings: assert transformed["model"] == "voyage-context-3" assert transformed["encoding_format"] == "float" + def test_contextual_flat_list_str_is_wrapped(self): + """Flat list[str] without input_type must be wrapped to list[list[str]].""" + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + + transformed = config.transform_embedding_request( + "voyage-context-4", ["Hello", "world"], {}, {} + ) + + # Voyage rejects a flat list[str] unless input_type="query", so each + # string becomes its own single-chunk document. + assert transformed["inputs"] == [["Hello"], ["world"]] + assert transformed["model"] == "voyage-context-4" + + def test_contextual_flat_list_str_query_stays_flat(self): + """Flat list[str] with input_type='query' is sent as-is (list[str]).""" + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + + transformed = config.transform_embedding_request( + "voyage-context-4", + ["Hello", "world"], + {"input_type": "query"}, + {}, + ) + + assert transformed["inputs"] == ["Hello", "world"] + assert transformed["input_type"] == "query" + + def test_contextual_single_string_is_wrapped(self): + """A single string is wrapped to a single document with a single chunk.""" + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + + transformed = config.transform_embedding_request( + "voyage-context-4", "Hello", {}, {} + ) + + assert transformed["inputs"] == [["Hello"]] + + def test_contextual_nested_input_passthrough(self): + """Already-nested list[list[str]] input is passed through unchanged.""" + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + + nested = [["Hello", "world"], ["Test"]] + transformed = config.transform_embedding_request( + "voyage-context-4", nested, {}, {} + ) + + assert transformed["inputs"] == nested + def test_contextual_embedding_response_transformation(self): """Test response transformation for contextual embeddings""" from litellm.llms.voyage.embedding.transformation_contextual import ( From 328b54d0042d213aee3e1e2c025480eae5a214d8 Mon Sep 17 00:00:00 2001 From: fzowl Date: Wed, 29 Jul 2026 15:22:31 +0200 Subject: [PATCH 07/60] fix(voyage): auto-chunk flat list[str] for contextual embeddings; fix mypy Rework per review: - Contextual inputs now prefer a flat list[str] and let the API auto-chunk (enable_auto_chunking=True, chunk_size=32000, input_type=document) instead of pre-wrapping into list[list[str]]. Matches the live API contract: * flat list[str] + input_type=query -> kept flat * flat list[str] otherwise / single str -> flat + auto-chunking * list[list[str]] -> passthrough - Fix mypy list-item error from the previous wrapping approach. - Update tests for the new normalization; all pass. --- .../embedding/transformation_contextual.py | 73 ++++++++++++++----- tests/llm_translation/test_voyage_ai.py | 47 ++++++++++-- 2 files changed, 94 insertions(+), 26 deletions(-) diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index 2d661610eaa..d555a1a2266 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -2,7 +2,7 @@ This module is used to transform the request and response for the Voyage contextualized embeddings API. This would be used for all the contextualized embeddings models in Voyage. """ -from typing import List, Optional, Union +from typing import Any, Dict, List, Optional, Tuple, Union import httpx @@ -98,6 +98,11 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): "Authorization": f"Bearer {api_key}", } + # Chunk size (in tokens) used when the API auto-chunks a flat ``list[str]``. + # Matches the voyage-context-4 context window so each string stays a single + # chunk instead of being split. + AUTO_CHUNK_SIZE = 32000 + def transform_embedding_request( self, model: str, @@ -105,42 +110,74 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): optional_params: dict, headers: dict, ) -> dict: + inputs, extra_params = self._transform_contextual_inputs(input, optional_params) return { - "inputs": self._transform_contextual_inputs(input, optional_params), + "inputs": inputs, "model": model, **optional_params, + **extra_params, } - @staticmethod + @classmethod def _transform_contextual_inputs( + cls, input: Union[AllEmbeddingInputValues, List[List[str]]], optional_params: dict, - ) -> List[List[str]]: + ) -> Tuple[Union[List[str], List[List[str]]], dict]: """ - Voyage's contextualized embeddings API expects ``inputs`` to be a - ``list[list[str]]`` (each inner list is a document made of chunks that - share context). + Normalize ``input`` for Voyage's contextualized embeddings API and + return ``(inputs, extra_params)`` where ``extra_params`` carries any + request fields (e.g. auto-chunking) needed for the chosen shape. - It also accepts a flat ``list[str]`` - but only when - ``input_type == "query"``. In every other case a flat list must be - wrapped so each string becomes its own single-chunk document, otherwise - the API rejects the request with a 400. + The API contract (verified against the live endpoint) is: + + - A flat ``list[str]`` is only accepted with ``input_type="query"`` or + with ``enable_auto_chunking=True`` (which itself requires + ``input_type="document"``). + - A ``list[list[str]]`` (each inner list = one document's chunks) is + always accepted. + + So we prefer to send a flat ``list[str]`` and let the API auto-chunk, + instead of pre-wrapping into ``list[list[str]]``: + + - ``str`` -> ``[str]`` + ``enable_auto_chunking`` (input_type=document) + - flat ``list[str]`` + ``input_type="query"`` -> kept flat, as-is + - flat ``list[str]`` otherwise -> kept flat + ``enable_auto_chunking`` + (input_type=document) + - ``list[list[str]]`` -> passed through unchanged Reference: https://docs.voyageai.com/reference/contextualized-embeddings-api """ - # Single string -> one document with a single chunk. + # Single string -> a one-element flat list, auto-chunked. if isinstance(input, str): - return [[input]] + return [input], cls._auto_chunk_params(optional_params) - # Flat list[str]: keep as list[str] when the API allows it - # (input_type="query"), otherwise wrap each string as its own document. + # Flat list[str]. if isinstance(input, list) and all(isinstance(i, str) for i in input): if optional_params.get("input_type") == "query": - return input # type: ignore[return-value] - return [[i] for i in input] + # The API accepts a flat query list as-is. + return input, {} # type: ignore[return-value] + # Otherwise let the API auto-chunk the flat list. + return input, cls._auto_chunk_params(optional_params) # type: ignore[return-value] # Already list[list[str]] (or another shape) -> pass through unchanged. - return input # type: ignore[return-value] + return input, {} # type: ignore[return-value] + + @classmethod + def _auto_chunk_params(cls, optional_params: dict) -> dict: + """ + Params required to send a flat ``list[str]`` to the contextualized API. + + ``enable_auto_chunking=True`` requires ``input_type="document"``, so set + it unless the caller already provided an ``input_type``. + """ + params: Dict[str, Any] = { + "enable_auto_chunking": True, + "chunk_size": cls.AUTO_CHUNK_SIZE, + } + if not optional_params.get("input_type"): + params["input_type"] = "document" + return params def transform_embedding_response( self, diff --git a/tests/llm_translation/test_voyage_ai.py b/tests/llm_translation/test_voyage_ai.py index 577ee808035..b7a2c41a979 100644 --- a/tests/llm_translation/test_voyage_ai.py +++ b/tests/llm_translation/test_voyage_ai.py @@ -219,8 +219,13 @@ class TestVoyageContextualEmbeddings: assert transformed["model"] == "voyage-context-3" assert transformed["encoding_format"] == "float" - def test_contextual_flat_list_str_is_wrapped(self): - """Flat list[str] without input_type must be wrapped to list[list[str]].""" + def test_contextual_flat_list_str_is_auto_chunked(self): + """Flat list[str] without input_type stays flat and is auto-chunked. + + The API accepts a flat list[str] only with input_type="query" or with + enable_auto_chunking=True (which requires input_type="document"), so we + keep the list flat and let the API auto-chunk it. + """ from litellm.llms.voyage.embedding.transformation_contextual import ( VoyageContextualEmbeddingConfig, ) @@ -231,10 +236,11 @@ class TestVoyageContextualEmbeddings: "voyage-context-4", ["Hello", "world"], {}, {} ) - # Voyage rejects a flat list[str] unless input_type="query", so each - # string becomes its own single-chunk document. - assert transformed["inputs"] == [["Hello"], ["world"]] + assert transformed["inputs"] == ["Hello", "world"] assert transformed["model"] == "voyage-context-4" + assert transformed["enable_auto_chunking"] is True + assert transformed["chunk_size"] == 32000 + assert transformed["input_type"] == "document" def test_contextual_flat_list_str_query_stays_flat(self): """Flat list[str] with input_type='query' is sent as-is (list[str]).""" @@ -253,9 +259,31 @@ class TestVoyageContextualEmbeddings: assert transformed["inputs"] == ["Hello", "world"] assert transformed["input_type"] == "query" + # A query list is accepted as-is, no auto-chunking needed. + assert "enable_auto_chunking" not in transformed - def test_contextual_single_string_is_wrapped(self): - """A single string is wrapped to a single document with a single chunk.""" + def test_contextual_flat_list_str_document_input_type_preserved(self): + """Explicit input_type='document' is preserved while auto-chunking.""" + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + + transformed = config.transform_embedding_request( + "voyage-context-4", + ["Hello", "world"], + {"input_type": "document"}, + {}, + ) + + assert transformed["inputs"] == ["Hello", "world"] + assert transformed["input_type"] == "document" + assert transformed["enable_auto_chunking"] is True + assert transformed["chunk_size"] == 32000 + + def test_contextual_single_string_is_auto_chunked(self): + """A single string becomes a one-element flat list and is auto-chunked.""" from litellm.llms.voyage.embedding.transformation_contextual import ( VoyageContextualEmbeddingConfig, ) @@ -266,7 +294,10 @@ class TestVoyageContextualEmbeddings: "voyage-context-4", "Hello", {}, {} ) - assert transformed["inputs"] == [["Hello"]] + assert transformed["inputs"] == ["Hello"] + assert transformed["enable_auto_chunking"] is True + assert transformed["chunk_size"] == 32000 + assert transformed["input_type"] == "document" def test_contextual_nested_input_passthrough(self): """Already-nested list[list[str]] input is passed through unchanged.""" From be6487d33c43dcf28de410fe209ede88712513b8 Mon Sep 17 00:00:00 2001 From: fzowl <160063452+fzowl@users.noreply.github.com> Date: Wed, 29 Jul 2026 17:20:27 +0200 Subject: [PATCH 08/60] Update litellm/llms/voyage/embedding/transformation_contextual.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --- litellm/llms/voyage/embedding/transformation_contextual.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index 84ed856e080..c4eb6f34283 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -148,6 +148,8 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): """ # Single string -> a one-element flat list, auto-chunked. if isinstance(input, str): + if optional_params.get("input_type") == "query": + return [input], {} return [input], cls._auto_chunk_params(optional_params) # Flat list[str]. From e8ae29f135d35bbea296d4fa935ef1820c65bb0d Mon Sep 17 00:00:00 2001 From: fzowl Date: Wed, 29 Jul 2026 17:45:08 +0200 Subject: [PATCH 09/60] fix(voyage): ruff format + add missing query-string test coverage --- .../embedding/transformation_contextual.py | 1 + tests/llm_translation/test_voyage_ai.py | 16 ++++++++++++++++ 2 files changed, 17 insertions(+) diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index c4eb6f34283..4d8bcc7ff80 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -2,6 +2,7 @@ This module is used to transform the request and response for the Voyage contextualized embeddings API. This would be used for all the contextualized embeddings models in Voyage. """ + from typing import Any, Dict, List, Optional, Tuple, Union import httpx diff --git a/tests/llm_translation/test_voyage_ai.py b/tests/llm_translation/test_voyage_ai.py index a0d291bc02f..e6c73240a0e 100644 --- a/tests/llm_translation/test_voyage_ai.py +++ b/tests/llm_translation/test_voyage_ai.py @@ -300,6 +300,22 @@ class TestVoyageContextualEmbeddings: assert transformed["chunk_size"] == 32000 assert transformed["input_type"] == "document" + def test_contextual_single_string_query_no_auto_chunk(self): + """A single string with input_type='query' is not auto-chunked.""" + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + + transformed = config.transform_embedding_request( + "voyage-context-4", "Hello", {"input_type": "query"}, {} + ) + + assert transformed["inputs"] == ["Hello"] + assert transformed["input_type"] == "query" + assert "enable_auto_chunking" not in transformed + def test_contextual_nested_input_passthrough(self): """Already-nested list[list[str]] input is passed through unchanged.""" from litellm.llms.voyage.embedding.transformation_contextual import ( From bb6db24d58bb93454a77bbcb283f1f6ad4953a1d Mon Sep 17 00:00:00 2001 From: fzowl Date: Wed, 29 Jul 2026 18:37:25 +0200 Subject: [PATCH 10/60] fix(voyage): use modern type annotations to satisfy strict ruff budget (TID251, UP006) --- .../embedding/transformation_contextual.py | 32 +++++++++---------- 1 file changed, 16 insertions(+), 16 deletions(-) diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index 4d8bcc7ff80..8f029c27be7 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -3,7 +3,9 @@ This module is used to transform the request and response for the Voyage context This would be used for all the contextualized embeddings models in Voyage. """ -from typing import Any, Dict, List, Optional, Tuple, Union +from __future__ import annotations + +from typing import Any import httpx @@ -20,7 +22,7 @@ class VoyageError(BaseLLMException): self, status_code: int, message: str, - headers: Union[dict, httpx.Headers] = {}, + headers: dict | httpx.Headers = {}, ): self.status_code = status_code self.message = message @@ -43,12 +45,12 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = None, ) -> str: if api_base: if not api_base.endswith("/contextualizedembeddings"): @@ -81,11 +83,11 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): self, headers: dict, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, ) -> dict: if api_key is None: api_key = ( @@ -105,7 +107,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): def transform_embedding_request( self, model: str, - input: Union[AllEmbeddingInputValues, List[List[str]]], + input: AllEmbeddingInputValues | list[list[str]], optional_params: dict, headers: dict, ) -> dict: @@ -120,9 +122,9 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): @classmethod def _transform_contextual_inputs( cls, - input: Union[AllEmbeddingInputValues, List[List[str]]], + input: AllEmbeddingInputValues | list[list[str]], optional_params: dict, - ) -> Tuple[Union[List[str], List[List[str]]], dict]: + ) -> tuple[list[str] | list[list[str]], dict]: """ Normalize ``input`` for Voyage's contextualized embeddings API and return ``(inputs, extra_params)`` where ``extra_params`` carries any @@ -172,7 +174,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): ``enable_auto_chunking=True`` requires ``input_type="document"``, so set it unless the caller already provided an ``input_type``. """ - params: Dict[str, Any] = { + params: dict[str, Any] = { "enable_auto_chunking": True, "chunk_size": cls.AUTO_CHUNK_SIZE, } @@ -186,7 +188,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): raw_response: httpx.Response, model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, - api_key: Optional[str] = None, + api_key: str | None = None, request_data: dict = {}, optional_params: dict = {}, litellm_params: dict = {}, @@ -208,9 +210,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): model_response.usage = usage return model_response - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return VoyageError(message=error_message, status_code=status_code, headers=headers) @staticmethod From eddb2225dc6dbc79238f4e6a7815c40cb7957a6d Mon Sep 17 00:00:00 2001 From: fzowl Date: Wed, 29 Jul 2026 21:29:38 +0200 Subject: [PATCH 11/60] fix(voyage): add mutable-ok pragmas and replace type: ignore with pyright: ignore (LIT001, LIT009) --- .../embedding/transformation_contextual.py | 84 ++++++++++--------- 1 file changed, 43 insertions(+), 41 deletions(-) diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index 8f029c27be7..04cb58bcfcc 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -22,11 +22,14 @@ class VoyageError(BaseLLMException): self, status_code: int, message: str, - headers: dict | httpx.Headers = {}, + headers: dict | httpx.Headers = {}, # mutable-ok: base class signature ): self.status_code = status_code self.message = message - self.request = httpx.Request(method="POST", url="https://api.voyageai.com/v1/contextualizedembeddings") + self.request = httpx.Request( + method="POST", + url="https://api.voyageai.com/v1/contextualizedembeddings", + ) self.response = httpx.Response(status_code=status_code, request=self.request) super().__init__( status_code=status_code, @@ -48,8 +51,8 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): api_base: str | None, api_key: str | None, model: str, - optional_params: dict, - litellm_params: dict, + optional_params: dict, # mutable-ok: base class signature + litellm_params: dict, # mutable-ok: base class signature stream: bool | None = None, ) -> str: if api_base: @@ -58,16 +61,16 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): return api_base return "https://api.voyageai.com/v1/contextualizedembeddings" - def get_supported_openai_params(self, model: str) -> list: + def get_supported_openai_params(self, model: str) -> list: # mutable-ok: base class signature return ["encoding_format", "dimensions"] def map_openai_params( self, - non_default_params: dict, - optional_params: dict, + non_default_params: dict, # mutable-ok: base class signature + optional_params: dict, # mutable-ok: base class signature model: str, drop_params: bool, - ) -> dict: + ) -> dict: # mutable-ok: base class signature """ Map OpenAI params to Voyage params @@ -81,14 +84,14 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): def validate_environment( self, - headers: dict, + headers: dict, # mutable-ok: base class signature model: str, - messages: list[AllMessageValues], - optional_params: dict, - litellm_params: dict, + messages: list[AllMessageValues], # mutable-ok: base class signature + optional_params: dict, # mutable-ok: base class signature + litellm_params: dict, # mutable-ok: base class signature api_key: str | None = None, api_base: str | None = None, - ) -> dict: + ) -> dict: # mutable-ok: base class signature if api_key is None: api_key = ( get_secret_str("VOYAGE_API_KEY") @@ -99,18 +102,15 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): "Authorization": f"Bearer {api_key}", } - # Chunk size (in tokens) used when the API auto-chunks a flat ``list[str]``. - # Matches the voyage-context-4 context window so each string stays a single - # chunk instead of being split. AUTO_CHUNK_SIZE = 32000 def transform_embedding_request( self, model: str, - input: AllEmbeddingInputValues | list[list[str]], - optional_params: dict, - headers: dict, - ) -> dict: + input: AllEmbeddingInputValues | list[list[str]], # mutable-ok: base class signature + optional_params: dict, # mutable-ok: base class signature + headers: dict, # mutable-ok: base class signature + ) -> dict: # mutable-ok: base class signature inputs, extra_params = self._transform_contextual_inputs(input, optional_params) return { "inputs": inputs, @@ -122,9 +122,9 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): @classmethod def _transform_contextual_inputs( cls, - input: AllEmbeddingInputValues | list[list[str]], - optional_params: dict, - ) -> tuple[list[str] | list[list[str]], dict]: + input: AllEmbeddingInputValues | list[list[str]], # mutable-ok: union with AllEmbeddingInputValues + optional_params: dict, # mutable-ok: matches public API + ) -> tuple[list[str] | list[list[str]], dict]: # mutable-ok: returned to caller who owns it """ Normalize ``input`` for Voyage's contextualized embeddings API and return ``(inputs, extra_params)`` where ``extra_params`` carries any @@ -149,32 +149,27 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): Reference: https://docs.voyageai.com/reference/contextualized-embeddings-api """ - # Single string -> a one-element flat list, auto-chunked. if isinstance(input, str): if optional_params.get("input_type") == "query": - return [input], {} - return [input], cls._auto_chunk_params(optional_params) + return [input], {} # mutable-ok: fresh list returned to caller + return [input], cls._auto_chunk_params(optional_params) # mutable-ok: fresh list returned to caller - # Flat list[str]. if isinstance(input, list) and all(isinstance(i, str) for i in input): if optional_params.get("input_type") == "query": - # The API accepts a flat query list as-is. - return input, {} # type: ignore[return-value] - # Otherwise let the API auto-chunk the flat list. - return input, cls._auto_chunk_params(optional_params) # type: ignore[return-value] + return input, {} # pyright: ignore[reportReturnType] # narrowed to list[str] by isinstance+all check + return input, cls._auto_chunk_params(optional_params) # pyright: ignore[reportReturnType] # narrowed to list[str] - # Already list[list[str]] (or another shape) -> pass through unchanged. - return input, {} # type: ignore[return-value] + return input, {} # pyright: ignore[reportReturnType] # list[list[str]] branch @classmethod - def _auto_chunk_params(cls, optional_params: dict) -> dict: + def _auto_chunk_params(cls, optional_params: dict) -> dict: # mutable-ok: matches public API """ Params required to send a flat ``list[str]`` to the contextualized API. ``enable_auto_chunking=True`` requires ``input_type="document"``, so set it unless the caller already provided an ``input_type``. """ - params: dict[str, Any] = { + params: dict[str, Any] = { # mutable-ok: building return value "enable_auto_chunking": True, "chunk_size": cls.AUTO_CHUNK_SIZE, } @@ -189,16 +184,18 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): model_response: EmbeddingResponse, logging_obj: LiteLLMLoggingObj, api_key: str | None = None, - request_data: dict = {}, - optional_params: dict = {}, - litellm_params: dict = {}, + request_data: dict = {}, # mutable-ok: base class signature + optional_params: dict = {}, # mutable-ok: base class signature + litellm_params: dict = {}, # mutable-ok: base class signature ) -> EmbeddingResponse: try: raw_response_json = raw_response.json() except Exception: - raise VoyageError(message=raw_response.text, status_code=raw_response.status_code) + raise VoyageError( + message=raw_response.text, + status_code=raw_response.status_code, + ) - # model_response.usage model_response.model = raw_response_json.get("model") model_response.data = raw_response_json.get("data") model_response.object = raw_response_json.get("object") @@ -210,7 +207,12 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): model_response.usage = usage return model_response - def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: + def get_error_class( + self, + error_message: str, + status_code: int, + headers: dict | httpx.Headers, # mutable-ok: base class signature + ) -> BaseLLMException: return VoyageError(message=error_message, status_code=status_code, headers=headers) @staticmethod From af634ab4265d00f3fe730623a7f28d4ddef86d42 Mon Sep 17 00:00:00 2001 From: fzowl Date: Wed, 29 Jul 2026 22:13:15 +0200 Subject: [PATCH 12/60] test(voyage): add unit tests for contextual embedding in test_litellm (CI coverage path) --- .../test_voyage_contextual_embedding.py | 222 ++++++++++++++++++ 1 file changed, 222 insertions(+) create mode 100644 tests/test_litellm/llms/voyage/test_voyage_contextual_embedding.py diff --git a/tests/test_litellm/llms/voyage/test_voyage_contextual_embedding.py b/tests/test_litellm/llms/voyage/test_voyage_contextual_embedding.py new file mode 100644 index 00000000000..865608e286f --- /dev/null +++ b/tests/test_litellm/llms/voyage/test_voyage_contextual_embedding.py @@ -0,0 +1,222 @@ +import json +from unittest.mock import MagicMock + +import pytest + + +class TestVoyageContextualEmbeddings: + def test_contextual_model_detection(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + assert VoyageContextualEmbeddingConfig.is_contextualized_embeddings("voyage-context-3") + assert VoyageContextualEmbeddingConfig.is_contextualized_embeddings("voyage-context-4") + assert not VoyageContextualEmbeddingConfig.is_contextualized_embeddings("voyage-3-lite") + + def test_url_generation(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + assert ( + config.get_complete_url(None, None, "voyage-context-4", {}, {}) + == "https://api.voyageai.com/v1/contextualizedembeddings" + ) + assert ( + config.get_complete_url("https://custom.api.com", None, "voyage-context-4", {}, {}) + == "https://custom.api.com/contextualizedembeddings" + ) + assert ( + config.get_complete_url( + "https://custom.api.com/contextualizedembeddings", + None, + "voyage-context-4", + {}, + {}, + ) + == "https://custom.api.com/contextualizedembeddings" + ) + + def test_get_supported_openai_params(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + assert config.get_supported_openai_params("voyage-context-4") == [ + "encoding_format", + "dimensions", + ] + + def test_map_openai_params(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + result = config.map_openai_params( + {"encoding_format": "float", "dimensions": 512}, {}, "voyage-context-4", False + ) + assert result["encoding_format"] == "float" + assert result["output_dimension"] == 512 + + def test_validate_environment_with_api_key(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + headers = config.validate_environment( + {}, "voyage-context-4", [], {}, {}, api_key="test-key" + ) + assert headers == {"Authorization": "Bearer test-key"} + + def test_validate_environment_secret_fallback(self, monkeypatch): + import litellm.llms.voyage.embedding.transformation_contextual as module + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + def fake_get_secret(name): + return "secret-key" if name == "VOYAGE_API_KEY" else None + + monkeypatch.setattr(module, "get_secret_str", fake_get_secret) + config = VoyageContextualEmbeddingConfig() + headers = config.validate_environment( + {}, "voyage-context-4", [], {}, {}, api_key=None + ) + assert headers == {"Authorization": "Bearer secret-key"} + + def test_nested_list_passthrough(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + nested = [["Hello", "world"], ["Test"]] + transformed = config.transform_embedding_request( + "voyage-context-4", nested, {}, {} + ) + assert transformed["inputs"] == nested + assert transformed["model"] == "voyage-context-4" + assert "enable_auto_chunking" not in transformed + + def test_flat_list_str_auto_chunked(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + transformed = config.transform_embedding_request( + "voyage-context-4", ["Hello", "world"], {}, {} + ) + assert transformed["inputs"] == ["Hello", "world"] + assert transformed["enable_auto_chunking"] is True + assert transformed["chunk_size"] == 32000 + assert transformed["input_type"] == "document" + + def test_flat_list_str_query_no_auto_chunk(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + transformed = config.transform_embedding_request( + "voyage-context-4", ["Hello", "world"], {"input_type": "query"}, {} + ) + assert transformed["inputs"] == ["Hello", "world"] + assert transformed["input_type"] == "query" + assert "enable_auto_chunking" not in transformed + + def test_flat_list_str_document_preserves_input_type(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + transformed = config.transform_embedding_request( + "voyage-context-4", ["Hello"], {"input_type": "document"}, {} + ) + assert transformed["input_type"] == "document" + assert transformed["enable_auto_chunking"] is True + + def test_single_string_auto_chunked(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + transformed = config.transform_embedding_request( + "voyage-context-4", "Hello", {}, {} + ) + assert transformed["inputs"] == ["Hello"] + assert transformed["enable_auto_chunking"] is True + assert transformed["input_type"] == "document" + + def test_single_string_query_no_auto_chunk(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + transformed = config.transform_embedding_request( + "voyage-context-4", "Hello", {"input_type": "query"}, {} + ) + assert transformed["inputs"] == ["Hello"] + assert transformed["input_type"] == "query" + assert "enable_auto_chunking" not in transformed + + def test_response_transformation(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + from litellm.types.utils import EmbeddingResponse + + config = VoyageContextualEmbeddingConfig() + response_payload = { + "object": "list", + "data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}], + "model": "voyage-context-4", + "usage": {"total_tokens": 24}, + } + raw_response = MagicMock() + raw_response.json.return_value = response_payload + raw_response.status_code = 200 + raw_response.text = json.dumps(response_payload) + + model_response = EmbeddingResponse() + transformed = config.transform_embedding_response( + "voyage-context-4", raw_response, model_response, MagicMock() + ) + assert transformed.model == "voyage-context-4" + assert transformed.object == "list" + assert transformed.data == response_payload["data"] + assert transformed.usage.prompt_tokens == 24 + assert transformed.usage.total_tokens == 24 + + def test_error_response_and_error_class(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + VoyageError, + ) + from litellm.types.utils import EmbeddingResponse + + config = VoyageContextualEmbeddingConfig() + raw_response = MagicMock() + raw_response.json.side_effect = ValueError("not json") + raw_response.status_code = 400 + raw_response.text = "bad request" + + with pytest.raises(VoyageError) as exc_info: + config.transform_embedding_response( + "voyage-context-4", raw_response, EmbeddingResponse(), MagicMock() + ) + assert exc_info.value.status_code == 400 + assert exc_info.value.message == "bad request" + + error = config.get_error_class("rate limited", 429, {"x-test": "1"}) + assert isinstance(error, VoyageError) + assert error.status_code == 429 + assert error.message == "rate limited" From a1be4236322d7615eece57ed6cc84867e194e866 Mon Sep 17 00:00:00 2001 From: fzowl Date: Sat, 15 Aug 2026 20:29:31 +0200 Subject: [PATCH 13/60] fix(voyage): drop banned typing.Any from contextual transform to satisfy strict lint gate --- litellm/llms/voyage/embedding/transformation_contextual.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index d473450a969..0c02f00fbd5 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -3,7 +3,7 @@ This module is used to transform the request and response for the Voyage context This would be used for all the contextualized embeddings models in Voyage. """ -from typing import Any, Final +from typing import Final import httpx @@ -167,7 +167,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): ``enable_auto_chunking=True`` requires ``input_type="document"``, so set it unless the caller already provided an ``input_type``. """ - params: dict[str, Any] = { # mutable-ok: building return value + params: dict[str, object] = { # mutable-ok: building return value "enable_auto_chunking": True, "chunk_size": cls.AUTO_CHUNK_SIZE, } From 8c8cd890adcc14ed8c13d2b315e8eebc6d5c031d Mon Sep 17 00:00:00 2001 From: fzowl Date: Sat, 15 Aug 2026 20:38:40 +0200 Subject: [PATCH 14/60] fix(voyage): drop redundant isinstance list check to satisfy basedpyright budget --- litellm/llms/voyage/embedding/transformation_contextual.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index 0c02f00fbd5..a9a4a01f8e3 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -152,9 +152,9 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): return [input], {} # mutable-ok: fresh list returned to caller return [input], cls._auto_chunk_params(optional_params) # mutable-ok: fresh list returned to caller - if isinstance(input, list) and all(isinstance(i, str) for i in input): + if all(isinstance(i, str) for i in input): if optional_params.get("input_type") == "query": - return input, {} # pyright: ignore[reportReturnType] # narrowed to list[str] by isinstance+all check + return input, {} # pyright: ignore[reportReturnType] # narrowed to list[str] by all(isinstance) check return input, cls._auto_chunk_params(optional_params) # pyright: ignore[reportReturnType] # narrowed to list[str] return input, {} # pyright: ignore[reportReturnType] # list[list[str]] branch From 88b3cfd789b5726af3acb7387dc7acb54dbacd85 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 29 Aug 2026 00:55:43 +0000 Subject: [PATCH 15/60] fix(mcp): bind per-user OAuth credentials to the authenticated LiteLLM caller Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 20 +- .../mcp_server/mcp_server_manager.py | 2 + .../mcp_server/oauth_identity_binding.py | 249 ++++++++++++++++++ .../mcp_management_endpoints.py | 22 ++ .../types/mcp_server/mcp_server_manager.py | 21 ++ .../mcp_server/test_oauth_identity_binding.py | 240 +++++++++++++++++ .../test_mcp_management_endpoints.py | 61 +++++ 7 files changed, 614 insertions(+), 1 deletion(-) create mode 100644 litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 93b85edd88d..370dee77083 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -54,6 +54,9 @@ from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( relative_request_url, revoke_refresh_token, ) +from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( + enforce_oauth_identity_binding, +) from litellm.proxy._experimental.mcp_server.oauth_utils import ( TOKEN_NO_CACHE_HEADERS, build_upstream_oauth2_token_request, @@ -1133,11 +1136,26 @@ async def exchange_token_with_server( server_id=resolved_server.server_id, ) + # Bind the exchanged token to the LiteLLM caller BEFORE it is returned, stored, or cached, so a + # token minted for a different upstream principal never becomes usable under the caller's user_id. + resolved_user_id: Final = ( + await _extract_user_id_from_request(request) + if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None + else None + ) + if isinstance(token_response, dict): + await enforce_oauth_identity_binding( + server=resolved_server, + token_response=token_response, + litellm_user_id=resolved_user_id, + grant_type=grant_type, + ) + # Store server-side when the server is configured for per-user OAuth and # the calling client has provided a valid LiteLLM identity. # Errors are non-fatal: the token is still returned to the client. if resolved_server.needs_user_oauth_token: - user_id: Final = await _extract_user_id_from_request(request) + user_id: Final = resolved_user_id if user_id: try: await _store_per_user_token_server_side( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 2330120adad..1d61aba8bc5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2174,6 +2174,8 @@ class MCPServerManager: allow_elicitation=bool(server_config.get("allow_elicitation", False)), timeout=server_config.get("timeout", None), max_concurrent_requests=server_config.get("max_concurrent_requests", None), + token_validation=server_config.get("token_validation", None), + oauth_identity_binding=server_config.get("oauth_identity_binding", None), ) self._assign_unique_short_prefix(new_server) _warn_internal_delegate_pkce_if_applicable(new_server, source="config") diff --git a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py new file mode 100644 index 00000000000..f132f59ae99 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py @@ -0,0 +1,249 @@ +"""Per-user OAuth identity binding: verify the upstream OIDC principal matches the LiteLLM caller. + +Closes the confused-deputy gap where a browser authenticated upstream as one principal produces a +token that the relay stores under a different, LiteLLM-authenticated principal: before the token +endpoint returns, stores, or caches an exchanged token for an identity-bound server, the id_token +is validated (signature via the pinned issuer's JWKS, issuer, audience, expiry) and its principal +claim is compared to the caller's trusted LiteLLM identity. Mismatches fail closed in enforce mode +and are logged in audit mode. +""" + +from collections.abc import Awaitable, Callable, Mapping +from dataclasses import dataclass +from typing import Final, Literal + +import jwt +from fastapi import HTTPException + +from litellm._logging import verbose_logger +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.types.llms.custom_http import httpxSpecialProvider +from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + +_ALLOWED_ID_TOKEN_ALGORITHMS: Final = ( + "RS256", + "RS384", + "RS512", + "ES256", + "ES384", + "ES512", + "PS256", + "PS384", + "PS512", +) +_JWKS_CACHE_TTL_SECONDS: Final = 3600 +_jwks_cache: Final = InMemoryCache(default_ttl=_JWKS_CACHE_TTL_SECONDS) + +JwksFetcher = Callable[[MCPOAuthIdentityBinding], Awaitable[list[Mapping[str, object]]]] +CallerPrincipalLoader = Callable[[str, MCPOAuthIdentityBinding], Awaitable[str | None]] + +_RejectionCode = Literal["oauth_principal_mismatch", "oauth_identity_binding_failed"] + + +@dataclass(frozen=True, slots=True) +class _BindingRejection: + code: _RejectionCode + description: str + + +async def _fetch_issuer_jwks(binding: MCPOAuthIdentityBinding) -> list[Mapping[str, object]]: + jwks_url: Final[str] = binding.jwks_url or await _discover_jwks_url(binding.issuer) + cached: Final = await _jwks_cache.async_get_cache(jwks_url) + if isinstance(cached, list): + return cached + client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) + response: Final = await client.get(jwks_url) + response.raise_for_status() + document: Final = response.json() + keys: Final = document.get("keys") if isinstance(document, dict) else None + if not isinstance(keys, list): + raise TypeError(f"JWKS document at {jwks_url} has no 'keys' array") + await _jwks_cache.async_set_cache(jwks_url, keys, ttl=_JWKS_CACHE_TTL_SECONDS) + return keys + + +async def _discover_jwks_url(issuer: str) -> str: + discovery_url: Final = f"{issuer.rstrip('/')}/.well-known/openid-configuration" + client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) + response: Final = await client.get(discovery_url) + response.raise_for_status() + metadata: Final = response.json() + jwks_uri: Final = metadata.get("jwks_uri") if isinstance(metadata, dict) else None + if not isinstance(jwks_uri, str) or not jwks_uri: + raise ValueError(f"OIDC discovery at {discovery_url} returned no jwks_uri") + return jwks_uri + + +def _select_signing_key(id_token: str, keys: list[Mapping[str, object]]) -> "jwt.PyJWK | _BindingRejection": + header: Final = jwt.get_unverified_header(id_token) + kid: Final = header.get("kid") + for key in keys: + if kid is None or key.get("kid") == kid: + return jwt.PyJWK(dict(key)) + return _BindingRejection( + code="oauth_identity_binding_failed", + description=f"id_token signing key (kid={kid!r}) not found in the issuer's JWKS", + ) + + +def _decode_id_token( + id_token: str, + binding: MCPOAuthIdentityBinding, + signing_key: "jwt.PyJWK", +) -> "Mapping[str, object] | _BindingRejection": + try: + return jwt.decode( + id_token, + signing_key.key, + algorithms=list(_ALLOWED_ID_TOKEN_ALGORITHMS), + issuer=binding.issuer, + audience=binding.audiences if binding.audiences else None, + options={ + "require": ["iss", "exp"], + "verify_aud": bool(binding.audiences), + }, + ) + except jwt.InvalidTokenError as exc: + return _BindingRejection( + code="oauth_identity_binding_failed", + description=f"id_token validation failed: {exc}", + ) + + +def _upstream_principal( + claims: Mapping[str, object], + binding: MCPOAuthIdentityBinding, +) -> "str | _BindingRejection": + principal: Final = claims.get(binding.principal_claim) + if not isinstance(principal, str) or not principal: + return _BindingRejection( + code="oauth_identity_binding_failed", + description=f"id_token has no usable '{binding.principal_claim}' claim", + ) + if binding.principal_claim == "email" and binding.require_email_verified and claims.get("email_verified") is not True: + return _BindingRejection( + code="oauth_identity_binding_failed", + description="id_token email is not verified (email_verified is not true)", + ) + return principal + + +async def _load_caller_principal(litellm_user_id: str, binding: MCPOAuthIdentityBinding) -> str | None: + if binding.caller_field == "user_id": + return litellm_user_id + from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( # noqa: PLC0415 # inline import avoids a module-load circular import + load_active_user_by_id, + ) + + loaded: Final = await load_active_user_by_id(litellm_user_id) + if isinstance(loaded, str): + return None + return loaded.user_email + + +def _principals_match(upstream: str, caller: str, binding: MCPOAuthIdentityBinding) -> bool: + if binding.principal_claim == "email" or binding.caller_field == "user_email": + return upstream.strip().casefold() == caller.strip().casefold() + return upstream == caller + + +async def _evaluate_binding( + binding: MCPOAuthIdentityBinding, + token_response: Mapping[str, object], + litellm_user_id: str | None, + grant_type: str, + jwks_fetcher: JwksFetcher, + caller_principal_loader: CallerPrincipalLoader, +) -> _BindingRejection | None: + id_token: Final = token_response.get("id_token") + if not isinstance(id_token, str) or not id_token: + if grant_type == "refresh_token": + return None + return _BindingRejection( + code="oauth_identity_binding_failed", + description="the upstream token response carries no id_token to bind the credential to a principal", + ) + if not litellm_user_id: + return _BindingRejection( + code="oauth_identity_binding_failed", + description="the request carries no resolvable LiteLLM user identity to bind the credential to", + ) + try: + keys: Final = await jwks_fetcher(binding) + except Exception as exc: # noqa: BLE001 # a JWKS fetch failure must fail closed, not surface as a 500 + return _BindingRejection( + code="oauth_identity_binding_failed", + description=f"could not fetch the issuer's JWKS: {exc}", + ) + signing_key: Final = _select_signing_key(id_token, keys) + if isinstance(signing_key, _BindingRejection): + return signing_key + claims: Final = _decode_id_token(id_token, binding, signing_key) + if isinstance(claims, _BindingRejection): + return claims + upstream: Final = _upstream_principal(claims, binding) + if isinstance(upstream, _BindingRejection): + return upstream + caller: Final = await caller_principal_loader(litellm_user_id, binding) + if not caller: + return _BindingRejection( + code="oauth_identity_binding_failed", + description=f"the LiteLLM user has no '{binding.caller_field}' to compare the upstream principal against", + ) + if not _principals_match(upstream, caller, binding): + return _BindingRejection( + code="oauth_principal_mismatch", + description="The browser account does not match the selected credential owner.", + ) + return None + + +async def enforce_oauth_identity_binding( + server: MCPServer, + token_response: Mapping[str, object], + litellm_user_id: str | None, + grant_type: str, + jwks_fetcher: JwksFetcher = _fetch_issuer_jwks, + caller_principal_loader: CallerPrincipalLoader = _load_caller_principal, +) -> None: + """Validate the exchanged token's upstream principal against the LiteLLM caller. + + No-op when the server has no binding or it is disabled. In enforce mode a failure raises 403 + before the caller returns, stores, or caches the token; in audit mode failures are logged only. + A refresh_token grant without an id_token is allowed in both modes: the stored credential keeps + the binding established at the original authorization_code exchange. + """ + binding: Final = server.oauth_identity_binding + if binding is None or binding.mode == "disabled": + return + rejection: Final = await _evaluate_binding( + binding=binding, + token_response=token_response, + litellm_user_id=litellm_user_id, + grant_type=grant_type, + jwks_fetcher=jwks_fetcher, + caller_principal_loader=caller_principal_loader, + ) + if rejection is None: + return + if binding.mode == "audit": + verbose_logger.warning( + "oauth_identity_binding audit: server=%s user=%s grant=%s rejected=%s (%s)", + server.server_id, + litellm_user_id, + grant_type, + rejection.code, + rejection.description, + ) + return + raise HTTPException( + status_code=403, + detail={ + "error": rejection.code, + "error_description": rejection.description, + "server_id": server.server_id, + "credential_owner": "caller", + "credential_stored": False, + }, + ) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 556a30d0b29..5e68e2440da 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -2135,6 +2135,28 @@ if MCP_AVAILABLE: """Persist the OAuth2 access token obtained by the calling user.""" prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") await _authorize_and_fetch_mcp_server(prisma_client, user_api_key_dict, server_id) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 + global_mcp_server_manager as _manager, + ) + + # This endpoint accepts an opaque token with no upstream identity validation, so it must be + # closed for identity-bound servers or it becomes a bypass of the token-relay binding check. + registry_server: Final = _manager.get_mcp_server_by_id(server_id) + binding: Final = registry_server.oauth_identity_binding if registry_server else None + if binding is not None and binding.mode == "enforce": + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": "oauth_identity_binding_enforced", + "error_description": ( + "Direct credential storage is disabled for this server: its OAuth identity " + "binding is enforced and this endpoint cannot validate the token's principal. " + "Complete the OAuth flow through the gateway instead." + ), + "server_id": server_id, + "credential_stored": False, + }, + ) user_id: Final = user_api_key_dict.user_id or "" if not user_id: raise HTTPException( diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 9bf3acc601c..cc3f6c8e8d8 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -38,6 +38,26 @@ class MCPOAuthMetadata(BaseModel): usable in memory but must never be persisted as configuration.""" +class MCPOAuthIdentityBinding(BaseModel): + """Per-server policy binding stored per-user OAuth credentials to the authenticated LiteLLM caller. + + When enabled for an interactive oauth2 server, the token relay validates the upstream OIDC + ``id_token`` (signature via the pinned issuer's JWKS, issuer, audience, expiry) and compares its + principal claim to the LiteLLM caller's trusted identity before the token is returned, stored, + or cached. ``audit`` logs mismatches without changing behavior; ``enforce`` fails closed with + 403 ``oauth_principal_mismatch`` and disables the direct ``oauth-user-credential`` POST, which + would otherwise bypass validation with an arbitrary opaque token. + """ + + mode: Literal["disabled", "audit", "enforce"] = "disabled" + issuer: str + jwks_url: str | None = None + audiences: list[str] = [] + principal_claim: str = "email" + caller_field: Literal["user_email", "user_id"] = "user_email" + require_email_verified: bool = True + + class MCPServer(BaseModel): server_id: str name: str @@ -172,6 +192,7 @@ class MCPServer(BaseModel): # response (supports dot-notation for nested fields, e.g. "team.enterprise_id"). # Tokens that fail validation are rejected before storage. token_validation: dict[str, Any] | None = None + oauth_identity_binding: MCPOAuthIdentityBinding | None = None # Optional TTL override (seconds) for the Redis per-user token cache, capped # at the token's expires_in minus the expiry buffer so a cached entry never # outlives the token. Defaults to the token's expires_in minus the expiry diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py new file mode 100644 index 00000000000..69eaabeda4a --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py @@ -0,0 +1,240 @@ +import time +from collections.abc import Mapping +from typing import Final + +import jwt +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from fastapi import HTTPException + +from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( + enforce_oauth_identity_binding, +) +from litellm.types.mcp import MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + +ISSUER: Final = "https://idp.example.com" +AUDIENCE: Final = "litellm-client" +KID: Final = "test-key" + +_PRIVATE_KEY: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) +_PRIVATE_PEM: Final = _PRIVATE_KEY.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), +) +_PUBLIC_JWK: Final = { + **jwt.algorithms.RSAAlgorithm.to_jwk(_PRIVATE_KEY.public_key(), as_dict=True), + "kid": KID, + "alg": "RS256", + "use": "sig", +} + + +def _sign_id_token(claims: Mapping[str, object]) -> str: + payload: Final = { + "iss": ISSUER, + "aud": AUDIENCE, + "exp": int(time.time()) + 300, + "iat": int(time.time()), + **claims, + } + return jwt.encode(payload, _PRIVATE_PEM, algorithm="RS256", headers={"kid": KID}) + + +async def _jwks_fetcher(_binding: MCPOAuthIdentityBinding) -> list[Mapping[str, object]]: + return [_PUBLIC_JWK] + + +def _caller_loader(email: str | None): + async def load(_user_id: str, _binding: MCPOAuthIdentityBinding) -> str | None: + return email + + return load + + +def _server(mode: str = "enforce", **binding_overrides: object) -> MCPServer: + return MCPServer( + server_id="srv-1", + name="srv-1", + url="https://mcp.example.com", + transport=MCPTransport.http, + oauth_identity_binding=MCPOAuthIdentityBinding( + mode=mode, + issuer=ISSUER, + audiences=[AUDIENCE], + **binding_overrides, + ), + ) + + +@pytest.mark.asyncio +async def test_matching_principal_passes(): + token: Final = _sign_id_token({"email": "Alice@Example.com", "email_verified": True}) + result: Final = await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert result is None + + +@pytest.mark.asyncio +async def test_mismatched_principal_rejected(): + token: Final = _sign_id_token({"email": "mallory@example.com", "email_verified": True}) + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "oauth_principal_mismatch" + assert exc_info.value.detail["credential_stored"] is False + + +@pytest.mark.asyncio +async def test_missing_id_token_rejected_on_authorization_code(): + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at"}, + litellm_user_id="user-a", + grant_type="authorization_code", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + + +@pytest.mark.asyncio +async def test_refresh_without_id_token_allowed(): + result: Final = await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at"}, + litellm_user_id="user-a", + grant_type="refresh_token", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert result is None + + +@pytest.mark.asyncio +async def test_refresh_with_mismatched_id_token_rejected(): + token: Final = _sign_id_token({"email": "mallory@example.com", "email_verified": True}) + with pytest.raises(HTTPException): + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="refresh_token", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + + +@pytest.mark.asyncio +async def test_audit_mode_logs_but_does_not_reject(): + token: Final = _sign_id_token({"email": "mallory@example.com", "email_verified": True}) + result: Final = await enforce_oauth_identity_binding( + server=_server(mode="audit"), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert result is None + + +@pytest.mark.asyncio +async def test_unverified_email_rejected(): + token: Final = _sign_id_token({"email": "alice@example.com", "email_verified": False}) + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + + +@pytest.mark.asyncio +async def test_wrong_issuer_rejected(): + payload: Final = { + "iss": "https://evil.example.com", + "aud": AUDIENCE, + "exp": int(time.time()) + 300, + "email": "alice@example.com", + "email_verified": True, + } + token: Final = jwt.encode(payload, _PRIVATE_PEM, algorithm="RS256", headers={"kid": KID}) + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + + +@pytest.mark.asyncio +async def test_no_litellm_identity_rejected(): + token: Final = _sign_id_token({"email": "alice@example.com", "email_verified": True}) + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id=None, + grant_type="authorization_code", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_disabled_binding_is_noop(): + result: Final = await enforce_oauth_identity_binding( + server=_server(mode="disabled"), + token_response={"access_token": "at"}, + litellm_user_id=None, + grant_type="authorization_code", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader(None), + ) + assert result is None + + +@pytest.mark.asyncio +async def test_no_binding_is_noop(): + server: Final = MCPServer( + server_id="srv-2", + name="srv-2", + url="https://mcp.example.com", + transport=MCPTransport.http, + ) + result: Final = await enforce_oauth_identity_binding( + server=server, + token_response={"access_token": "at"}, + litellm_user_id=None, + grant_type="authorization_code", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader(None), + ) + assert result is None diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index ceb44de5576..1cb36918f11 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -4638,6 +4638,67 @@ async def test_store_mcp_oauth_user_credential_returns_status(): assert result.expires_at == "2099-01-01T00:00:00+00:00" +@pytest.mark.asyncio +async def test_store_mcp_oauth_user_credential_blocked_when_identity_binding_enforced(): + """The direct opaque-token POST must be closed for enforce-mode identity-bound servers, + otherwise it bypasses the token-relay principal check.""" + from litellm.proxy._types import MCPOAuthUserCredentialRequest + from litellm.types.mcp import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + store_mcp_oauth_user_credential, + ) + + server_id = "srv-binding-1" + bound_server = MCPServer( + server_id=server_id, + name=server_id, + url="https://mcp.example.com", + transport=MCPTransport.http, + oauth_identity_binding=MCPOAuthIdentityBinding(mode="enforce", issuer="https://idp.example.com"), + ) + store_mock = AsyncMock(return_value=None) + + with ( + patch( # test-quality-ok: mirrors the existing store-credential tests in this file + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( # test-quality-ok: mirrors the existing store-credential tests in this file + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server", + new=AsyncMock(return_value=generate_mock_mcp_server_db_record(server_id=server_id)), + ), + patch( # test-quality-ok: mirrors the existing store-credential tests in this file + "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + return_value=True, + ), + patch.object( # test-quality-ok: registry is a module-level singleton; injecting it would change the endpoint signature + manager_module.global_mcp_server_manager, + "get_mcp_server_by_id", + return_value=bound_server, + ), + patch( # test-quality-ok: asserting the DB write is never reached is the point of the test + "litellm.proxy.management_endpoints.mcp_management_endpoints.store_user_oauth_credential", + new=store_mock, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await store_mcp_oauth_user_credential( + server_id=server_id, + payload=MCPOAuthUserCredentialRequest(access_token="opaque-tok", expires_in=3600), + user_api_key_dict=_make_user_auth("user-123"), + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "oauth_identity_binding_enforced" + store_mock.assert_not_called() + + @pytest.mark.asyncio async def test_delete_mcp_oauth_user_credential_only_deletes_oauth(): """delete_mcp_oauth_user_credential only deletes OAuth2 credentials, not BYOK.""" From 997a813ab99d6bd4a6b8d52fd9b42e5945c409d4 Mon Sep 17 00:00:00 2001 From: Wolfram Ravenwolf Date: Tue, 1 Sep 2026 23:15:18 +0200 Subject: [PATCH 16/60] fix(wandb): preserve reasoning_effort in chat completions --- litellm/llms/wandb/chat/transformation.py | 3 + .../wandb/test_wandb_chat_transformation.py | 55 ++++++++++++++++++- 2 files changed, 57 insertions(+), 1 deletion(-) diff --git a/litellm/llms/wandb/chat/transformation.py b/litellm/llms/wandb/chat/transformation.py index a477c00f643..cb9ac7b4c9b 100644 --- a/litellm/llms/wandb/chat/transformation.py +++ b/litellm/llms/wandb/chat/transformation.py @@ -10,6 +10,9 @@ from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig class WandbConfig(OpenAIGPTConfig): + def get_supported_openai_params(self, model: str) -> list: # mutable-ok: inherited contract + return super().get_supported_openai_params(model) + ["reasoning_effort"] # mutable-ok: inherited contract + def map_openai_params( self, non_default_params: dict, diff --git a/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py b/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py index a5d1eccebe0..4fd66448009 100644 --- a/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py +++ b/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py @@ -5,7 +5,7 @@ These tests validate the WandbInferenceConfig class which extends OpenAIGPTConfi Nebius AI Studio is an OpenAI-compatible provider with minor customizations. """ - +import json import pytest @@ -17,6 +17,20 @@ from litellm.llms.wandb.chat.transformation import WandbConfig class TestWandbConfig: """Test class for WandB Inference functionality""" + def test_map_openai_params_preserves_reasoning_effort(self): + supported_params = litellm.get_supported_openai_params(model="wandb/openai/gpt-oss-20b") + assert supported_params is not None + assert "reasoning_effort" in supported_params + + result = WandbConfig().map_openai_params( + non_default_params={"reasoning_effort": "medium"}, + optional_params={}, + model="openai/gpt-oss-20b", + drop_params=True, + ) + + assert result == {"reasoning_effort": "medium"} + def test_default_api_base(self): """Test that default API base is used when none is provided""" config = WandbConfig() @@ -139,3 +153,42 @@ class TestWandbConfig: # Check for specific content in the response assert "```python" in content assert "Hey from LiteLLM" in content + + @pytest.mark.respx() + def test_wandb_completion_preserves_reasoning_effort_with_drop_params(self, respx_mock, monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + api_base = "https://api.inference.wandb.ai/v1" + request_mock = respx_mock.post(f"{api_base}/chat/completions").respond( + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "gpt-oss-20b", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Done"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + }, + status_code=200, + ) + + completion( + model="wandb/openai/gpt-oss-20b", + messages=[{"role": "user", "content": "Hello"}], + api_key="fake-wandb-key", + api_base=api_base, + reasoning_effort="medium", + drop_params=True, + ) + + request_body = json.loads(request_mock.calls[0].request.content) + assert request_body["reasoning_effort"] == "medium" From 1bf9418917af093d4d3d66885e1c5baa956995cb Mon Sep 17 00:00:00 2001 From: Wolfram Ravenwolf Date: Wed, 2 Sep 2026 02:48:16 +0200 Subject: [PATCH 17/60] fix(wandb): type supported parameter list --- litellm/llms/wandb/chat/transformation.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/llms/wandb/chat/transformation.py b/litellm/llms/wandb/chat/transformation.py index cb9ac7b4c9b..9c3ff818331 100644 --- a/litellm/llms/wandb/chat/transformation.py +++ b/litellm/llms/wandb/chat/transformation.py @@ -10,7 +10,7 @@ from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig class WandbConfig(OpenAIGPTConfig): - def get_supported_openai_params(self, model: str) -> list: # mutable-ok: inherited contract + def get_supported_openai_params(self, model: str) -> list[str]: # mutable-ok: inherited contract return super().get_supported_openai_params(model) + ["reasoning_effort"] # mutable-ok: inherited contract def map_openai_params( From c60c60e6fe4a860377d9fcd9e611d71234b91e36 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Wed, 2 Sep 2026 10:35:49 +0200 Subject: [PATCH 18/60] fix(proxy): resolve /v1/models limits from the deployment, not the alias --- litellm/proxy/utils.py | 108 +++++++++++---- litellm/router.py | 42 ++++-- litellm/types/router.py | 17 +++ litellm/utils.py | 27 ++-- .../test_model_management_endpoints.py | 4 +- .../proxy/proxy_server/test_routes_models.py | 8 +- .../test_team_model_name_translation.py | 12 +- tests/test_litellm/proxy/test_proxy_utils.py | 127 +++++++++++++++++- tests/test_litellm/test_router.py | 99 ++++++++++++++ 9 files changed, 388 insertions(+), 56 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 051d36c4d0f..b86e6d857cc 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -7481,6 +7481,64 @@ async def get_available_models_for_user( return all_models +def _safe_get_model_info(model: str, get_model_info: Callable[[str], ModelInfo]) -> ModelInfo | None: + try: + return get_model_info(model) + except Exception as e: + verbose_proxy_logger.debug( + "create_model_info_response: cost map lookup failed for %s: %s", + model, + e, + ) + return None + + +def _resolve_listing_model_info( + deployment_model: str | None, + listed_model: str, + get_model_info: Callable[[str], ModelInfo], +) -> tuple[ModelInfo, ...]: + """ + Cost-map entries describing a listed model, best source first. + + The name a model is listed under is an arbitrary public alias, so it often misses the + cost map and lands on a fallback-generalization rule that answers with a conservative + family baseline instead of the real model's limits; the deployment's underlying model + is what the request actually reaches. Both names are kept because either can + generalize, and because a deployment's own model is registered into the cost map as a + stub that carries no limits of its own. Exact entries are consulted before generalized + ones, and each field is then taken from the first entry that has it. + """ + listed_info: Final = _safe_get_model_info(listed_model, get_model_info) + + # Fast path, and the only one a wildcard-expanded name takes: with a single name + # there is nothing to order, so skip the generalization test entirely. This keeps + # the per-model cost of the listing on the hot path #33721 exists to protect. + if deployment_model is None or deployment_model == listed_model: + return () if listed_info is None else (listed_info,) + + deployment_info: Final = _safe_get_model_info(deployment_model, get_model_info) + if deployment_info is None: + return () if listed_info is None else (listed_info,) + if listed_info is None: + return (deployment_info,) + + from litellm.utils import is_generalized_model_info + + # Both names resolved: the deployment's model leads unless it only generalized + # while the listed name is an exact cost-map entry. + if is_generalized_model_info(deployment_info) and not is_generalized_model_info(listed_info): + return (listed_info, deployment_info) + return (deployment_info, listed_info) + + +def _first_token_limit(candidates: tuple[ModelInfo, ...], field: str) -> int | None: + return next( + (limit for limit in (coerce_token_limit(info.get(field)) for info in candidates) if limit is not None), + None, + ) + + def create_model_info_response( model_id: str, provider: str, @@ -7505,31 +7563,35 @@ def create_model_info_response( "owned_by": provider, } - try: - model_cost_info: ModelInfo | None = get_model_info(model_id) - except Exception as e: - verbose_proxy_logger.debug( - "create_model_info_response: cost map lookup failed for %s: %s", - model_id, - e, - ) - model_cost_info = None + listing_info: Final = llm_router.get_model_listing_info(model_id) if llm_router is not None else None - max_input_tokens: int | None = None - max_output_tokens: int | None = None - if model_cost_info is not None: - max_input_tokens = coerce_token_limit(model_cost_info.get("max_input_tokens")) - max_output_tokens = coerce_token_limit(model_cost_info.get("max_output_tokens")) - mode: Final = model_cost_info.get("mode") - if isinstance(mode, str): - base["mode"] = mode + candidates: Final = _resolve_listing_model_info( + deployment_model=listing_info.cost_map_key if listing_info is not None else None, + listed_model=model_id, + get_model_info=get_model_info, + ) - if llm_router is not None: - configured_input, configured_output = llm_router.get_configured_token_limits(model_id) - if configured_input is not None: - max_input_tokens = configured_input - if configured_output is not None: - max_output_tokens = configured_output + max_input_tokens: int | None = _first_token_limit(candidates, "max_input_tokens") + max_output_tokens: int | None = _first_token_limit(candidates, "max_output_tokens") + mode: Final = next( + ( + m + for m in ( + cast("Mapping[str, object]", info).get("mode") # cast-ok: an entry need not carry "mode" + for info in candidates + ) + if isinstance(m, str) + ), + None, + ) + if mode is not None: + base["mode"] = mode + + if listing_info is not None: + if listing_info.max_input_tokens is not None: + max_input_tokens = listing_info.max_input_tokens + if listing_info.max_output_tokens is not None: + max_output_tokens = listing_info.max_output_tokens if max_input_tokens is not None: base["max_input_tokens"] = max_input_tokens diff --git a/litellm/router.py b/litellm/router.py index 23d8907fb49..7bc6a6c714a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -200,6 +200,7 @@ from litellm.types.router import ( CredentialLiteLLMParams, CustomRoutingStrategyBase, Deployment, + DeploymentModelListingInfo, DeploymentTypedDict, FallbackAccessCheck, GuardrailTypedDict, @@ -9726,25 +9727,44 @@ class Router: return None return Deployment(**first_usable) if isinstance(first_usable, dict) else first_usable + def get_model_listing_info(self, model_name: str) -> DeploymentModelListingInfo | None: + """ + Return what the concrete deployment behind model_name contributes to its + /v1/models entry: the cost-map key for its underlying model, plus any token + limits explicitly configured in its model_info. Resolved via O(1) index lookup. + + Returns None for wildcard-expanded or unknown names, where the listed name is + the real model name and no deployment-specific information exists, and treats a + malformed configured limit as absent rather than failing the listing. Unlike + get_model_group_info, this never triggers pattern matching or deep copies, so it + is safe to call per listed model on the /v1/models hot path. + """ + deployment: Final = self.get_deployment_by_model_group_name(model_group_name=model_name) + if deployment is None: + return None + + model_info: Final = deployment.model_info + # base_model is a declared field, so read it as one: an unset or blank value + # means the deployment's own model name is the cost-map key. + base_model: Final = model_info.base_model + return DeploymentModelListingInfo( + cost_map_key=base_model or deployment.litellm_params.model, + max_input_tokens=coerce_token_limit(model_info.get("max_input_tokens")), + max_output_tokens=coerce_token_limit(model_info.get("max_output_tokens")), + ) + def get_configured_token_limits(self, model_name: str) -> "tuple[int | None, int | None]": """ Return (max_input_tokens, max_output_tokens) explicitly configured in a concrete deployment's model_info for model_name, via O(1) index lookup. Returns (None, None) for wildcard-expanded or unknown names, and treats a - malformed configured value as absent rather than failing the listing. Unlike - get_model_group_info, this never triggers pattern matching or deep copies, so it - is safe to call per listed model on the /v1/models hot path. + malformed configured value as absent rather than failing the caller. """ - deployment: Final = self.get_deployment_by_model_group_name(model_group_name=model_name) - if deployment is None: + listing_info: Final = self.get_model_listing_info(model_name=model_name) + if listing_info is None: return (None, None) - - model_info: Final = deployment.model_info - return ( - coerce_token_limit(model_info.get("max_input_tokens")), - coerce_token_limit(model_info.get("max_output_tokens")), - ) + return (listing_info.max_input_tokens, listing_info.max_output_tokens) def get_deployment_credentials_with_provider( self, model_id: str, team_id: str | None = None diff --git a/litellm/types/router.py b/litellm/types/router.py index e0957383aac..6f9fb4bf6ee 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -575,6 +575,23 @@ class Deployment(BaseModel): setattr(self, key, value) +@dataclass(frozen=True, slots=True) +class DeploymentModelListingInfo: + """What a concrete deployment contributes to its OpenAI-compatible listing entry. + + ``cost_map_key`` is the name the deployment's underlying model is known by in + ``litellm.model_cost`` (``model_info.base_model`` when set, else + ``litellm_params.model``), which is what the request actually reaches; the public + model name it is listed under is an arbitrary alias and often absent from the cost + map. The token limits are the ones explicitly set in ``model_info``, which outrank + anything the cost map says. + """ + + cost_map_key: str + max_input_tokens: int | None + max_output_tokens: int | None + + class RouterErrors(enum.Enum): """ Enum for router specific errors with common codes diff --git a/litellm/utils.py b/litellm/utils.py index 252b6756937..80d2399c518 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2948,24 +2948,33 @@ def _resolve_builtin_model_cost_entry(key: str, provider: str) -> dict[str, obje return None +def is_generalized_model_info(model_info: ModelInfo) -> bool: + """Whether ``model_info`` came from a fallback-generalization capability rule. + + Detected as the resolved key missing ``litellm.model_cost`` while matching a + capability rule. A rule-derived entry carries no pricing and only a conservative + family-baseline context window, so callers holding a second candidate name should + prefer an exact cost-map entry from that name over this one. + """ + key: Final = cast("Mapping[str, object]", model_info).get("key") # cast-ok: partial dicts may omit "key" + if not isinstance(key, str): + return False + return key not in litellm.model_cost and match_capability_generalizations(key) is not None + + def _get_builtin_model_info_for_registration(model: str) -> ModelInfo | None: """Resolve ``model`` to its built-in cost-map entry for registration merging. Returns ``None`` when the lookup raises or when it resolved via a - fallback-generalization capability rule, detected as the resolved key missing - ``litellm.model_cost`` while matching a capability rule. A rule-derived entry - carries no pricing, so treating it as a hit would skip the built-in - cache-pricing inheritance for prefix-mangled keys. + fallback-generalization capability rule. A rule-derived entry carries no + pricing, so treating it as a hit would skip the built-in cache-pricing + inheritance for prefix-mangled keys. """ try: info: Final = get_model_info(model=model) except Exception: return None - if info["key"] in litellm.model_cost: - return info - if match_capability_generalizations(info["key"]) is None: - return info - return None + return None if is_generalized_model_info(info) else info _runtime_registered_model_cost: Final[dict[str, dict[str, object]]] = {} # mutable-ok: replayed on reload diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 4661cc17dbc..55e195ff14b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -1944,7 +1944,7 @@ class TestModelInfoEndpoint: ): mock_router.get_fully_blocked_model_names.return_value = set() mock_router.get_model_list.return_value = [] - mock_router.get_configured_token_limits.return_value = (None, None) + mock_router.get_model_listing_info.return_value = None mock_router.get_deployment_by_model_group_name.return_value = Deployment( model_name="gpt-4", litellm_params=LiteLLM_Params(model="openai/gpt-4"), @@ -2021,7 +2021,7 @@ class TestModelInfoEndpoint: ): mock_router.get_fully_blocked_model_names.return_value = set() mock_router.get_model_list.return_value = [] - mock_router.get_configured_token_limits.return_value = (None, None) + mock_router.get_model_listing_info.return_value = None mock_router.get_deployment_by_model_group_name.return_value = Deployment( model_name="team-model-1", litellm_params=LiteLLM_Params(model="custom/team-model-1"), diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_models.py b/tests/test_litellm/proxy/proxy_server/test_routes_models.py index 2b126b1ea95..9063010b366 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_models.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_models.py @@ -17,6 +17,7 @@ import litellm from litellm.proxy import proxy_server from litellm.proxy import utils as proxy_utils from litellm.proxy.utils import create_model_info_response +from litellm.types.router import DeploymentModelListingInfo from .conftest import normalize # type: ignore[import-not-found] @@ -165,13 +166,16 @@ def test_anthropic_format_carries_router_configured_token_limits(client, auth_as built from another entry's lookup shows up as the wrong numbers.""" def _configured(model_name): - return (300000, 32000) if model_name == "gpt-4" else (500000, 4096) + max_input, max_output = (300000, 32000) if model_name == "gpt-4" else (500000, 4096) + return DeploymentModelListingInfo( + cost_map_key=model_name, max_input_tokens=max_input, max_output_tokens=max_output + ) def _cost_map_lookup(model_id): max_input, max_output = (200000, 64000) if model_id == "gpt-4" else (100000, 8000) return {"max_input_tokens": max_input, "max_output_tokens": max_output, "mode": "chat"} - patched_models.get_configured_token_limits = MagicMock(side_effect=_configured) + patched_models.get_model_listing_info = MagicMock(side_effect=_configured) def _resolved(**kwargs): return create_model_info_response(**kwargs, get_model_info=_cost_map_lookup) diff --git a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py index aa35fd64f18..6c1f2eb9bcf 100644 --- a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py +++ b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py @@ -877,7 +877,7 @@ async def test_v1_models_translates_team_model_for_access_group_key(monkeypatch) router.get_model_names.return_value = ["model_name_teamX_uuid9"] router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]} router.get_fully_blocked_model_names.return_value = set() - router.get_configured_token_limits.return_value = (None, None) + router.get_model_listing_info.return_value = None router.model_list = [team_dep] router.get_model_list.return_value = [team_dep] @@ -919,7 +919,7 @@ async def test_v1_models_keeps_internal_names_when_public_name_flag_disabled( router.get_model_names.return_value = ["model_name_teamX_uuid9"] router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]} router.get_fully_blocked_model_names.return_value = set() - router.get_configured_token_limits.return_value = (None, None) + router.get_model_listing_info.return_value = None router.model_list = [team_dep] router.get_model_list.return_value = [team_dep] @@ -954,7 +954,7 @@ async def test_v1_models_translates_team_model_with_metadata(monkeypatch): router.get_model_names.return_value = ["model_name_teamX_uuid9"] router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]} router.get_fully_blocked_model_names.return_value = set() - router.get_configured_token_limits.return_value = (None, None) + router.get_model_listing_info.return_value = None router.model_list = [team_dep] router.get_model_list.return_value = [team_dep] router.get_model_group_info.return_value = None @@ -1000,7 +1000,7 @@ async def test_v1_models_metadata_fallbacks_use_internal_routing_key(monkeypatch router.get_model_names.return_value = ["model_name_teamX_uuid9"] router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]} router.get_fully_blocked_model_names.return_value = set() - router.get_configured_token_limits.return_value = (None, None) + router.get_model_listing_info.return_value = None router.model_list = [team_dep] router.get_model_list.return_value = [team_dep] # Fallbacks are keyed on the internal routing name, as the router stores them. @@ -1057,7 +1057,7 @@ async def test_v1_models_metadata_does_not_leak_other_team_fallbacks(monkeypatch router.get_model_names.return_value = ["model_name_teamX_uuid9"] router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]} router.get_fully_blocked_model_names.return_value = set() - router.get_configured_token_limits.return_value = (None, None) + router.get_model_listing_info.return_value = None router.model_list = [team_x, team_y] router.get_model_list.return_value = [team_x, team_y] router.fallbacks = [ @@ -1312,7 +1312,7 @@ def test_translate_team_model_names_for_listing_respects_legacy_flag(): def _public_named_router(*team_rows: dict) -> MagicMock: router = MagicMock() router.get_model_list.return_value = list(team_rows) - router.get_configured_token_limits.return_value = (None, None) + router.get_model_listing_info.return_value = None return router diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index dcaad968663..58f62bc4928 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -886,6 +886,7 @@ from typing import cast import litellm from litellm.proxy.utils import create_model_info_response +from litellm.types.router import DeploymentModelListingInfo from litellm.types.utils import ModelInfo @@ -913,7 +914,7 @@ def test_create_model_info_response_includes_max_tokens_from_lookup(): def test_create_model_info_response_does_not_call_router_group_info(): router = MagicMock() - router.get_configured_token_limits.return_value = (None, None) + router.get_model_listing_info.return_value = None response = create_model_info_response( model_id="some-model", @@ -928,7 +929,9 @@ def test_create_model_info_response_does_not_call_router_group_info(): def test_create_model_info_response_uses_deployment_limits_when_not_in_cost_map(): router = MagicMock() - router.get_configured_token_limits.return_value = (32000, 8000) + router.get_model_listing_info.return_value = DeploymentModelListingInfo( + cost_map_key="my-custom-deployment", max_input_tokens=32000, max_output_tokens=8000 + ) response = create_model_info_response( model_id="my-custom-deployment", @@ -944,7 +947,9 @@ def test_create_model_info_response_uses_deployment_limits_when_not_in_cost_map( def test_create_model_info_response_deployment_limits_override_cost_map(): router = MagicMock() - router.get_configured_token_limits.return_value = (200000, None) + router.get_model_listing_info.return_value = DeploymentModelListingInfo( + cost_map_key="gpt-4o", max_input_tokens=200000, max_output_tokens=None + ) response = create_model_info_response( model_id="gpt-4o", @@ -1878,3 +1883,119 @@ async def test_proxy_only_error_5xx_keeps_traceback_and_runs_sync_callbacks(monk Logging.failure_handler = orig_sync_failure assert "test_proxy_utils" in captured["async_traceback"] + + +def test_create_model_info_response_resolves_alias_to_deployment_model(): + """A public model name that is not itself a cost-map key must not be resolved through + the fallback-generalization rules: `bedrock-claude-opus-5` matches the generic + claude-family baseline (200k/64k) by substring, while the deployment it fronts really + accepts 1M/128k. Regression for the /v1/models alias resolution introduced in v1.94.0.""" + from litellm import Router + + saved_model_cost = dict(litellm.model_cost) + try: + router = Router( + model_list=[ + { + "model_name": "bedrock-claude-opus-5", + "litellm_params": { + "custom_llm_provider": "bedrock", + "model": "bedrock/eu.anthropic.claude-opus-5", + }, + "model_info": {"base_model": "eu.anthropic.claude-opus-5"}, + } + ] + ) + + response = create_model_info_response( + model_id="bedrock-claude-opus-5", provider="openai", llm_router=router + ) + finally: + litellm.model_cost.clear() + litellm.model_cost.update(saved_model_cost) + + assert response["max_input_tokens"] == 1000000 + assert response["max_output_tokens"] == 128000 + + +def test_create_model_info_response_keeps_exact_alias_over_generalized_deployment_model(): + """Mirror of the alias bug: when the deployment points at a custom backend name that + only matches a generalization rule, the listed name's exact cost-map entry is the + better answer and must win.""" + from litellm import Router + + saved_model_cost = dict(litellm.model_cost) + try: + router = Router( + model_list=[ + { + "model_name": "claude-opus-5", + "litellm_params": { + "custom_llm_provider": "bedrock", + "model": "bedrock/my-claude-opus-5-provisioned", + }, + } + ] + ) + + response = create_model_info_response( + model_id="claude-opus-5", provider="openai", llm_router=router + ) + finally: + litellm.model_cost.clear() + litellm.model_cost.update(saved_model_cost) + + assert response["max_input_tokens"] == 1000000 + + +def test_create_model_info_response_falls_back_to_alias_for_opaque_deployment_name(): + """An Azure deployment named after the resource rather than the model has no cost-map + entry; the listed name still does, and must keep answering.""" + from litellm import Router + + saved_model_cost = dict(litellm.model_cost) + try: + router = Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "azure/my-gpt4o-deployment"}, + } + ] + ) + + response = create_model_info_response( + model_id="gpt-4o", provider="openai", llm_router=router + ) + finally: + litellm.model_cost.clear() + litellm.model_cost.update(saved_model_cost) + + assert response["max_input_tokens"] == 128000 + assert response["max_output_tokens"] == 16384 + + +def test_create_model_info_response_resolves_mode_through_deployment_model(): + """`mode` is derived from the same lookup, so an aliased embedding deployment + currently reports no mode at all; it must report `embedding`.""" + from litellm import Router + + saved_model_cost = dict(litellm.model_cost) + try: + router = Router( + model_list=[ + { + "model_name": "my-embeddings", + "litellm_params": {"model": "openai/text-embedding-3-small"}, + } + ] + ) + + response = create_model_info_response( + model_id="my-embeddings", provider="openai", llm_router=router + ) + finally: + litellm.model_cost.clear() + litellm.model_cost.update(saved_model_cost) + + assert response["mode"] == "embedding" diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 84f6344be35..37805436f8d 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -7271,6 +7271,105 @@ def test_get_configured_token_limits_coerces_numeric_strings(): assert router.get_configured_token_limits("quoted-limits-model") == (32000, 8000) +def test_get_model_listing_info_prefers_base_model_over_litellm_params_model(): + """The cost-map key comes from base_model when set, so a deployment pointing at an + opaque backend name still resolves the real catalog entry.""" + router = litellm.Router( + model_list=[ + { + "model_name": "bedrock-claude-opus-5", + "litellm_params": {"model": "bedrock/eu.anthropic.claude-opus-5"}, + "model_info": {"base_model": "eu.anthropic.claude-opus-5"}, + } + ] + ) + + info = router.get_model_listing_info("bedrock-claude-opus-5") + assert info is not None + assert info.cost_map_key == "eu.anthropic.claude-opus-5" + + +def test_get_model_listing_info_falls_back_to_litellm_params_model(): + router = litellm.Router( + model_list=[ + { + "model_name": "bedrock-claude-opus-5", + "litellm_params": {"model": "bedrock/eu.anthropic.claude-opus-5"}, + } + ] + ) + + info = router.get_model_listing_info("bedrock-claude-opus-5") + assert info is not None + assert info.cost_map_key == "bedrock/eu.anthropic.claude-opus-5" + + +def test_get_model_listing_info_ignores_blank_base_model(): + """A base_model set to an empty string is absent, not a cost-map key.""" + router = litellm.Router( + model_list=[ + { + "model_name": "bedrock-claude-opus-5", + "litellm_params": {"model": "bedrock/eu.anthropic.claude-opus-5"}, + "model_info": {"base_model": ""}, + } + ] + ) + + info = router.get_model_listing_info("bedrock-claude-opus-5") + assert info is not None + assert info.cost_map_key == "bedrock/eu.anthropic.claude-opus-5" + + +def test_get_model_listing_info_returns_none_for_unknown_name(): + router = litellm.Router( + model_list=[ + { + "model_name": "no-limits-model", + "litellm_params": {"model": "openai/some-unmapped-model"}, + } + ] + ) + + assert router.get_model_listing_info("not-a-real-model") is None + + +def test_get_model_listing_info_carries_configured_limits(): + router = litellm.Router( + model_list=[ + { + "model_name": "my-custom-model", + "litellm_params": {"model": "openai/some-unmapped-model"}, + "model_info": {"max_input_tokens": 32000, "max_output_tokens": 8000}, + } + ] + ) + + info = router.get_model_listing_info("my-custom-model") + assert info is not None + assert (info.max_input_tokens, info.max_output_tokens) == (32000, 8000) + + +def test_get_model_listing_info_skips_wildcard_pattern_matching(): + router = litellm.Router( + model_list=[ + { + "model_name": "bedrock/*", + "litellm_params": {"model": "bedrock/*"}, + "model_info": {"max_input_tokens": 12345}, + } + ] + ) + + with patch.object( + router.pattern_router, "route", side_effect=AssertionError("pattern route called") + ): + assert ( + router.get_model_listing_info("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0") + is None + ) + + @pytest.mark.asyncio async def test_acreate_batch_disable_fallbacks_surfaces_owning_provider_error(): router = litellm.Router( From d5270890c409fe5942ff98c5b3a69370fc0bff81 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Wed, 2 Sep 2026 12:03:36 +0200 Subject: [PATCH 19/60] fix(proxy): report the widest window across a model group, not the first deployment's --- litellm/proxy/utils.py | 45 +++++++-- litellm/router.py | 79 +++++++++++----- litellm/types/router.py | 17 ++-- .../proxy/proxy_server/test_routes_models.py | 2 +- tests/test_litellm/proxy/test_proxy_utils.py | 52 ++++++++++- tests/test_litellm/test_router.py | 91 ++++++++++++++++++- 6 files changed, 242 insertions(+), 44 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index b86e6d857cc..874dca4238f 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -7496,10 +7496,11 @@ def _safe_get_model_info(model: str, get_model_info: Callable[[str], ModelInfo]) def _resolve_listing_model_info( deployment_model: str | None, listed_model: str, + listed_info: ModelInfo | None, get_model_info: Callable[[str], ModelInfo], ) -> tuple[ModelInfo, ...]: """ - Cost-map entries describing a listed model, best source first. + Cost-map entries describing one deployment behind a listed model, best source first. The name a model is listed under is an arbitrary public alias, so it often misses the cost map and lands on a fallback-generalization rule that answers with a conservative @@ -7508,9 +7509,10 @@ def _resolve_listing_model_info( generalize, and because a deployment's own model is registered into the cost map as a stub that carries no limits of its own. Exact entries are consulted before generalized ones, and each field is then taken from the first entry that has it. - """ - listed_info: Final = _safe_get_model_info(listed_model, get_model_info) + ``listed_info`` is resolved once by the caller, since a group with several distinct + underlying models resolves the same alias for each of them. + """ # Fast path, and the only one a wildcard-expanded name takes: with a single name # there is nothing to order, so skip the generalization test entirely. This keeps # the per-model cost of the listing on the hot path #33721 exists to protect. @@ -7539,6 +7541,20 @@ def _first_token_limit(candidates: tuple[ModelInfo, ...], field: str) -> int | N ) +def _group_token_limit(candidate_sets: tuple[tuple[ModelInfo, ...], ...], field: str) -> int | None: + """The widest limit any deployment behind the listed name declares for ``field``. + + A model group is normally one model behind several interchangeable deployments, so + there is a single value to report. When a group genuinely mixes models, reporting the + widest window keeps the listing independent of config order and agreeing with + ``/model_group/info``, which aggregates the same way for the Admin UI. + """ + limits: Final = tuple( + limit for limit in (_first_token_limit(candidates, field) for candidates in candidate_sets) if limit is not None + ) + return max(limits) if limits else None + + def create_model_info_response( model_id: str, provider: str, @@ -7565,19 +7581,30 @@ def create_model_info_response( listing_info: Final = llm_router.get_model_listing_info(model_id) if llm_router is not None else None - candidates: Final = _resolve_listing_model_info( - deployment_model=listing_info.cost_map_key if listing_info is not None else None, - listed_model=model_id, - get_model_info=get_model_info, + # One entry per distinct model behind the listed name; (None,) when the router knows + # nothing about it, so the listed name is resolved on its own as before. + deployment_models: Final[tuple[str | None, ...]] = ( + listing_info.cost_map_keys if listing_info is not None and listing_info.cost_map_keys else (None,) + ) + listed_info: Final = _safe_get_model_info(model_id, get_model_info) + candidate_sets: Final = tuple( + _resolve_listing_model_info( + deployment_model=deployment_model, + listed_model=model_id, + listed_info=listed_info, + get_model_info=get_model_info, + ) + for deployment_model in deployment_models ) - max_input_tokens: int | None = _first_token_limit(candidates, "max_input_tokens") - max_output_tokens: int | None = _first_token_limit(candidates, "max_output_tokens") + max_input_tokens: int | None = _group_token_limit(candidate_sets, "max_input_tokens") + max_output_tokens: int | None = _group_token_limit(candidate_sets, "max_output_tokens") mode: Final = next( ( m for m in ( cast("Mapping[str, object]", info).get("mode") # cast-ok: an entry need not carry "mode" + for candidates in candidate_sets for info in candidates ) if isinstance(m, str) diff --git a/litellm/router.py b/litellm/router.py index 7bc6a6c714a..7011f7f504c 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -9729,29 +9729,57 @@ class Router: def get_model_listing_info(self, model_name: str) -> DeploymentModelListingInfo | None: """ - Return what the concrete deployment behind model_name contributes to its - /v1/models entry: the cost-map key for its underlying model, plus any token - limits explicitly configured in its model_info. Resolved via O(1) index lookup. + Return what the concrete deployments behind model_name contribute to its + /v1/models entry: the cost-map keys for their underlying models, plus the widest + token limits explicitly configured in their model_info. Resolved via O(1) index + lookup. - Returns None for wildcard-expanded or unknown names, where the listed name is - the real model name and no deployment-specific information exists, and treats a - malformed configured limit as absent rather than failing the listing. Unlike - get_model_group_info, this never triggers pattern matching or deep copies, so it - is safe to call per listed model on the /v1/models hot path. + Returns None for wildcard-expanded or unknown names, where the listed name is the + real model name and no deployment-specific information exists, and treats a + malformed configured limit as absent rather than failing the listing. + + The whole group is read rather than just its first deployment, so a group that + mixes models does not advertise a window that depends on config order; the widest + one is reported, which is what get_model_group_info already shows the Admin UI. + Keys are deduplicated, so the ordinary group of interchangeable deployments of one + model still costs the caller a single cost-map lookup. Unlike get_model_group_info, + this never triggers pattern matching or deep copies, so it is safe to call per + listed model on the /v1/models hot path. """ - deployment: Final = self.get_deployment_by_model_group_name(model_group_name=model_name) - if deployment is None: + indices: Final = self.model_name_to_deployment_indices.get(model_name) + if not indices: return None - model_info: Final = deployment.model_info - # base_model is a declared field, so read it as one: an unset or blank value - # means the deployment's own model name is the cost-map key. - base_model: Final = model_info.base_model - return DeploymentModelListingInfo( - cost_map_key=base_model or deployment.litellm_params.model, - max_input_tokens=coerce_token_limit(model_info.get("max_input_tokens")), - max_output_tokens=coerce_token_limit(model_info.get("max_output_tokens")), + deployments: Final = tuple(self.model_list[index] for index in indices) + model_infos: Final = tuple(deployment.get("model_info") or MappingProxyType({}) for deployment in deployments) + params: Final = tuple(deployment.get("litellm_params") or MappingProxyType({}) for deployment in deployments) + # base_model resolution mirrors get_router_model_info: unset or blank means the + # deployment's own model name is the cost-map key. + cost_map_keys: Final = tuple( + dict.fromkeys( # deduplicates while preserving config order + key + for key in ( + model_info.get("base_model") or litellm_params.get("base_model") or litellm_params.get("model") + for model_info, litellm_params in zip(model_infos, params) + ) + if isinstance(key, str) and key + ) ) + return DeploymentModelListingInfo( + cost_map_keys=cost_map_keys, + max_input_tokens=self._widest_configured_limit(model_infos, "max_input_tokens"), + max_output_tokens=self._widest_configured_limit(model_infos, "max_output_tokens"), + ) + + @staticmethod + def _widest_configured_limit(model_infos: Sequence[Mapping[str, Any]], field: str) -> int | None: + """The largest usable value of ``field`` across a group's configured model_info blocks.""" + limits: Final = tuple( + limit + for limit in (coerce_token_limit(model_info.get(field)) for model_info in model_infos) + if limit is not None + ) + return max(limits) if limits else None def get_configured_token_limits(self, model_name: str) -> "tuple[int | None, int | None]": """ @@ -9760,11 +9788,20 @@ class Router: Returns (None, None) for wildcard-expanded or unknown names, and treats a malformed configured value as absent rather than failing the caller. + + Deliberately reads one deployment rather than aggregating the group the way + get_model_listing_info does: its caller truncates an embedding input to this + value, so the widest window in a mixed group would be the wrong answer there. """ - listing_info: Final = self.get_model_listing_info(model_name=model_name) - if listing_info is None: + deployment: Final = self.get_deployment_by_model_group_name(model_group_name=model_name) + if deployment is None: return (None, None) - return (listing_info.max_input_tokens, listing_info.max_output_tokens) + + model_info: Final = deployment.model_info + return ( + coerce_token_limit(model_info.get("max_input_tokens")), + coerce_token_limit(model_info.get("max_output_tokens")), + ) def get_deployment_credentials_with_provider( self, model_id: str, team_id: str | None = None diff --git a/litellm/types/router.py b/litellm/types/router.py index 6f9fb4bf6ee..a773959f5f1 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -577,17 +577,18 @@ class Deployment(BaseModel): @dataclass(frozen=True, slots=True) class DeploymentModelListingInfo: - """What a concrete deployment contributes to its OpenAI-compatible listing entry. + """What the deployments behind a model name contribute to its OpenAI-compatible listing entry. - ``cost_map_key`` is the name the deployment's underlying model is known by in - ``litellm.model_cost`` (``model_info.base_model`` when set, else - ``litellm_params.model``), which is what the request actually reaches; the public - model name it is listed under is an arbitrary alias and often absent from the cost - map. The token limits are the ones explicitly set in ``model_info``, which outrank - anything the cost map says. + ``cost_map_keys`` are the names those deployments' underlying models are known by in + ``litellm.model_cost`` (``base_model`` when set, else ``litellm_params.model``), which + is what a request actually reaches; the public model name they are listed under is an + arbitrary alias and often absent from the cost map. Keys are deduplicated in config + order, so the ordinary group -- several interchangeable deployments of one model -- + carries exactly one. The token limits are the widest explicitly set in any + deployment's ``model_info``, which outrank anything the cost map says. """ - cost_map_key: str + cost_map_keys: tuple[str, ...] max_input_tokens: int | None max_output_tokens: int | None diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_models.py b/tests/test_litellm/proxy/proxy_server/test_routes_models.py index 9063010b366..54768fef2ce 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_models.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_models.py @@ -168,7 +168,7 @@ def test_anthropic_format_carries_router_configured_token_limits(client, auth_as def _configured(model_name): max_input, max_output = (300000, 32000) if model_name == "gpt-4" else (500000, 4096) return DeploymentModelListingInfo( - cost_map_key=model_name, max_input_tokens=max_input, max_output_tokens=max_output + cost_map_keys=(model_name,), max_input_tokens=max_input, max_output_tokens=max_output ) def _cost_map_lookup(model_id): diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 58f62bc4928..90e7431a22d 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -930,7 +930,7 @@ def test_create_model_info_response_does_not_call_router_group_info(): def test_create_model_info_response_uses_deployment_limits_when_not_in_cost_map(): router = MagicMock() router.get_model_listing_info.return_value = DeploymentModelListingInfo( - cost_map_key="my-custom-deployment", max_input_tokens=32000, max_output_tokens=8000 + cost_map_keys=("my-custom-deployment",), max_input_tokens=32000, max_output_tokens=8000 ) response = create_model_info_response( @@ -948,7 +948,7 @@ def test_create_model_info_response_uses_deployment_limits_when_not_in_cost_map( def test_create_model_info_response_deployment_limits_override_cost_map(): router = MagicMock() router.get_model_listing_info.return_value = DeploymentModelListingInfo( - cost_map_key="gpt-4o", max_input_tokens=200000, max_output_tokens=None + cost_map_keys=("gpt-4o",), max_input_tokens=200000, max_output_tokens=None ) response = create_model_info_response( @@ -962,6 +962,54 @@ def test_create_model_info_response_deployment_limits_override_cost_map(): assert response["max_output_tokens"] == 16384 +def test_create_model_info_response_reports_widest_window_in_a_mixed_group(): + """A group mixing models advertises the widest window, not whichever is listed first.""" + limits = { + "small-model": _fake_model_info(max_input_tokens=200000, max_output_tokens=4096, mode="chat"), + "large-model": _fake_model_info(max_input_tokens=1000000, max_output_tokens=128000, mode="chat"), + } + + for keys in (("small-model", "large-model"), ("large-model", "small-model")): + router = MagicMock() + router.get_model_listing_info.return_value = DeploymentModelListingInfo( + cost_map_keys=keys, max_input_tokens=None, max_output_tokens=None + ) + + response = create_model_info_response( + model_id="house-claude", + provider="openai", + llm_router=router, + get_model_info=lambda model: limits[model], + ) + + assert response["max_input_tokens"] == 1000000, keys + assert response["max_output_tokens"] == 128000, keys + + +def test_create_model_info_response_resolves_alias_once_per_listing(): + """The alias is the same for every deployment in the group, so it is looked up once.""" + seen: list[str] = [] + + def _tracking_get_model_info(model: str) -> ModelInfo: + seen.append(model) + return _fake_model_info(max_input_tokens=128000) + + router = MagicMock() + router.get_model_listing_info.return_value = DeploymentModelListingInfo( + cost_map_keys=("model-a", "model-b"), max_input_tokens=None, max_output_tokens=None + ) + + create_model_info_response( + model_id="house-model", + provider="openai", + llm_router=router, + get_model_info=_tracking_get_model_info, + ) + + assert seen.count("house-model") == 1 + assert sorted(seen) == ["house-model", "model-a", "model-b"] + + def test_create_model_info_response_survives_malformed_configured_limits(): from litellm import Router diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 37805436f8d..588fd98abd1 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -7286,7 +7286,7 @@ def test_get_model_listing_info_prefers_base_model_over_litellm_params_model(): info = router.get_model_listing_info("bedrock-claude-opus-5") assert info is not None - assert info.cost_map_key == "eu.anthropic.claude-opus-5" + assert info.cost_map_keys == ("eu.anthropic.claude-opus-5",) def test_get_model_listing_info_falls_back_to_litellm_params_model(): @@ -7301,7 +7301,7 @@ def test_get_model_listing_info_falls_back_to_litellm_params_model(): info = router.get_model_listing_info("bedrock-claude-opus-5") assert info is not None - assert info.cost_map_key == "bedrock/eu.anthropic.claude-opus-5" + assert info.cost_map_keys == ("bedrock/eu.anthropic.claude-opus-5",) def test_get_model_listing_info_ignores_blank_base_model(): @@ -7318,7 +7318,7 @@ def test_get_model_listing_info_ignores_blank_base_model(): info = router.get_model_listing_info("bedrock-claude-opus-5") assert info is not None - assert info.cost_map_key == "bedrock/eu.anthropic.claude-opus-5" + assert info.cost_map_keys == ("bedrock/eu.anthropic.claude-opus-5",) def test_get_model_listing_info_returns_none_for_unknown_name(): @@ -7350,6 +7350,91 @@ def test_get_model_listing_info_carries_configured_limits(): assert (info.max_input_tokens, info.max_output_tokens) == (32000, 8000) +def test_get_model_listing_info_dedupes_interchangeable_deployments(): + """The ordinary group is N deployments of one model, so it yields exactly one key.""" + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-a"}, + }, + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-b"}, + }, + ] + ) + + info = router.get_model_listing_info("gpt-4o") + assert info is not None + assert info.cost_map_keys == ("openai/gpt-4o",) + + +def test_get_model_listing_info_collects_every_model_in_a_mixed_group(): + router = litellm.Router( + model_list=[ + { + "model_name": "house-claude", + "litellm_params": {"model": "anthropic/claude-3-haiku-20240307"}, + }, + { + "model_name": "house-claude", + "litellm_params": {"model": "bedrock/eu.anthropic.claude-opus-5"}, + }, + ] + ) + + info = router.get_model_listing_info("house-claude") + assert info is not None + assert info.cost_map_keys == ( + "anthropic/claude-3-haiku-20240307", + "bedrock/eu.anthropic.claude-opus-5", + ) + + +def test_get_model_listing_info_reports_widest_configured_limits_in_a_mixed_group(): + """Matches how get_model_group_info aggregates for the Admin UI, so the two agree.""" + router = litellm.Router( + model_list=[ + { + "model_name": "house-model", + "litellm_params": {"model": "openai/some-unmapped-model"}, + "model_info": {"max_input_tokens": 32000, "max_output_tokens": 4096}, + }, + { + "model_name": "house-model", + "litellm_params": {"model": "openai/another-unmapped-model"}, + "model_info": {"max_input_tokens": 128000, "max_output_tokens": 16384}, + }, + ] + ) + + info = router.get_model_listing_info("house-model") + assert info is not None + assert (info.max_input_tokens, info.max_output_tokens) == (128000, 16384) + + +def test_get_model_listing_info_reads_base_model_from_litellm_params(): + """base_model resolution mirrors get_router_model_info, which also accepts it there.""" + router = litellm.Router( + model_list=[ + { + "model_name": "azure-deployment", + "litellm_params": { + "model": "azure/my-azure-deployment-name", + "base_model": "azure/gpt-4o", + "api_key": "sk-a", + "api_base": "https://example.openai.azure.com", + }, + } + ] + ) + + info = router.get_model_listing_info("azure-deployment") + assert info is not None + assert info.cost_map_keys == ("azure/gpt-4o",) + + def test_get_model_listing_info_skips_wildcard_pattern_matching(): router = litellm.Router( model_list=[ From f0f6ff8de4c5bd09b092206003f1f1aa00b7bb41 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Wed, 2 Sep 2026 14:12:23 +0200 Subject: [PATCH 20/60] test(router): cover _widest_configured_limit directly for the coverage gate --- litellm/proxy/utils.py | 14 +++++++++++--- tests/test_litellm/test_router.py | 14 ++++++++++++++ 2 files changed, 25 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 874dca4238f..f891c85320a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -7545,9 +7545,17 @@ def _group_token_limit(candidate_sets: tuple[tuple[ModelInfo, ...], ...], field: """The widest limit any deployment behind the listed name declares for ``field``. A model group is normally one model behind several interchangeable deployments, so - there is a single value to report. When a group genuinely mixes models, reporting the - widest window keeps the listing independent of config order and agreeing with - ``/model_group/info``, which aggregates the same way for the Admin UI. + there is a single value to report and the choice of aggregate does not arise. + + When a group genuinely mixes models no single number is right, and the widest is the + deliberate pick over the narrowest for two reasons. It is what ``/model_group/info`` + has long reported to the Admin UI, so the two surfaces agree; disagreeing is the very + complaint this resolution path exists to fix. And of the two ways to be wrong, + under-advertising is worse: a client that trusts a narrowed window silently refuses + prompts the group would have served, while an over-long prompt that reaches a smaller + deployment comes back as a legible context-length error -- and does not reach one at + all when ``enable_pre_call_checks`` is set, which filters deployments the prompt does + not fit. """ limits: Final = tuple( limit for limit in (_first_token_limit(candidates, field) for candidates in candidate_sets) if limit is not None diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 588fd98abd1..eba9a85feb6 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -7350,6 +7350,20 @@ def test_get_model_listing_info_carries_configured_limits(): assert (info.max_input_tokens, info.max_output_tokens) == (32000, 8000) +def test_widest_configured_limit_ignores_absent_and_malformed_values(): + model_infos = ( + {"max_input_tokens": 32000}, + {}, + {"max_input_tokens": "not-a-number"}, + {"max_input_tokens": "128000"}, + {"max_output_tokens": 4096}, + ) + + assert litellm.Router._widest_configured_limit(model_infos, "max_input_tokens") == 128000 + assert litellm.Router._widest_configured_limit(model_infos, "max_output_tokens") == 4096 + assert litellm.Router._widest_configured_limit((), "max_input_tokens") is None + + def test_get_model_listing_info_dedupes_interchangeable_deployments(): """The ordinary group is N deployments of one model, so it yields exactly one key.""" router = litellm.Router( From 50874f9fc9cdc9e925d0fc7abda06a178a303980 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 2 Sep 2026 15:52:47 +0000 Subject: [PATCH 21/60] fix(mcp): require audiences and prove refresh-token ownership for identity binding Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 9 ++ .../mcp_server/oauth_identity_binding.py | 93 +++++++++++--- .../types/mcp_server/mcp_server_manager.py | 4 +- .../mcp_server/test_oauth_identity_binding.py | 117 +++++++++++++++++- .../test_mcp_management_endpoints.py | 6 +- 5 files changed, 211 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 92128a82856..88f2c47f5fc 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -56,6 +56,8 @@ from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( revoke_refresh_token, ) from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( + RefreshOwnershipProven, + RefreshTokenPresented, enforce_oauth_identity_binding, ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( @@ -1046,7 +1048,13 @@ async def exchange_token_with_server( refresh_request_scope = scope or bridge_upstream_scope if refresh_request_scope: token_data["scope"] = refresh_request_scope + refresh_ownership = ( + RefreshOwnershipProven() + if bridge_upstream_refresh is not None + else RefreshTokenPresented(upstream_refresh_token) + ) else: + refresh_ownership = None if not code: raise HTTPException( status_code=400, @@ -1145,6 +1153,7 @@ async def exchange_token_with_server( token_response=token_response, litellm_user_id=resolved_user_id, grant_type=grant_type, + refresh_ownership=refresh_ownership, ) # Store server-side when the server is configured for per-user OAuth and diff --git a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py index f132f59ae99..d4339c15e1a 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py @@ -14,6 +14,7 @@ from typing import Final, Literal import jwt from fastapi import HTTPException +from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm.caching.in_memory_cache import InMemoryCache @@ -37,6 +38,7 @@ _jwks_cache: Final = InMemoryCache(default_ttl=_JWKS_CACHE_TTL_SECONDS) JwksFetcher = Callable[[MCPOAuthIdentityBinding], Awaitable[list[Mapping[str, object]]]] CallerPrincipalLoader = Callable[[str, MCPOAuthIdentityBinding], Awaitable[str | None]] +StoredRefreshTokenLoader = Callable[[str, str], Awaitable[str | None]] _RejectionCode = Literal["oauth_principal_mismatch", "oauth_identity_binding_failed"] @@ -47,6 +49,19 @@ class _BindingRejection: description: str +@dataclass(frozen=True, slots=True) +class RefreshOwnershipProven: + """The gateway itself unwrapped the upstream refresh token from a sealed per-user envelope.""" + + +@dataclass(frozen=True, slots=True) +class RefreshTokenPresented: + refresh_token: str + + +RefreshOwnership = RefreshOwnershipProven | RefreshTokenPresented | None + + async def _fetch_issuer_jwks(binding: MCPOAuthIdentityBinding) -> list[Mapping[str, object]]: jwks_url: Final[str] = binding.jwks_url or await _discover_jwks_url(binding.issuer) cached: Final = await _jwks_cache.async_get_cache(jwks_url) @@ -98,11 +113,8 @@ def _decode_id_token( signing_key.key, algorithms=list(_ALLOWED_ID_TOKEN_ALGORITHMS), issuer=binding.issuer, - audience=binding.audiences if binding.audiences else None, - options={ - "require": ["iss", "exp"], - "verify_aud": bool(binding.audiences), - }, + audience=binding.audiences, + options={"require": ["iss", "exp"]}, ) except jwt.InvalidTokenError as exc: return _BindingRejection( @@ -121,7 +133,11 @@ def _upstream_principal( code="oauth_identity_binding_failed", description=f"id_token has no usable '{binding.principal_claim}' claim", ) - if binding.principal_claim == "email" and binding.require_email_verified and claims.get("email_verified") is not True: + if ( + binding.principal_claim == "email" + and binding.require_email_verified + and claims.get("email_verified") is not True + ): return _BindingRejection( code="oauth_identity_binding_failed", description="id_token email is not verified (email_verified is not true)", @@ -142,6 +158,26 @@ async def _load_caller_principal(litellm_user_id: str, binding: MCPOAuthIdentity return loaded.user_email +async def _load_stored_refresh_token(litellm_user_id: str, server_id: str) -> str | None: + try: + from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 + get_user_oauth_credential, + ) + from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 + + prisma_client: Final = get_prisma_client_or_throw( + "Database not connected. Cannot verify OAuth refresh token ownership." + ) + cred: Final = await get_user_oauth_credential( + prisma_client=prisma_client, + user_id=litellm_user_id, + server_id=server_id, + ) + return cred.get("refresh_token") if cred else None + except Exception: # noqa: BLE001 # a credential lookup failure must fail closed + return None + + def _principals_match(upstream: str, caller: str, binding: MCPOAuthIdentityBinding) -> bool: if binding.principal_claim == "email" or binding.caller_field == "user_email": return upstream.strip().casefold() == caller.strip().casefold() @@ -153,17 +189,41 @@ async def _evaluate_binding( token_response: Mapping[str, object], litellm_user_id: str | None, grant_type: str, + server_id: str, + refresh_ownership: RefreshOwnership, jwks_fetcher: JwksFetcher, caller_principal_loader: CallerPrincipalLoader, + stored_refresh_token_loader: StoredRefreshTokenLoader, ) -> _BindingRejection | None: id_token: Final = token_response.get("id_token") if not isinstance(id_token, str) or not id_token: - if grant_type == "refresh_token": - return None - return _BindingRejection( - code="oauth_identity_binding_failed", - description="the upstream token response carries no id_token to bind the credential to a principal", - ) + if grant_type != "refresh_token": + return _BindingRejection( + code="oauth_identity_binding_failed", + description="the upstream token response carries no id_token to bind the credential to a principal", + ) + match refresh_ownership: + case RefreshOwnershipProven(): + return None + case None: + return _BindingRejection( + code="oauth_identity_binding_failed", + description="refresh_token grant without an id_token carries no refresh token to prove ownership of", + ) + case RefreshTokenPresented(refresh_token): + if not litellm_user_id: + return _BindingRejection( + code="oauth_identity_binding_failed", + description="the request carries no resolvable LiteLLM user identity to bind the credential to", + ) + stored: Final = await stored_refresh_token_loader(litellm_user_id, server_id) + if stored is None or stored != refresh_token: + return _BindingRejection( + code="oauth_identity_binding_failed", + description="the presented refresh_token is not the caller's stored credential for this server", + ) + return None + assert_never(refresh_ownership) if not litellm_user_id: return _BindingRejection( code="oauth_identity_binding_failed", @@ -204,15 +264,17 @@ async def enforce_oauth_identity_binding( token_response: Mapping[str, object], litellm_user_id: str | None, grant_type: str, + refresh_ownership: RefreshOwnership, jwks_fetcher: JwksFetcher = _fetch_issuer_jwks, caller_principal_loader: CallerPrincipalLoader = _load_caller_principal, + stored_refresh_token_loader: StoredRefreshTokenLoader = _load_stored_refresh_token, ) -> None: """Validate the exchanged token's upstream principal against the LiteLLM caller. No-op when the server has no binding or it is disabled. In enforce mode a failure raises 403 before the caller returns, stores, or caches the token; in audit mode failures are logged only. - A refresh_token grant without an id_token is allowed in both modes: the stored credential keeps - the binding established at the original authorization_code exchange. + A refresh_token grant without an id_token is allowed only when the presented refresh token matches + the caller's stored credential or a sealed bridge envelope already proved ownership. """ binding: Final = server.oauth_identity_binding if binding is None or binding.mode == "disabled": @@ -222,8 +284,11 @@ async def enforce_oauth_identity_binding( token_response=token_response, litellm_user_id=litellm_user_id, grant_type=grant_type, + server_id=server.server_id, + refresh_ownership=refresh_ownership, jwks_fetcher=jwks_fetcher, caller_principal_loader=caller_principal_loader, + stored_refresh_token_loader=stored_refresh_token_loader, ) if rejection is None: return diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index cc3f6c8e8d8..cebc7decac4 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -1,7 +1,7 @@ from datetime import datetime from typing import Any, Final, Literal -from pydantic import BaseModel, ConfigDict, field_validator +from pydantic import BaseModel, ConfigDict, Field, field_validator from litellm.types.mcp import ( DEFAULT_SUBJECT_TOKEN_TYPE, @@ -52,7 +52,7 @@ class MCPOAuthIdentityBinding(BaseModel): mode: Literal["disabled", "audit", "enforce"] = "disabled" issuer: str jwks_url: str | None = None - audiences: list[str] = [] + audiences: list[str] = Field(min_length=1) principal_claim: str = "email" caller_field: Literal["user_email", "user_id"] = "user_email" require_email_verified: bool = True diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py index 69eaabeda4a..de59020c44c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py @@ -7,8 +7,11 @@ import pytest from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import rsa from fastapi import HTTPException +from pydantic import ValidationError from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( + RefreshOwnershipProven, + RefreshTokenPresented, enforce_oauth_identity_binding, ) from litellm.types.mcp import MCPTransport @@ -54,6 +57,13 @@ def _caller_loader(email: str | None): return load +def _stored_refresh_token_loader(refresh_token: str | None): + async def load(_user_id: str, _server_id: str) -> str | None: + return refresh_token + + return load + + def _server(mode: str = "enforce", **binding_overrides: object) -> MCPServer: return MCPServer( server_id="srv-1", @@ -77,6 +87,7 @@ async def test_matching_principal_passes(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) @@ -92,6 +103,7 @@ async def test_mismatched_principal_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) @@ -108,6 +120,7 @@ async def test_missing_id_token_rejected_on_authorization_code(): token_response={"access_token": "at"}, litellm_user_id="user-a", grant_type="authorization_code", + refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) @@ -116,18 +129,89 @@ async def test_missing_id_token_rejected_on_authorization_code(): @pytest.mark.asyncio -async def test_refresh_without_id_token_allowed(): +async def test_refresh_without_id_token_allowed_when_presented_token_matches_stored_credential(): result: Final = await enforce_oauth_identity_binding( server=_server(), token_response={"access_token": "at"}, litellm_user_id="user-a", grant_type="refresh_token", + refresh_ownership=RefreshTokenPresented("rt-1"), jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), + stored_refresh_token_loader=_stored_refresh_token_loader("rt-1"), ) assert result is None +@pytest.mark.asyncio +async def test_refresh_without_id_token_rejects_different_presented_token(): + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at"}, + litellm_user_id="user-a", + grant_type="refresh_token", + refresh_ownership=RefreshTokenPresented("rt-stolen"), + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + stored_refresh_token_loader=_stored_refresh_token_loader("rt-1"), + ) + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + + +@pytest.mark.asyncio +async def test_refresh_without_id_token_rejects_missing_stored_token(): + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at"}, + litellm_user_id="user-a", + grant_type="refresh_token", + refresh_ownership=RefreshTokenPresented("rt-1"), + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + stored_refresh_token_loader=_stored_refresh_token_loader(None), + ) + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + + +@pytest.mark.asyncio +async def test_refresh_without_id_token_passes_when_bridge_proves_ownership(): + async def fail_if_called(_user_id: str, _server_id: str) -> str | None: + raise AssertionError("stored refresh token loader should not be called") + + result: Final = await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at"}, + litellm_user_id="user-a", + grant_type="refresh_token", + refresh_ownership=RefreshOwnershipProven(), + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + stored_refresh_token_loader=fail_if_called, + ) + assert result is None + + +@pytest.mark.asyncio +async def test_refresh_without_id_token_rejects_without_ownership_proof(): + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at"}, + litellm_user_id="user-a", + grant_type="refresh_token", + refresh_ownership=None, + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + stored_refresh_token_loader=_stored_refresh_token_loader("rt-1"), + ) + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + + @pytest.mark.asyncio async def test_refresh_with_mismatched_id_token_rejected(): token: Final = _sign_id_token({"email": "mallory@example.com", "email_verified": True}) @@ -137,6 +221,7 @@ async def test_refresh_with_mismatched_id_token_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="refresh_token", + refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) @@ -150,6 +235,7 @@ async def test_audit_mode_logs_but_does_not_reject(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) @@ -165,6 +251,7 @@ async def test_unverified_email_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) @@ -187,6 +274,7 @@ async def test_wrong_issuer_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) @@ -202,6 +290,7 @@ async def test_no_litellm_identity_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id=None, grant_type="authorization_code", + refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) @@ -215,6 +304,7 @@ async def test_disabled_binding_is_noop(): token_response={"access_token": "at"}, litellm_user_id=None, grant_type="authorization_code", + refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader(None), ) @@ -234,7 +324,32 @@ async def test_no_binding_is_noop(): token_response={"access_token": "at"}, litellm_user_id=None, grant_type="authorization_code", + refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader(None), ) assert result is None + + +def test_identity_binding_requires_non_empty_audiences(): + with pytest.raises(ValidationError): + MCPOAuthIdentityBinding(mode="enforce", issuer=ISSUER, audiences=[]) + with pytest.raises(ValidationError): + MCPOAuthIdentityBinding(mode="enforce", issuer=ISSUER) + + +@pytest.mark.asyncio +async def test_wrong_audience_rejected(): + token: Final = _sign_id_token({"aud": "other-client", "email": "alice@example.com", "email_verified": True}) + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + refresh_ownership=None, + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index bfbbf816e69..f3510e7e0a8 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -4674,7 +4674,11 @@ async def test_store_mcp_oauth_user_credential_blocked_when_identity_binding_enf name=server_id, url="https://mcp.example.com", transport=MCPTransport.http, - oauth_identity_binding=MCPOAuthIdentityBinding(mode="enforce", issuer="https://idp.example.com"), + oauth_identity_binding=MCPOAuthIdentityBinding( + mode="enforce", + issuer="https://idp.example.com", + audiences=["litellm-client"], + ), ) store_mock = AsyncMock(return_value=None) From 35db36ee936643d36391eab1cb3c9f5031284b0d Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 2 Sep 2026 16:04:53 +0000 Subject: [PATCH 22/60] fix(mcp): satisfy type discipline lint gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 4 +- .../mcp_server/oauth_identity_binding.py | 39 ++++++++++++------- .../mcp_management_endpoints.py | 6 +-- .../types/mcp_server/mcp_server_manager.py | 2 +- 4 files changed, 31 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 88f2c47f5fc..5f54bbb4958 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1048,13 +1048,13 @@ async def exchange_token_with_server( refresh_request_scope = scope or bridge_upstream_scope if refresh_request_scope: token_data["scope"] = refresh_request_scope - refresh_ownership = ( + refresh_ownership = ( # rebind-ok: grant-specific branches assign one ownership value RefreshOwnershipProven() if bridge_upstream_refresh is not None else RefreshTokenPresented(upstream_refresh_token) ) else: - refresh_ownership = None + refresh_ownership = None # rebind-ok: grant-specific branches assign one ownership value if not code: raise HTTPException( status_code=400, diff --git a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py index d4339c15e1a..b73227f3413 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py @@ -8,12 +8,13 @@ claim is compared to the caller's trusted LiteLLM identity. Mismatches fail clos and are logged in audit mode. """ -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass -from typing import Final, Literal +from typing import Final, Literal, TypeAlias import jwt from fastapi import HTTPException +from jwt.types import Options from typing_extensions import assert_never from litellm._logging import verbose_logger @@ -36,11 +37,20 @@ _ALLOWED_ID_TOKEN_ALGORITHMS: Final = ( _JWKS_CACHE_TTL_SECONDS: Final = 3600 _jwks_cache: Final = InMemoryCache(default_ttl=_JWKS_CACHE_TTL_SECONDS) -JwksFetcher = Callable[[MCPOAuthIdentityBinding], Awaitable[list[Mapping[str, object]]]] -CallerPrincipalLoader = Callable[[str, MCPOAuthIdentityBinding], Awaitable[str | None]] -StoredRefreshTokenLoader = Callable[[str, str], Awaitable[str | None]] +JwksFetcher: TypeAlias = Callable[ + [MCPOAuthIdentityBinding], # mutable-ok: Callable parameter syntax requires a list + Awaitable[Sequence[Mapping[str, object]]], +] +CallerPrincipalLoader: TypeAlias = Callable[ + [str, MCPOAuthIdentityBinding], # mutable-ok: Callable parameter syntax requires a list + Awaitable[str | None], +] +StoredRefreshTokenLoader: TypeAlias = Callable[ + [str, str], # mutable-ok: Callable parameter syntax requires a list + Awaitable[str | None], +] -_RejectionCode = Literal["oauth_principal_mismatch", "oauth_identity_binding_failed"] +_RejectionCode: TypeAlias = Literal["oauth_principal_mismatch", "oauth_identity_binding_failed"] @dataclass(frozen=True, slots=True) @@ -59,10 +69,10 @@ class RefreshTokenPresented: refresh_token: str -RefreshOwnership = RefreshOwnershipProven | RefreshTokenPresented | None +RefreshOwnership: TypeAlias = RefreshOwnershipProven | RefreshTokenPresented | None -async def _fetch_issuer_jwks(binding: MCPOAuthIdentityBinding) -> list[Mapping[str, object]]: +async def _fetch_issuer_jwks(binding: MCPOAuthIdentityBinding) -> Sequence[Mapping[str, object]]: jwks_url: Final[str] = binding.jwks_url or await _discover_jwks_url(binding.issuer) cached: Final = await _jwks_cache.async_get_cache(jwks_url) if isinstance(cached, list): @@ -90,12 +100,12 @@ async def _discover_jwks_url(issuer: str) -> str: return jwks_uri -def _select_signing_key(id_token: str, keys: list[Mapping[str, object]]) -> "jwt.PyJWK | _BindingRejection": +def _select_signing_key(id_token: str, keys: Sequence[Mapping[str, object]]) -> "jwt.PyJWK | _BindingRejection": header: Final = jwt.get_unverified_header(id_token) kid: Final = header.get("kid") for key in keys: if kid is None or key.get("kid") == kid: - return jwt.PyJWK(dict(key)) + return jwt.PyJWK(dict(key)) # mutable-ok: PyJWT requires a concrete JWK dictionary return _BindingRejection( code="oauth_identity_binding_failed", description=f"id_token signing key (kid={kid!r}) not found in the issuer's JWKS", @@ -108,13 +118,14 @@ def _decode_id_token( signing_key: "jwt.PyJWK", ) -> "Mapping[str, object] | _BindingRejection": try: + decode_options: Final[Options] = {"require": ("iss", "exp")} return jwt.decode( id_token, signing_key.key, - algorithms=list(_ALLOWED_ID_TOKEN_ALGORITHMS), + algorithms=_ALLOWED_ID_TOKEN_ALGORITHMS, issuer=binding.issuer, audience=binding.audiences, - options={"require": ["iss", "exp"]}, + options=decode_options, ) except jwt.InvalidTokenError as exc: return _BindingRejection( @@ -160,10 +171,10 @@ async def _load_caller_principal(litellm_user_id: str, binding: MCPOAuthIdentity async def _load_stored_refresh_token(litellm_user_id: str, server_id: str) -> str | None: try: - from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 + from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # keep database imports lazy get_user_oauth_credential, ) - from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 + from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 # keep database imports lazy prisma_client: Final = get_prisma_client_or_throw( "Database not connected. Cannot verify OAuth refresh token ownership." diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 92a6997abff..109030a2df5 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -259,7 +259,7 @@ if MCP_AVAILABLE: if normalize_upstream_header_name(raw) is None: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail={ + detail={ # mutable-ok: FastAPI exception detail requires a JSON-serializable dictionary "error": ( f"Invalid upstream_token_header {raw!r}: must be a valid HTTP header name " "(RFC 7230 token, e.g. 'esb-oauth')" @@ -2256,7 +2256,7 @@ if MCP_AVAILABLE: """Persist the OAuth2 access token obtained by the calling user.""" prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") await _authorize_and_fetch_mcp_server(prisma_client, user_api_key_dict, server_id) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # keep manager import lazy global_mcp_server_manager as _manager, ) @@ -2267,7 +2267,7 @@ if MCP_AVAILABLE: if binding is not None and binding.mode == "enforce": raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail={ + detail={ # mutable-ok: FastAPI exception detail requires a JSON-serializable dictionary "error": "oauth_identity_binding_enforced", "error_description": ( "Direct credential storage is disabled for this server: its OAuth identity " diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index cebc7decac4..c38f1c366b1 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -52,7 +52,7 @@ class MCPOAuthIdentityBinding(BaseModel): mode: Literal["disabled", "audit", "enforce"] = "disabled" issuer: str jwks_url: str | None = None - audiences: list[str] = Field(min_length=1) + audiences: list[str] = Field(min_length=1) # mutable-ok: public Pydantic schema requires list values principal_claim: str = "email" caller_field: Literal["user_email", "user_id"] = "user_email" require_email_verified: bool = True From 22e3fcba54c2409a46e04ba81bd6885cbfc0442a Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 2 Sep 2026 16:05:37 +0000 Subject: [PATCH 23/60] fix(mcp): restore unrelated exception formatting Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/management_endpoints/mcp_management_endpoints.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 109030a2df5..fd6294fd77e 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -259,7 +259,7 @@ if MCP_AVAILABLE: if normalize_upstream_header_name(raw) is None: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail={ # mutable-ok: FastAPI exception detail requires a JSON-serializable dictionary + detail={ "error": ( f"Invalid upstream_token_header {raw!r}: must be a valid HTTP header name " "(RFC 7230 token, e.g. 'esb-oauth')" From 7ca9f28575021d026d28bb0bfb11a667ba3add77 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 2 Sep 2026 16:33:59 +0000 Subject: [PATCH 24/60] test(mcp): cover OAuth identity binding paths Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/oauth_identity_binding.py | 2 +- .../mcp_server/test_discoverable_endpoints.py | 126 ++++++++++ .../mcp_server/test_oauth_identity_binding.py | 229 ++++++++++++++++++ 3 files changed, 356 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py index b73227f3413..58ec3cfdb4f 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py @@ -234,7 +234,7 @@ async def _evaluate_binding( description="the presented refresh_token is not the caller's stored credential for this server", ) return None - assert_never(refresh_ownership) + assert_never(refresh_ownership) # pragma: no cover if not litellm_user_id: return _BindingRejection( code="oauth_identity_binding_failed", diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 598e9276423..4ae1b105dbf 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -7816,6 +7816,132 @@ async def test_token_exchange_pairs_client_secret_with_server_client_id(): assert "client_secret" not in sent +@pytest.mark.asyncio +async def test_token_exchange_refresh_passes_presented_refresh_ownership(): + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + exchange_token_with_server, + ) + from litellm.proxy._experimental.mcp_server.oauth_identity_binding import RefreshTokenPresented + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + + server = MCPServer( + server_id="srv-1", + name="srv-1", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="cid", + token_url="https://provider.example/token", + oauth_identity_binding=MCPOAuthIdentityBinding( + mode="enforce", + issuer="https://provider.example", + audiences=["cid"], + ), + ) + request = MagicMock(spec=Request) + request.base_url = "https://litellm.example.com/" + request.headers = {} + response = MagicMock() + response.json.return_value = {"access_token": "at"} + response.raise_for_status = MagicMock() + client = MagicMock() + client.post = AsyncMock(return_value=response) + enforce = AsyncMock() + + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=client, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + new=AsyncMock(return_value="user-a"), + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.enforce_oauth_identity_binding", + new=enforce, + ), + ): + await exchange_token_with_server( + request=request, + mcp_server=server, + grant_type="refresh_token", + code=None, + redirect_uri=None, + client_id="cid", + client_secret=None, + code_verifier=None, + refresh_token="rt-1", + ) + + ownership = enforce.await_args.kwargs["refresh_ownership"] + assert isinstance(ownership, RefreshTokenPresented) + assert ownership.refresh_token == "rt-1" + + +@pytest.mark.asyncio +async def test_token_exchange_authorization_code_passes_no_refresh_ownership(): + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + exchange_token_with_server, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + + server = MCPServer( + server_id="srv-1", + name="srv-1", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="cid", + token_url="https://provider.example/token", + oauth_identity_binding=MCPOAuthIdentityBinding( + mode="enforce", + issuer="https://provider.example", + audiences=["cid"], + ), + ) + request = MagicMock(spec=Request) + request.base_url = "https://litellm.example.com/" + request.headers = {} + response = MagicMock() + response.json.return_value = {"access_token": "at"} + response.raise_for_status = MagicMock() + client = MagicMock() + client.post = AsyncMock(return_value=response) + enforce = AsyncMock() + + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=client, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + new=AsyncMock(return_value="user-a"), + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.enforce_oauth_identity_binding", + new=enforce, + ), + ): + await exchange_token_with_server( + request=request, + mcp_server=server, + grant_type="authorization_code", + code="auth-code", + redirect_uri="https://litellm.example.com/callback", + client_id="cid", + client_secret=None, + code_verifier=None, + ) + + assert enforce.await_args.kwargs["refresh_ownership"] is None + + def _upstream_token_response(status_code: int, *, json_body: object = None, text_body: str = "") -> "httpx.Response": import httpx diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py index de59020c44c..c2662f9f298 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py @@ -1,6 +1,8 @@ import time from collections.abc import Mapping +from types import SimpleNamespace from typing import Final +from unittest.mock import AsyncMock, MagicMock, patch import jwt import pytest @@ -12,6 +14,11 @@ from pydantic import ValidationError from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( RefreshOwnershipProven, RefreshTokenPresented, + _discover_jwks_url, + _fetch_issuer_jwks, + _load_caller_principal, + _load_stored_refresh_token, + _select_signing_key, enforce_oauth_identity_binding, ) from litellm.types.mcp import MCPTransport @@ -79,6 +86,121 @@ def _server(mode: str = "enforce", **binding_overrides: object) -> MCPServer: ) +@pytest.mark.asyncio +async def test_fetch_issuer_jwks_fetches_and_caches_keys(): + binding: Final = _server(jwks_url="https://idp.example.com/jwks-cache").oauth_identity_binding + assert binding is not None + response: Final = MagicMock() + response.json.return_value = {"keys": [_PUBLIC_JWK]} + client: Final = MagicMock() + client.get = AsyncMock(return_value=response) + + with patch( + "litellm.proxy._experimental.mcp_server.oauth_identity_binding.get_async_httpx_client", + return_value=client, + ): + first: Final = await _fetch_issuer_jwks(binding) + second: Final = await _fetch_issuer_jwks(binding) + + assert first == second == [_PUBLIC_JWK] + client.get.assert_awaited_once_with("https://idp.example.com/jwks-cache") + + +@pytest.mark.asyncio +async def test_fetch_issuer_jwks_rejects_malformed_document(): + binding: Final = _server(jwks_url="https://idp.example.com/jwks-invalid").oauth_identity_binding + assert binding is not None + response: Final = MagicMock() + response.json.return_value = {} + client: Final = MagicMock() + client.get = AsyncMock(return_value=response) + + with patch( + "litellm.proxy._experimental.mcp_server.oauth_identity_binding.get_async_httpx_client", + return_value=client, + ): + with pytest.raises(TypeError, match="has no 'keys' array"): + await _fetch_issuer_jwks(binding) + + +@pytest.mark.asyncio +async def test_discover_jwks_url_returns_provider_uri(): + response: Final = MagicMock() + response.json.return_value = {"jwks_uri": "https://idp.example.com/jwks"} + client: Final = MagicMock() + client.get = AsyncMock(return_value=response) + + with patch( + "litellm.proxy._experimental.mcp_server.oauth_identity_binding.get_async_httpx_client", + return_value=client, + ): + result: Final = await _discover_jwks_url("https://idp.example.com/") + + assert result == "https://idp.example.com/jwks" + client.get.assert_awaited_once_with("https://idp.example.com/.well-known/openid-configuration") + + +@pytest.mark.asyncio +async def test_discover_jwks_url_rejects_missing_provider_uri(): + response: Final = MagicMock() + response.json.return_value = {} + client: Final = MagicMock() + client.get = AsyncMock(return_value=response) + + with patch( + "litellm.proxy._experimental.mcp_server.oauth_identity_binding.get_async_httpx_client", + return_value=client, + ): + with pytest.raises(ValueError, match="returned no jwks_uri"): + await _discover_jwks_url("https://idp.example.com") + + +def test_select_signing_key_returns_matching_key_or_rejection(): + token: Final = _sign_id_token({"email": "alice@example.com", "email_verified": True}) + + selected: Final = _select_signing_key(token, [_PUBLIC_JWK]) + rejected: Final = _select_signing_key(token, []) + + assert selected.__class__.__name__ == "PyJWK" + assert rejected.code == "oauth_identity_binding_failed" + + +@pytest.mark.asyncio +async def test_load_caller_principal_supports_user_id_and_database_email(): + user_id_binding: Final = _server(caller_field="user_id").oauth_identity_binding + assert user_id_binding is not None + assert await _load_caller_principal("user-a", user_id_binding) == "user-a" + + with patch( + "litellm.proxy._experimental.mcp_server.bridge_token_flow.load_active_user_by_id", + new=AsyncMock(side_effect=["no_active_key", SimpleNamespace(user_email="alice@example.com")]), + ): + assert await _load_caller_principal("user-a", _server().oauth_identity_binding) is None + assert await _load_caller_principal("user-a", _server().oauth_identity_binding) == "alice@example.com" + + +@pytest.mark.asyncio +async def test_load_stored_refresh_token_returns_credential_and_fails_closed(): + get_credential: Final = AsyncMock(return_value={"refresh_token": "rt-1"}) + with ( + patch( + "litellm.proxy._experimental.mcp_server.db.get_user_oauth_credential", + new=get_credential, + ), + patch( + "litellm.proxy.utils.get_prisma_client_or_throw", + return_value="prisma", + ), + ): + assert await _load_stored_refresh_token("user-a", "srv-1") == "rt-1" + + with patch( + "litellm.proxy.utils.get_prisma_client_or_throw", + side_effect=RuntimeError("database unavailable"), + ): + assert await _load_stored_refresh_token("user-a", "srv-1") is None + + @pytest.mark.asyncio async def test_matching_principal_passes(): token: Final = _sign_id_token({"email": "Alice@Example.com", "email_verified": True}) @@ -112,6 +234,113 @@ async def test_mismatched_principal_rejected(): assert exc_info.value.detail["credential_stored"] is False +@pytest.mark.asyncio +async def test_missing_upstream_principal_rejected(): + token: Final = _sign_id_token({"email_verified": True}) + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + refresh_ownership=None, + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + assert "no usable" in exc_info.value.detail["error_description"] + + +@pytest.mark.asyncio +async def test_jwks_fetch_failure_is_rejected(): + async def fail(_binding: MCPOAuthIdentityBinding) -> list[Mapping[str, object]]: + raise RuntimeError("jwks unavailable") + + token: Final = _sign_id_token({"email": "alice@example.com", "email_verified": True}) + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + refresh_ownership=None, + jwks_fetcher=fail, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + assert "jwks unavailable" in exc_info.value.detail["error_description"] + + +@pytest.mark.asyncio +async def test_missing_signing_key_is_rejected(): + token: Final = _sign_id_token({"email": "alice@example.com", "email_verified": True}) + + async def no_keys(_binding: MCPOAuthIdentityBinding) -> tuple[Mapping[str, object], ...]: + return () + + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + refresh_ownership=None, + jwks_fetcher=no_keys, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + assert "signing key" in exc_info.value.detail["error_description"] + + +@pytest.mark.asyncio +async def test_missing_caller_principal_is_rejected(): + token: Final = _sign_id_token({"email": "alice@example.com", "email_verified": True}) + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + refresh_ownership=None, + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader(None), + ) + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + assert "has no" in exc_info.value.detail["error_description"] + + +@pytest.mark.asyncio +async def test_refresh_without_id_token_requires_litellm_identity_for_presented_token(): + with pytest.raises(HTTPException) as exc_info: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at"}, + litellm_user_id=None, + grant_type="refresh_token", + refresh_ownership=RefreshTokenPresented("rt-1"), + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + stored_refresh_token_loader=_stored_refresh_token_loader("rt-1"), + ) + assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + assert "no resolvable LiteLLM user identity" in exc_info.value.detail["error_description"] + + +@pytest.mark.asyncio +async def test_user_id_principal_matching_uses_exact_comparison(): + token: Final = _sign_id_token({"sub": "user-a"}) + result: Final = await enforce_oauth_identity_binding( + server=_server(principal_claim="sub", caller_field="user_id"), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + refresh_ownership=None, + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("user-a"), + ) + assert result is None + + @pytest.mark.asyncio async def test_missing_id_token_rejected_on_authorization_code(): with pytest.raises(HTTPException) as exc_info: From 540e528d2190e6ed7292174e1985673a602a3018 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 2 Sep 2026 16:43:18 +0000 Subject: [PATCH 25/60] test(mcp): satisfy test quality gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/test_discoverable_endpoints.py | 16 +++++++++------- .../mcp_server/test_oauth_identity_binding.py | 16 ++++++++-------- 2 files changed, 17 insertions(+), 15 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 4ae1b105dbf..33b71f8d50f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -7851,15 +7851,15 @@ async def test_token_exchange_refresh_passes_presented_refresh_ownership(): enforce = AsyncMock() with ( - patch( + patch( # test-quality-ok: no injection seam exists for the exchange's HTTP and identity collaborators "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", return_value=client, ), - patch( + patch( # test-quality-ok: no injection seam exists for request identity extraction "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", new=AsyncMock(return_value="user-a"), ), - patch( + patch( # test-quality-ok: captures the ownership value at the exchange boundary "litellm.proxy._experimental.mcp_server.discoverable_endpoints.enforce_oauth_identity_binding", new=enforce, ), @@ -7915,20 +7915,20 @@ async def test_token_exchange_authorization_code_passes_no_refresh_ownership(): enforce = AsyncMock() with ( - patch( + patch( # test-quality-ok: no injection seam exists for the exchange's HTTP and identity collaborators "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", return_value=client, ), - patch( + patch( # test-quality-ok: no injection seam exists for request identity extraction "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", new=AsyncMock(return_value="user-a"), ), - patch( + patch( # test-quality-ok: captures the ownership value at the exchange boundary "litellm.proxy._experimental.mcp_server.discoverable_endpoints.enforce_oauth_identity_binding", new=enforce, ), ): - await exchange_token_with_server( + result = await exchange_token_with_server( request=request, mcp_server=server, grant_type="authorization_code", @@ -7939,6 +7939,8 @@ async def test_token_exchange_authorization_code_passes_no_refresh_ownership(): code_verifier=None, ) + assert result.status_code == 200 + assert json.loads(result.body)["access_token"] == "at" assert enforce.await_args.kwargs["refresh_ownership"] is None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py index c2662f9f298..36364c32ae5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py @@ -95,7 +95,7 @@ async def test_fetch_issuer_jwks_fetches_and_caches_keys(): client: Final = MagicMock() client.get = AsyncMock(return_value=response) - with patch( + with patch( # test-quality-ok: no HTTP dependency injection seam exists for JWKS fetching "litellm.proxy._experimental.mcp_server.oauth_identity_binding.get_async_httpx_client", return_value=client, ): @@ -115,7 +115,7 @@ async def test_fetch_issuer_jwks_rejects_malformed_document(): client: Final = MagicMock() client.get = AsyncMock(return_value=response) - with patch( + with patch( # test-quality-ok: no HTTP dependency injection seam exists for JWKS fetching "litellm.proxy._experimental.mcp_server.oauth_identity_binding.get_async_httpx_client", return_value=client, ): @@ -130,7 +130,7 @@ async def test_discover_jwks_url_returns_provider_uri(): client: Final = MagicMock() client.get = AsyncMock(return_value=response) - with patch( + with patch( # test-quality-ok: no HTTP dependency injection seam exists for OIDC discovery "litellm.proxy._experimental.mcp_server.oauth_identity_binding.get_async_httpx_client", return_value=client, ): @@ -147,7 +147,7 @@ async def test_discover_jwks_url_rejects_missing_provider_uri(): client: Final = MagicMock() client.get = AsyncMock(return_value=response) - with patch( + with patch( # test-quality-ok: no HTTP dependency injection seam exists for OIDC discovery "litellm.proxy._experimental.mcp_server.oauth_identity_binding.get_async_httpx_client", return_value=client, ): @@ -171,7 +171,7 @@ async def test_load_caller_principal_supports_user_id_and_database_email(): assert user_id_binding is not None assert await _load_caller_principal("user-a", user_id_binding) == "user-a" - with patch( + with patch( # test-quality-ok: caller loading is a lazy database boundary without injection "litellm.proxy._experimental.mcp_server.bridge_token_flow.load_active_user_by_id", new=AsyncMock(side_effect=["no_active_key", SimpleNamespace(user_email="alice@example.com")]), ): @@ -183,18 +183,18 @@ async def test_load_caller_principal_supports_user_id_and_database_email(): async def test_load_stored_refresh_token_returns_credential_and_fails_closed(): get_credential: Final = AsyncMock(return_value={"refresh_token": "rt-1"}) with ( - patch( + patch( # test-quality-ok: stored-token loading is a lazy database boundary without injection "litellm.proxy._experimental.mcp_server.db.get_user_oauth_credential", new=get_credential, ), - patch( + patch( # test-quality-ok: stored-token loading is a lazy database boundary without injection "litellm.proxy.utils.get_prisma_client_or_throw", return_value="prisma", ), ): assert await _load_stored_refresh_token("user-a", "srv-1") == "rt-1" - with patch( + with patch( # test-quality-ok: stored-token loading is a lazy database boundary without injection "litellm.proxy.utils.get_prisma_client_or_throw", side_effect=RuntimeError("database unavailable"), ): From eddb29d90ed45849967c2318f15c2edd3e991d8f Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 3 Sep 2026 11:33:44 +0200 Subject: [PATCH 26/60] fix(proxy): dedup latest health checks in SQL and gate the DB save per window --- litellm/constants.py | 1 + litellm/proxy/db/health_check_latest.py | 96 +++++++++ litellm/proxy/proxy_server.py | 58 +++++- litellm/proxy/utils.py | 52 +---- .../proxy/db/test_health_check_latest.py | 111 ++++++++++ .../proxy_server/test_background_health.py | 81 ++++++++ .../proxy/test_health_check_functions.py | 190 ++++++++++-------- .../test_prisma_client_health.py | 74 ++++--- 8 files changed, 502 insertions(+), 161 deletions(-) create mode 100644 litellm/proxy/db/health_check_latest.py create mode 100644 tests/test_litellm/proxy/db/test_health_check_latest.py diff --git a/litellm/constants.py b/litellm/constants.py index ef9329b9dfc..1f5bd673997 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1619,6 +1619,7 @@ CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME: Final = "cloudzero_export_usage_data" MAVVRIK_FOCUS_EXPORT_JOB_NAME: Final = "mavvrik_focus_export_usage_data" CLOUDZERO_MAX_FETCHED_DATA_RECORDS: Final = int(os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000)) SPEND_LOG_CLEANUP_JOB_NAME: Final = "spend_log_cleanup" +BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME: Final = "background_health_check_db_save" KEY_ROTATION_JOB_NAME: Final = "litellm_key_rotation_job" EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME: Final = "litellm_expired_ui_session_key_cleanup_job" WEEKLY_SPEND_REPORT_JOB_ID: Final = "weekly_spend_report_job" diff --git a/litellm/proxy/db/health_check_latest.py b/litellm/proxy/db/health_check_latest.py new file mode 100644 index 00000000000..35bc838379c --- /dev/null +++ b/litellm/proxy/db/health_check_latest.py @@ -0,0 +1,96 @@ +""" +Latest health-check row per model, deduplicated by Postgres. + +prisma-client-py's ``find_many(distinct=...)`` dedups client-side: the emitted +SQL carries no DISTINCT, so the whole append-only history table streams to the +worker on every call. ``SELECT DISTINCT ON`` keeps the transfer at one row per +(model_id, model_name) and is served by the matching descending index. +""" + +from __future__ import annotations + +import json +from collections.abc import Sequence +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Final + +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, field_validator + +from litellm._logging import verbose_proxy_logger + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +LATEST_HEALTH_CHECKS_SQL: Final = """ +SELECT DISTINCT ON ("model_id", "model_name") + "health_check_id", "model_name", "model_id", "status", + "healthy_count", "unhealthy_count", "error_message", + "response_time_ms", "details", "checked_by", + "checked_at", "created_at", "updated_at" +FROM "LiteLLM_HealthCheckTable" +ORDER BY "model_id" ASC, "model_name" ASC, "checked_at" DESC +""" + +LATEST_HEALTH_CHECKS_FOR_MODELS_SQL: Final = """ +SELECT DISTINCT ON ("model_id", "model_name") + "health_check_id", "model_name", "model_id", "status", + "healthy_count", "unhealthy_count", "error_message", + "response_time_ms", "details", "checked_by", + "checked_at", "created_at", "updated_at" +FROM "LiteLLM_HealthCheckTable" +WHERE "model_name" = ANY($1) +ORDER BY "model_id" ASC, "model_name" ASC, "checked_at" DESC +""" + + +class LatestHealthCheckRow(BaseModel): + model_config = ConfigDict(frozen=True, protected_namespaces=()) + + health_check_id: str + model_name: str + model_id: str | None = None + status: str + healthy_count: int = 0 + unhealthy_count: int = 0 + error_message: str | None = None + response_time_ms: float | None = None + details: JsonValue | None = None + checked_by: str | None = None + checked_at: datetime + created_at: datetime + updated_at: datetime + + @field_validator("details", mode="before") + @classmethod + def _decode_json_text(cls, value: object) -> object: + return json.loads(value) if isinstance(value, str) else value + + @field_validator("checked_at", "created_at", "updated_at") + @classmethod + def _assume_utc(cls, value: datetime) -> datetime: + return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value + + +_ROWS_ADAPTER: Final = TypeAdapter(tuple[LatestHealthCheckRow, ...]) + + +async def fetch_latest_health_checks(prisma_client: PrismaClient) -> tuple[LatestHealthCheckRow, ...]: + try: + rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_SQL) + return _ROWS_ADAPTER.validate_python(rows) + except Exception as query_err: # noqa: BLE001 # health decorates other reads; a driver error must not fail them + verbose_proxy_logger.error("Error getting all latest health checks: %s", query_err) + return () + + +async def fetch_latest_health_checks_for_models( + prisma_client: PrismaClient, model_names: Sequence[str] +) -> tuple[LatestHealthCheckRow, ...]: + if not model_names: + return () + try: + rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_FOR_MODELS_SQL, list(model_names)) + return _ROWS_ADAPTER.validate_python(rows) + except Exception as query_err: # noqa: BLE001 # a paged model list must not fail on its health decoration + verbose_proxy_logger.error("Error getting latest health checks for models: %s", query_err) + return () diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 27132c90e05..0e330272bdf 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -16,7 +16,16 @@ import threading import time import traceback import warnings -from collections.abc import AsyncGenerator, AsyncIterator, Callable, Collection, Mapping, MutableMapping, Sequence +from collections.abc import ( + AsyncGenerator, + AsyncIterator, + Awaitable, + Callable, + Collection, + Mapping, + MutableMapping, + Sequence, +) from datetime import datetime, timedelta, timezone from types import MappingProxyType, UnionType from typing import ( @@ -50,6 +59,7 @@ from litellm.constants import ( AIOHTTP_NEEDS_CLEANUP_CLOSED, AIOHTTP_TTL_DNS_CACHE, AUDIO_SPEECH_CHUNK_SIZE, + BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME, BASE_MCP_ROUTE, DAILY_TAG_SPEND_BATCH_MULTIPLIER, DEFAULT_MAX_RECURSE_DEPTH, @@ -224,7 +234,7 @@ def generate_feedback_box(): import contextlib from collections import defaultdict from contextlib import asynccontextmanager -from functools import lru_cache +from functools import lru_cache, partial import litellm import litellm._redis @@ -402,6 +412,7 @@ from litellm.proxy.config_resolvers.alerting import ( ) from litellm.proxy.container_endpoints.endpoints import router as container_router from litellm.proxy.credential_endpoints.endpoints import router as credential_router +from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import ( SPEND_LOG_CLEANUP_BOUND_SETTINGS, SpendLogCleanup, @@ -3580,12 +3591,35 @@ async def _run_direct_health_check_with_instrumentation( raise AssertionError("perform_health_check rejected every optional argument") +async def _window_gated_health_check_db_save( + save: Callable[[], Awaitable[None]], + pod_lock_manager: PodLockManager | None, + lock_ttl: int | None, +) -> None: + """ + Persist at most once per window fleet-wide: the lock is the "this window's save is + done" marker, so it is deliberately never released and expires with the interval. + """ + if pod_lock_manager is not None and pod_lock_manager.redis_cache is not None: + acquired: Final = await pod_lock_manager.acquire_lock( + cronjob_id=BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME, + ttl=lock_ttl, + allow_reentrant=False, + ) + if not acquired: + verbose_proxy_logger.debug("background_health_check_db_save_skipped another pod persisted this window") + return + await save() + + def _schedule_background_health_check_db_save( prisma_client, shared_health_manager, model_list: list, healthy_endpoints: list, unhealthy_endpoints: list, + pod_lock_manager: PodLockManager | None = None, + lock_ttl: int | None = None, ): """Fire-and-forget: persist health check results to DB if prisma is available.""" if prisma_client is None: @@ -3598,16 +3632,16 @@ def _schedule_background_health_check_db_save( checked_by: Final = shared_health_manager.pod_id if shared_health_manager is not None else "background_health_check" start_time: Final = time_module.time() - asyncio.create_task( - _save_background_health_checks_to_db( - prisma_client, - model_list, - healthy_endpoints, - unhealthy_endpoints, - start_time, - checked_by=checked_by, - ) + save: Final = partial( + _save_background_health_checks_to_db, + prisma_client, + model_list, + healthy_endpoints, + unhealthy_endpoints, + start_time, + checked_by=checked_by, ) + asyncio.create_task(_window_gated_health_check_db_save(save, pod_lock_manager, lock_ttl)) def _get_endpoint_exception_status(endpoint: dict, exceptions: dict) -> int: @@ -3914,6 +3948,8 @@ async def _run_background_health_check(): _llm_model_list, healthy_endpoints, unhealthy_endpoints, + pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager, + lock_ttl=health_check_interval, ) # Write health state to router cache for health-check-driven routing diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index cab2bd6d9db..01530762800 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -121,6 +121,11 @@ from litellm.proxy.db.exception_handler import ( PrismaDBExceptionHandler, call_with_db_reconnect_retry, ) +from litellm.proxy.db.health_check_latest import ( + LatestHealthCheckRow, + fetch_latest_health_checks, + fetch_latest_health_checks_for_models, +) from litellm.proxy.db.log_db_metrics import log_db_metrics from litellm.proxy.db.prisma_client import ( PrismaWrapper, @@ -5962,48 +5967,13 @@ class PrismaClient: verbose_proxy_logger.error("Error getting health check history: %s", e) return [] - async def get_all_latest_health_checks(self) -> "Sequence[prisma_models.LiteLLM_HealthCheckTable]": - """ - Get the latest health check for each model. + async def get_all_latest_health_checks(self) -> tuple[LatestHealthCheckRow, ...]: + """Latest health check per (model_id, model_name), deduplicated in Postgres.""" + return await fetch_latest_health_checks(self) - Uses DB-level DISTINCT ON (model_id, model_name) with ORDER BY checked_at DESC - (via Prisma ``distinct`` + ``order``) so we never load the full history into memory. - """ - try: - return await HealthCheckRepository(self).table.find_many( - distinct=["model_id", "model_name"], - order=[ - {"model_id": "asc"}, - {"model_name": "asc"}, - {"checked_at": "desc"}, - ], - ) - except Exception as e: - verbose_proxy_logger.error("Error getting all latest health checks: %s", e) - return [] - - async def get_latest_health_checks_for_models( - self, model_names: "Sequence[str]" - ) -> "Sequence[prisma_models.LiteLLM_HealthCheckTable]": - """ - Get the latest health check for each of the named models. - - Same DISTINCT ON as ``get_all_latest_health_checks``, bounded to the models asked - about, so a paged caller reads health for its page instead of for the whole table. - """ - if not model_names: - return () - latest_first: Final = (("model_id", "asc"), ("model_name", "asc"), ("checked_at", "desc")) - order: Final = [{field: direction} for field, direction in latest_first] # mutable-ok: prisma order is a list - try: - return await HealthCheckRepository(self).table.find_many( - where={"model_name": {"in": list(model_names)}}, # mutable-ok: prisma filters are dicts and lists - distinct=["model_id", "model_name"], # mutable-ok: prisma distinct takes a list - order=order, - ) - except Exception as e: # noqa: BLE001 # health decorates a list; a driver error must not fail the page - verbose_proxy_logger.error("Error getting latest health checks for models: %s", e) - return () + async def get_latest_health_checks_for_models(self, model_names: Sequence[str]) -> tuple[LatestHealthCheckRow, ...]: + """Same as ``get_all_latest_health_checks``, bounded to the named models.""" + return await fetch_latest_health_checks_for_models(self, model_names) ### HELPER FUNCTIONS ### diff --git a/tests/test_litellm/proxy/db/test_health_check_latest.py b/tests/test_litellm/proxy/db/test_health_check_latest.py new file mode 100644 index 00000000000..529e4d8f2e9 --- /dev/null +++ b/tests/test_litellm/proxy/db/test_health_check_latest.py @@ -0,0 +1,111 @@ +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.db.health_check_latest import ( + LATEST_HEALTH_CHECKS_FOR_MODELS_SQL, + LATEST_HEALTH_CHECKS_SQL, + fetch_latest_health_checks, + fetch_latest_health_checks_for_models, +) + + +def _prisma(rows): + prisma = MagicMock() + prisma.db.query_raw = AsyncMock(return_value=rows) + return prisma + + +def _raw_row(**overrides): + row = { + "health_check_id": "hc-1", + "model_name": "gpt-4", + "model_id": "deployment-abc", + "status": "healthy", + "healthy_count": 1, + "unhealthy_count": 0, + "error_message": None, + "response_time_ms": 12.5, + "details": None, + "checked_by": "pod-1", + "checked_at": "2026-08-25T00:00:00+00:00", + "created_at": "2026-08-25T00:00:00+00:00", + "updated_at": "2026-08-25T00:00:00+00:00", + } + return {**row, **overrides} + + +@pytest.mark.asyncio +async def test_fetch_all_runs_one_distinct_on_query_with_no_parameters(): + """The dedup must be in the SQL: prisma find_many(distinct=...) streams the whole history table.""" + prisma = _prisma([]) + assert await fetch_latest_health_checks(prisma) == () + assert prisma.db.query_raw.await_args.args == (LATEST_HEALTH_CHECKS_SQL,) + assert 'DISTINCT ON ("model_id", "model_name")' in LATEST_HEALTH_CHECKS_SQL + assert '"checked_at" DESC' in LATEST_HEALTH_CHECKS_SQL + + +@pytest.mark.asyncio +async def test_raw_datetimes_come_back_tz_aware_with_or_without_an_offset(): + """The save path subtracts checked_at from datetime.now(timezone.utc); a naive value would TypeError.""" + naive = _raw_row(health_check_id="hc-naive", model_id=None, checked_at="2026-08-25T00:00:00") + aware = _raw_row(health_check_id="hc-aware", checked_at="2026-08-25T01:00:00+02:00") + rows = await fetch_latest_health_checks(_prisma([naive, aware])) + assert {row.health_check_id: (row.model_id, row.checked_at) for row in rows} == { + "hc-naive": (None, datetime(2026, 8, 25, 0, 0, tzinfo=timezone.utc)), + "hc-aware": ("deployment-abc", datetime(2026, 8, 24, 23, 0, tzinfo=timezone.utc)), + } + + +@pytest.mark.asyncio +async def test_json_details_decode_from_text_and_pass_through_as_dict(): + rows = await fetch_latest_health_checks( + _prisma( + [ + _raw_row(health_check_id="text", details='{"region": "eu"}'), + _raw_row(health_check_id="dict", details={"region": "us"}), + _raw_row(health_check_id="none", details=None), + ] + ) + ) + assert {row.health_check_id: row.details for row in rows} == { + "text": {"region": "eu"}, + "dict": {"region": "us"}, + "none": None, + } + + +@pytest.mark.asyncio +async def test_fetch_all_degrades_to_no_rows_when_the_query_fails(): + prisma = _prisma([]) + prisma.db.query_raw.side_effect = RuntimeError("db down") + assert await fetch_latest_health_checks(prisma) == () + + +@pytest.mark.asyncio +async def test_fetch_all_degrades_to_no_rows_for_a_malformed_row(): + assert await fetch_latest_health_checks(_prisma([{"unexpected": "shape"}])) == () + + +@pytest.mark.asyncio +async def test_fetch_for_models_binds_the_page_as_the_only_parameter(): + prisma = _prisma([_raw_row()]) + rows = await fetch_latest_health_checks_for_models(prisma, ("gpt-4", "claude-opus")) + assert prisma.db.query_raw.await_args.args == (LATEST_HEALTH_CHECKS_FOR_MODELS_SQL, ["gpt-4", "claude-opus"]) + assert [row.model_name for row in rows] == ["gpt-4"] + assert 'WHERE "model_name" = ANY($1)' in LATEST_HEALTH_CHECKS_FOR_MODELS_SQL + + +@pytest.mark.asyncio +async def test_fetch_for_models_skips_the_database_for_an_empty_page(): + prisma = _prisma([]) + assert await fetch_latest_health_checks_for_models(prisma, ()) == () + prisma.db.query_raw.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_fetch_for_models_degrades_to_no_rows_when_the_query_fails(): + prisma = _prisma([]) + prisma.db.query_raw.side_effect = RuntimeError("db down") + assert await fetch_latest_health_checks_for_models(prisma, ("gpt-4",)) == () diff --git a/tests/test_litellm/proxy/proxy_server/test_background_health.py b/tests/test_litellm/proxy/proxy_server/test_background_health.py index 990844369f7..4874f75752c 100644 --- a/tests/test_litellm/proxy/proxy_server/test_background_health.py +++ b/tests/test_litellm/proxy/proxy_server/test_background_health.py @@ -20,6 +20,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest import litellm.proxy.proxy_server as proxy_server +from litellm.constants import BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME from litellm.proxy.proxy_server import ( _adaptive_router_flusher_loop, _get_endpoint_exception_status, @@ -245,6 +246,86 @@ async def test_schedule_background_health_check_db_save_invalid_no_event_loop_ra ) +def _lock_manager(redis_cache, acquired): + manager = MagicMock() + manager.redis_cache = redis_cache + manager.acquire_lock = AsyncMock(return_value=acquired) + manager.release_lock = AsyncMock() + return manager + + +def _capture_saves(monkeypatch): + saves = [] + + async def _fake_save(*_args, **kwargs): + saves.append(kwargs) + + import litellm.proxy.health_endpoints._health_endpoints as he + + monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + return saves + + +def _schedule_with(lock_manager): + _schedule_background_health_check_db_save( + prisma_client=MagicMock(), + shared_health_manager=None, + model_list=[], + healthy_endpoints=[], + unhealthy_endpoints=[], + pod_lock_manager=lock_manager, + lock_ttl=300, + ) + + +@pytest.mark.asyncio +async def test_schedule_background_health_check_db_save_skips_a_window_another_pod_persisted(monkeypatch): + saves = _capture_saves(monkeypatch) + lock_manager = _lock_manager(redis_cache=MagicMock(), acquired=False) + + _schedule_with(lock_manager) + await asyncio.sleep(0) + + assert saves == [] + + +@pytest.mark.asyncio +async def test_schedule_background_health_check_db_save_holds_the_window_lock_for_the_whole_interval(monkeypatch): + """The lock is the "saved this window" marker: never reentrant, TTL = interval, and never released.""" + saves = _capture_saves(monkeypatch) + lock_manager = _lock_manager(redis_cache=MagicMock(), acquired=True) + + _schedule_with(lock_manager) + await asyncio.sleep(0) + + assert normalize( + { + "saves": len(saves), + "lock_request": lock_manager.acquire_lock.await_args.kwargs, + "released": lock_manager.release_lock.await_count, + } + ) == { + "saves": 1, + "lock_request": { + "cronjob_id": BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME, + "ttl": 300, + "allow_reentrant": False, + }, + "released": 0, + } + + +@pytest.mark.asyncio +async def test_schedule_background_health_check_db_save_runs_ungated_without_redis(monkeypatch): + saves = _capture_saves(monkeypatch) + lock_manager = _lock_manager(redis_cache=None, acquired=True) + + _schedule_with(lock_manager) + await asyncio.sleep(0) + + assert (len(saves), lock_manager.acquire_lock.await_count) == (1, 0) + + # --------------------------------------------------------------------------- # _get_endpoint_exception_status # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/test_health_check_functions.py b/tests/test_litellm/proxy/test_health_check_functions.py index fdae11d517a..9adaddb032d 100644 --- a/tests/test_litellm/proxy/test_health_check_functions.py +++ b/tests/test_litellm/proxy/test_health_check_functions.py @@ -6,6 +6,8 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.db.health_check_latest import LatestHealthCheckRow from litellm.proxy.health_endpoints._health_endpoints import ( _aggregate_health_check_results, _build_model_param_to_info_mapping, @@ -13,6 +15,7 @@ from litellm.proxy.health_endpoints._health_endpoints import ( _save_background_health_checks_to_db, _save_health_check_results_if_changed, _save_health_check_to_db, + latest_health_checks_endpoint, ) from litellm.proxy.utils import PrismaClient @@ -451,96 +454,125 @@ async def test_save_background_health_checks_to_db_exception_handling(): # Function should complete without raising -@pytest.mark.asyncio -async def test_get_all_latest_health_checks_with_model_id(mock_prisma): - """Test get_all_latest_health_checks properly groups by model_id""" - mock_check2 = MagicMock() - mock_check2.model_id = "model-456" - mock_check2.model_name = "gpt-3.5-turbo" - mock_check2.checked_at = datetime.now(timezone.utc) - timedelta(minutes=5) - - mock_check3 = MagicMock() - mock_check3.model_id = "model-123" - mock_check3.model_name = "gpt-3.5-turbo" - mock_check3.checked_at = datetime.now(timezone.utc) - timedelta( - minutes=1 - ) # Latest for model-123 - - # Order by checked_at desc - mock_prisma.db.litellm_healthchecktable.find_many = AsyncMock( - return_value=[mock_check3, mock_check2] - ) - - result = await mock_prisma.get_all_latest_health_checks() - - # Should return 2 unique models (by model_id) - assert len(result) == 2 - - # Should have latest check for each model_id - model_ids = {check.model_id for check in result} - assert "model-123" in model_ids - assert "model-456" in model_ids - - # model-123 should have the latest check (1 minute ago) - model123_check = next(c for c in result if c.model_id == "model-123") - assert model123_check.checked_at == mock_check3.checked_at +def _raw_latest_row(model_name: str, model_id, checked_at: datetime) -> dict: + return { + "health_check_id": f"hc-{model_id or 'no-id'}-{model_name}", + "model_name": model_name, + "model_id": model_id, + "status": "healthy", + "healthy_count": 1, + "unhealthy_count": 0, + "error_message": None, + "response_time_ms": 10.0, + "details": None, + "checked_by": "pod-1", + "checked_at": checked_at.isoformat(), + "created_at": checked_at.isoformat(), + "updated_at": checked_at.isoformat(), + } @pytest.mark.asyncio -async def test_get_all_latest_health_checks_without_model_id(mock_prisma): - """Test get_all_latest_health_checks groups by model_name when model_id is None""" - mock_check2 = MagicMock() - mock_check2.model_id = None - mock_check2.model_name = "gpt-3.5-turbo" - mock_check2.checked_at = datetime.now(timezone.utc) - timedelta(minutes=1) # Latest - - mock_prisma.db.litellm_healthchecktable.find_many = AsyncMock( - return_value=[mock_check2] - ) - - result = await mock_prisma.get_all_latest_health_checks() - - # Should return 1 unique model (by model_name) - assert len(result) == 1 - assert result[0].model_name == "gpt-3.5-turbo" - assert result[0].checked_at == mock_check2.checked_at # Latest - - -@pytest.mark.asyncio -async def test_get_all_latest_health_checks_same_name_with_and_without_model_id( - mock_prisma, -): +async def test_get_all_latest_health_checks_keeps_every_distinct_group_with_its_own_checked_at(mock_prisma): """ - Same model_name can appear twice after DISTINCT ON: once keyed by (model_id, name) - and once by (NULL, name) — different Postgres groups than a single row with id. + Postgres owns the dedup. (id, name), (other id, name) and (NULL, name) are distinct groups and each row + must arrive typed, with its own checked_at, for the 1h re-save compare and the id-or-name lookup key. """ now = datetime.now(timezone.utc) - with_id = MagicMock() - with_id.model_id = "deployment-abc" - with_id.model_name = "gpt-4" - with_id.checked_at = now - timedelta(minutes=2) - - without_id = MagicMock() - without_id.model_id = None - without_id.model_name = "gpt-4" - without_id.checked_at = now - timedelta(minutes=1) - - mock_prisma.db.litellm_healthchecktable.find_many = AsyncMock( - return_value=[without_id, with_id] + mock_prisma.db.query_raw = AsyncMock( + return_value=[ + _raw_latest_row("gpt-3.5-turbo", "model-123", now - timedelta(minutes=1)), + _raw_latest_row("gpt-3.5-turbo", "model-456", now - timedelta(minutes=5)), + _raw_latest_row("gpt-4", "deployment-abc", now - timedelta(minutes=2)), + _raw_latest_row("gpt-4", None, now - timedelta(minutes=3)), + ] ) result = await mock_prisma.get_all_latest_health_checks() - assert len(result) == 2 - names = {r.model_name for r in result} - assert names == {"gpt-4"} - ids = {r.model_id for r in result} - assert "deployment-abc" in ids - assert None in ids + assert {(check.model_id, check.model_name): check.checked_at for check in result} == { + ("model-123", "gpt-3.5-turbo"): now - timedelta(minutes=1), + ("model-456", "gpt-3.5-turbo"): now - timedelta(minutes=5), + ("deployment-abc", "gpt-4"): now - timedelta(minutes=2), + (None, "gpt-4"): now - timedelta(minutes=3), + } - by_key = {(r.model_id, r.model_name): r for r in result} - assert by_key[("deployment-abc", "gpt-4")].checked_at == with_id.checked_at - assert by_key[(None, "gpt-4")].checked_at == without_id.checked_at + +@pytest.mark.asyncio +async def test_save_background_health_checks_compares_raw_checked_at_against_utc_now(mock_prisma): + """ + Raw rows carry ISO strings and the engine may omit the offset. A naive checked_at would TypeError + inside the 1h compare, be swallowed, and silently stop every save; a stale row must still re-save. + """ + stale = (datetime.now(timezone.utc) - timedelta(hours=2)).replace(tzinfo=None) + fresh = datetime.now(timezone.utc) - timedelta(minutes=5) + mock_prisma.db.query_raw = AsyncMock( + return_value=[ + _raw_latest_row("stale-model", "stale-id", stale), + _raw_latest_row("fresh-model", "fresh-id", fresh), + ] + ) + mock_prisma.save_health_check_result = AsyncMock() + model_list = [ + {"model_name": "stale-model", "model_info": {"id": "stale-id"}, "litellm_params": {"model": "openai/stale"}}, + {"model_name": "fresh-model", "model_info": {"id": "fresh-id"}, "litellm_params": {"model": "openai/fresh"}}, + ] + + await _save_background_health_checks_to_db( + mock_prisma, + model_list, + [{"model": "openai/stale"}, {"model": "openai/fresh"}], + [], + time.time(), + "pod-1", + ) + await asyncio.sleep(0) + + assert [call.kwargs["model_id"] for call in mock_prisma.save_health_check_result.await_args_list] == ["stale-id"] + + +@pytest.mark.asyncio +async def test_latest_health_checks_endpoint_serialises_raw_rows(monkeypatch): + row = LatestHealthCheckRow( + health_check_id="hc-1", + model_name="gpt-4", + model_id="deployment-abc", + status="healthy", + healthy_count=1, + unhealthy_count=0, + error_message=None, + response_time_ms=12.5, + details='{"region": "eu"}', + checked_by="pod-1", + checked_at=datetime(2026, 8, 25), + created_at=datetime(2026, 8, 25, tzinfo=timezone.utc), + updated_at=datetime(2026, 8, 25, tzinfo=timezone.utc), + ) + prisma = MagicMock() + prisma.get_all_latest_health_checks = AsyncMock(return_value=(row,)) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) + + response = await latest_health_checks_endpoint(user_api_key_dict=UserAPIKeyAuth()) + + assert response == { + "latest_health_checks": { + "deployment-abc": { + "health_check_id": "hc-1", + "model_name": "gpt-4", + "model_id": "deployment-abc", + "status": "healthy", + "healthy_count": 1, + "unhealthy_count": 0, + "error_message": None, + "response_time_ms": 12.5, + "details": {"region": "eu"}, + "checked_by": "pod-1", + "checked_at": "2026-08-25T00:00:00+00:00", + "created_at": "2026-08-25T00:00:00+00:00", + } + }, + "total_models": 1, + } @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py index 9f48ba68b4f..fbe9934f06c 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py @@ -22,6 +22,10 @@ from unittest.mock import AsyncMock, MagicMock import pytest +from litellm.proxy.db.health_check_latest import ( + LATEST_HEALTH_CHECKS_FOR_MODELS_SQL, + LATEST_HEALTH_CHECKS_SQL, +) from litellm.proxy.utils import PrismaClient @@ -261,36 +265,49 @@ async def test_get_health_check_history_db_error_returns_empty_list( assert await prisma_client.get_health_check_history() == [] +def _raw_health_check_row(model_name: str = "gpt-4", model_id: str | None = "deployment-abc") -> dict[str, Any]: + return { + "health_check_id": f"hc-{model_name}", + "model_name": model_name, + "model_id": model_id, + "status": "healthy", + "healthy_count": 1, + "unhealthy_count": 0, + "error_message": None, + "response_time_ms": 12.5, + "details": None, + "checked_by": "pod-1", + "checked_at": "2026-08-25T00:00:00+00:00", + "created_at": "2026-08-25T00:00:00+00:00", + "updated_at": "2026-08-25T00:00:00+00:00", + } + + @pytest.mark.asyncio -async def test_get_all_latest_health_checks_uses_distinct( +async def test_get_all_latest_health_checks_dedups_in_postgres( prisma_client: PrismaClient, ) -> None: - rows = [MagicMock(name=f"row-{i}") for i in range(3)] - prisma_client.db.litellm_healthchecktable.find_many = AsyncMock(return_value=rows) + """A revert to prisma find_many(distinct=...) streams the whole history table; the SQL must own the DISTINCT.""" + prisma_client.db.query_raw = AsyncMock(return_value=[_raw_health_check_row()]) result = await prisma_client.get_all_latest_health_checks() - kwargs = prisma_client.db.litellm_healthchecktable.find_many.await_args.kwargs actual = { - "len": len(result), - "distinct": kwargs["distinct"], - "order_len": len(kwargs["order"]), - "first_order": kwargs["order"][0], + "query": prisma_client.db.query_raw.await_args.args, + "rows": [(row.model_id, row.model_name, row.status) for row in result], + "row_type": type(result[0]).__name__, } assert actual == { - "len": 3, - "distinct": ["model_id", "model_name"], - "order_len": 3, - "first_order": {"model_id": "asc"}, + "query": (LATEST_HEALTH_CHECKS_SQL,), + "rows": [("deployment-abc", "gpt-4", "healthy")], + "row_type": "LatestHealthCheckRow", } @pytest.mark.asyncio -async def test_get_all_latest_health_checks_db_error_returns_empty_list( +async def test_get_all_latest_health_checks_db_error_returns_no_rows( prisma_client: PrismaClient, ) -> None: - prisma_client.db.litellm_healthchecktable.find_many = AsyncMock( - side_effect=RuntimeError("oops") - ) - assert await prisma_client.get_all_latest_health_checks() == [] + prisma_client.db.query_raw = AsyncMock(side_effect=RuntimeError("oops")) + assert await prisma_client.get_all_latest_health_checks() == () @pytest.mark.asyncio @@ -298,18 +315,15 @@ async def test_get_latest_health_checks_for_models_bounds_the_query_to_those_mod prisma_client: PrismaClient, ) -> None: """A paged caller reads health for its page; an unbounded read is the bug this exists to avoid.""" - prisma_client.db.litellm_healthchecktable.find_many = AsyncMock(return_value=[]) - await prisma_client.get_latest_health_checks_for_models(["gpt-5", "claude-opus"]) - kwargs = prisma_client.db.litellm_healthchecktable.find_many.await_args.kwargs + prisma_client.db.query_raw = AsyncMock(return_value=[_raw_health_check_row(model_name="gpt-5")]) + result = await prisma_client.get_latest_health_checks_for_models(["gpt-5", "claude-opus"]) actual = { - "where": kwargs["where"], - "distinct": kwargs["distinct"], - "order": kwargs["order"], + "query": prisma_client.db.query_raw.await_args.args, + "rows": [row.model_name for row in result], } assert actual == { - "where": {"model_name": {"in": ["gpt-5", "claude-opus"]}}, - "distinct": ["model_id", "model_name"], - "order": [{"model_id": "asc"}, {"model_name": "asc"}, {"checked_at": "desc"}], + "query": (LATEST_HEALTH_CHECKS_FOR_MODELS_SQL, ["gpt-5", "claude-opus"]), + "rows": ["gpt-5"], } @@ -317,14 +331,14 @@ async def test_get_latest_health_checks_for_models_bounds_the_query_to_those_mod async def test_get_latest_health_checks_for_models_does_not_query_for_an_empty_page( prisma_client: PrismaClient, ) -> None: - prisma_client.db.litellm_healthchecktable.find_many = AsyncMock(return_value=[]) + prisma_client.db.query_raw = AsyncMock(return_value=[]) assert await prisma_client.get_latest_health_checks_for_models([]) == () - assert prisma_client.db.litellm_healthchecktable.find_many.await_count == 0 + assert prisma_client.db.query_raw.await_count == 0 @pytest.mark.asyncio -async def test_get_latest_health_checks_for_models_db_error_returns_empty_list( +async def test_get_latest_health_checks_for_models_db_error_returns_no_rows( prisma_client: PrismaClient, ) -> None: - prisma_client.db.litellm_healthchecktable.find_many = AsyncMock(side_effect=RuntimeError("oops")) + prisma_client.db.query_raw = AsyncMock(side_effect=RuntimeError("oops")) assert await prisma_client.get_latest_health_checks_for_models(["gpt-5"]) == () From f980b94923e25e5c01607fcdf2088bd0c232fc9e Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 3 Sep 2026 12:36:04 +0200 Subject: [PATCH 27/60] fix(proxy): release the health check save window lock on failure or cancel --- .../health_endpoints/_health_endpoints.py | 105 +++++++----- litellm/proxy/proxy_server.py | 45 +++-- .../proxy_server/test_background_health.py | 109 +++++++----- .../proxy/test_health_check_functions.py | 155 +++++++++++++----- 4 files changed, 278 insertions(+), 136 deletions(-) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 8b57bdca2fe..91825a5606c 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -7,7 +7,7 @@ import secrets import time import traceback from collections.abc import Iterable, Mapping -from datetime import datetime, timedelta +from datetime import datetime, timedelta, timezone from typing import Any, Final, Literal, TypedDict, cast import fastapi @@ -41,6 +41,7 @@ from litellm.proxy.auth.auth_utils import ( ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler +from litellm.proxy.db.health_check_latest import LatestHealthCheckRow from litellm.proxy.db.proxy_worker_heartbeat import count_live_proxy_workers from litellm.proxy.health_check import ( ADMIN_ONLY_HEALTH_DISPLAY_PARAMS, @@ -697,13 +698,42 @@ def _aggregate_health_check_results( return model_results +class _AggregatedHealthResult(TypedDict): + """One entry of ``_aggregate_health_check_results``: a model's counts for this cycle.""" + + model_name: ReadOnly[str] + model_id: ReadOnly[str | None] + healthy_count: ReadOnly[int] + unhealthy_count: ReadOnly[int] + error_message: ReadOnly[str | None] + + +def _new_health_status(result: _AggregatedHealthResult) -> str: + return "healthy" if result["healthy_count"] > 0 else "unhealthy" + + +def _should_persist_health_check_result( + result: _AggregatedHealthResult, latest_checks_map: Mapping[str, LatestHealthCheckRow] +) -> bool: + """ + True when this result has to be written: no previous row, the status changed, or the + previous row is older than one hour (periodic refresh while the status is stable). + """ + lookup_key: Final = result["model_id"] if result["model_id"] else result["model_name"] + last_check: Final = latest_checks_map.get(lookup_key) + if last_check is None or last_check.status != _new_health_status(result): + return True + time_since_last_check: Final = (datetime.now(timezone.utc) - last_check.checked_at).total_seconds() + return time_since_last_check >= 3600 # 1 hour threshold + + async def _save_health_check_results_if_changed( prisma_client, model_results: dict, latest_checks_map: dict, start_time: float, checked_by: str | None = None, -): +) -> bool: """ Save health check results to database, but only if status changed or >1 hour since last save. @@ -714,47 +744,39 @@ async def _save_health_check_results_if_changed( - Status changes: Immediate write (no delay) - Result: ~92% reduction in DB writes for stable systems, while maintaining real-time updates on changes + The writes are awaited rather than detached so the caller learns whether this cycle's + persistence completed. + Args: prisma_client: Database client model_results: Dictionary of aggregated health check results per model latest_checks_map: Dictionary mapping model_id/model_name to latest health check start_time: Start time of health check for calculating response time checked_by: Identifier for who/what performed the check + + Returns: + True when every row that needed writing was written (including when nothing needed + writing); False when any write failed. """ - for result in model_results.values(): - new_status = "healthy" if result["healthy_count"] > 0 else "unhealthy" - - # Check if we should save this result - should_save = True - lookup_key = result["model_id"] if result["model_id"] else result["model_name"] - if lookup_key in latest_checks_map: - last_check = latest_checks_map[lookup_key] - # Only save if status changed or if it's been a while since last check - if last_check.status == new_status: - # Check if last check was recent (within 1 hour) - if last_check.checked_at: - from datetime import datetime, timezone - - time_since_last_check = (datetime.now(timezone.utc) - last_check.checked_at).total_seconds() - # Only skip if status unchanged AND checked recently (within 1 hour) - # This ensures we still get periodic updates even if status is stable - if time_since_last_check < 3600: # 1 hour threshold - should_save = False - - if should_save: - asyncio.create_task( - prisma_client.save_health_check_result( - model_name=result["model_name"], - model_id=result["model_id"], - status=new_status, - healthy_count=result["healthy_count"], - unhealthy_count=result["unhealthy_count"], - error_message=result["error_message"], - response_time_ms=(time.time() - start_time) * 1000, - details=None, - checked_by=checked_by, - ) - ) + to_write: Final = tuple( + result for result in model_results.values() if _should_persist_health_check_result(result, latest_checks_map) + ) + writes: Final = tuple( + prisma_client.save_health_check_result( + model_name=result["model_name"], + model_id=result["model_id"], + status=_new_health_status(result), + healthy_count=result["healthy_count"], + unhealthy_count=result["unhealthy_count"], + error_message=result["error_message"], + response_time_ms=(time.time() - start_time) * 1000, + details=None, + checked_by=checked_by, + ) + for result in to_write + ) + rows: Final = await asyncio.gather(*writes) + return all(row is not None for row in rows) async def _save_background_health_checks_to_db( @@ -764,7 +786,7 @@ async def _save_background_health_checks_to_db( unhealthy_endpoints: list, start_time: float, checked_by: str | None = None, -): +) -> bool: """ Save background health check results to database for each model. @@ -773,9 +795,13 @@ async def _save_background_health_checks_to_db( OPTIMIZATION: Only saves to database if the status has changed from the last saved check. This dramatically reduces database writes when health status remains stable. + + Returns: + True when this cycle's persistence completed; False when it was skipped or any step + failed. Never raises: a database failure must not break the health check loop. """ if prisma_client is None: - return + return False try: # Step 1: Build mapping from model parameter to model info @@ -798,7 +824,7 @@ async def _save_background_health_checks_to_db( latest_checks_map[key] = check # Step 4: Save aggregated results, but only if status changed - await _save_health_check_results_if_changed( + return await _save_health_check_results_if_changed( prisma_client, model_results, latest_checks_map, @@ -808,6 +834,7 @@ async def _save_background_health_checks_to_db( except Exception as db_error: verbose_proxy_logger.warning("Failed to save background health checks to database: %s", db_error) # Continue execution - don't let database save failure break health checks + return False _PROXY_ADMIN_ROLES: Final = frozenset( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0e330272bdf..88ab36b3d1e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -152,6 +152,7 @@ if TYPE_CHECKING: from prisma import models as prisma_models from litellm.integrations.opentelemetry import OpenTelemetry + from litellm.proxy.health_check_utils.shared_health_check_manager import SharedHealthCheckManager Span = _Span | Any else: @@ -3592,35 +3593,49 @@ async def _run_direct_health_check_with_instrumentation( async def _window_gated_health_check_db_save( - save: Callable[[], Awaitable[None]], + save: Callable[[], Awaitable[bool]], pod_lock_manager: PodLockManager | None, lock_ttl: int | None, ) -> None: """ - Persist at most once per window fleet-wide: the lock is the "this window's save is - done" marker, so it is deliberately never released and expires with the interval. + Persist at most once per window fleet-wide. A completed save keeps the lock as the + "this window's save is done" marker, so it is deliberately never released and expires + with the interval. A save that reports failure or is cancelled releases the lock so + another pod's cycle in the same window can retry, instead of the fleet going a whole + window without a write. """ - if pod_lock_manager is not None and pod_lock_manager.redis_cache is not None: - acquired: Final = await pod_lock_manager.acquire_lock( - cronjob_id=BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME, - ttl=lock_ttl, - allow_reentrant=False, + if pod_lock_manager is None or pod_lock_manager.redis_cache is None: + await save() + return + acquired: Final = await pod_lock_manager.acquire_lock( + cronjob_id=BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME, + ttl=lock_ttl, + allow_reentrant=False, + ) + if not acquired: + verbose_proxy_logger.debug("background_health_check_db_save_skipped another pod persisted this window") + return + try: + persisted: Final = await save() + except BaseException: + await pod_lock_manager.release_lock(cronjob_id=BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME) + raise + if not persisted: + verbose_proxy_logger.warning( + "background_health_check_db_save_incomplete released the window lock so another pod can retry" ) - if not acquired: - verbose_proxy_logger.debug("background_health_check_db_save_skipped another pod persisted this window") - return - await save() + await pod_lock_manager.release_lock(cronjob_id=BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME) def _schedule_background_health_check_db_save( - prisma_client, - shared_health_manager, + prisma_client: PrismaClient | None, + shared_health_manager: "SharedHealthCheckManager | None", model_list: list, healthy_endpoints: list, unhealthy_endpoints: list, pod_lock_manager: PodLockManager | None = None, lock_ttl: int | None = None, -): +) -> None: """Fire-and-forget: persist health check results to DB if prisma is available.""" if prisma_client is None: return diff --git a/tests/test_litellm/proxy/proxy_server/test_background_health.py b/tests/test_litellm/proxy/proxy_server/test_background_health.py index 4874f75752c..d99aa637955 100644 --- a/tests/test_litellm/proxy/proxy_server/test_background_health.py +++ b/tests/test_litellm/proxy/proxy_server/test_background_health.py @@ -112,13 +112,11 @@ async def test_run_direct_health_check_with_instrumentation_returns_results( lambda _gs: {}, ) - healthy, unhealthy, exceptions = ( - await _run_direct_health_check_with_instrumentation( - model_list=[{"model_name": "gpt-4"}], - details=False, - max_concurrency=1, - instrumentation_context={"source": "test"}, - ) + healthy, unhealthy, exceptions = await _run_direct_health_check_with_instrumentation( + model_list=[{"model_name": "gpt-4"}], + details=False, + max_concurrency=1, + instrumentation_context={"source": "test"}, ) assert normalize( @@ -254,11 +252,12 @@ def _lock_manager(redis_cache, acquired): return manager -def _capture_saves(monkeypatch): +def _capture_saves(monkeypatch, persisted=True): saves = [] async def _fake_save(*_args, **kwargs): saves.append(kwargs) + return persisted import litellm.proxy.health_endpoints._health_endpoints as he @@ -266,6 +265,15 @@ def _capture_saves(monkeypatch): return saves +def _cancel_during_save(monkeypatch): + async def _fake_save(*_args, **_kwargs): + raise asyncio.CancelledError() + + import litellm.proxy.health_endpoints._health_endpoints as he + + monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + + def _schedule_with(lock_manager): _schedule_background_health_check_db_save( prisma_client=MagicMock(), @@ -315,15 +323,56 @@ async def test_schedule_background_health_check_db_save_holds_the_window_lock_fo } +@pytest.mark.asyncio +async def test_schedule_background_health_check_db_save_releases_the_window_lock_when_the_save_reports_failure( + monkeypatch, +): + """A failed save must not burn the window: release the lock so another pod's cycle can retry.""" + saves = _capture_saves(monkeypatch, persisted=False) + lock_manager = _lock_manager(redis_cache=MagicMock(), acquired=True) + + _schedule_with(lock_manager) + await asyncio.sleep(0) + + assert normalize( + { + "saves": len(saves), + "release_request": lock_manager.release_lock.await_args.kwargs, + "release_count": lock_manager.release_lock.await_count, + } + ) == { + "saves": 1, + "release_request": {"cronjob_id": BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME}, + "release_count": 1, + } + + +@pytest.mark.asyncio +async def test_schedule_background_health_check_db_save_releases_the_window_lock_when_the_save_is_cancelled( + monkeypatch, +): + """A pod shutting down mid-save releases the lock instead of holding it until the TTL.""" + _cancel_during_save(monkeypatch) + lock_manager = _lock_manager(redis_cache=MagicMock(), acquired=True) + + _schedule_with(lock_manager) + await asyncio.sleep(0) + + assert ( + lock_manager.release_lock.await_args.kwargs, + lock_manager.release_lock.await_count, + ) == ({"cronjob_id": BACKGROUND_HEALTH_CHECK_DB_SAVE_JOB_NAME}, 1) + + @pytest.mark.asyncio async def test_schedule_background_health_check_db_save_runs_ungated_without_redis(monkeypatch): - saves = _capture_saves(monkeypatch) + saves = _capture_saves(monkeypatch, persisted=False) lock_manager = _lock_manager(redis_cache=None, acquired=True) _schedule_with(lock_manager) await asyncio.sleep(0) - assert (len(saves), lock_manager.acquire_lock.await_count) == (1, 0) + assert (len(saves), lock_manager.acquire_lock.await_count, lock_manager.release_lock.await_count) == (1, 0, 0) # --------------------------------------------------------------------------- @@ -400,13 +449,9 @@ def test_write_health_state_to_router_cache_sets_states(monkeypatch): _write_health_state_to_router_cache(healthy, unhealthy, exceptions) - fake_router.health_state_cache.set_deployment_health_states.assert_called_once_with( - fake_states - ) + fake_router.health_state_cache.set_deployment_health_states.assert_called_once_with(fake_states) - call_args = fake_router.health_state_cache.set_deployment_health_states.call_args[ - 0 - ][0] + call_args = fake_router.health_state_cache.set_deployment_health_states.call_args[0][0] assert normalize( { "states_keys": sorted(call_args.keys()), @@ -448,9 +493,7 @@ def test_write_health_state_to_router_cache_populates_for_listing_filter(monkeyp fake_router.cooldown_time = 30 monkeypatch.setattr(proxy_server, "llm_router", fake_router) - monkeypatch.setattr( - proxy_server, "general_settings", {"model_list_healthy_only": True} - ) + monkeypatch.setattr(proxy_server, "general_settings", {"model_list_healthy_only": True}) fake_states = {"m1": {"is_healthy": True}, "m2": {"is_healthy": False}} @@ -484,9 +527,7 @@ def test_write_health_state_to_router_cache_populates_for_listing_filter(monkeyp {"m2": SimpleNamespace(status_code=500)}, ) - fake_router.health_state_cache.set_deployment_health_states.assert_called_once_with( - fake_states - ) + fake_router.health_state_cache.set_deployment_health_states.assert_called_once_with(fake_states) assert cooldowns == [] assert failures == [] @@ -496,9 +537,7 @@ def test_write_health_state_to_router_cache_swallows_internal_failures(monkeypat fake_router = MagicMock() fake_router.enable_health_check_routing = True fake_router.health_check_ignore_transient_errors = False - fake_router.health_state_cache.set_deployment_health_states.side_effect = ( - RuntimeError("cache exploded") - ) + fake_router.health_state_cache.set_deployment_health_states.side_effect = RuntimeError("cache exploded") monkeypatch.setattr(proxy_server, "llm_router", fake_router) @@ -528,9 +567,7 @@ async def test_adaptive_router_flusher_loop_flushes_each_router(monkeypatch): from litellm.types.router import TaggedPreRoutingStrategy fake_router = MagicMock() - fake_router.adaptive_routers = { - "alpha": [TaggedPreRoutingStrategy(tags=(), strategy=fake_ar)] - } + fake_router.adaptive_routers = {"alpha": [TaggedPreRoutingStrategy(tags=(), strategy=fake_ar)]} monkeypatch.setattr(proxy_server, "llm_router", fake_router) monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) @@ -628,12 +665,8 @@ async def test_run_background_health_check_runs_one_cycle_then_cancels(monkeypat "_run_direct_health_check_with_instrumentation", _fake_direct, ) - monkeypatch.setattr( - proxy_server, "_schedule_background_health_check_db_save", lambda *a, **kw: None - ) - monkeypatch.setattr( - proxy_server, "_write_health_state_to_router_cache", lambda *a, **kw: None - ) + monkeypatch.setattr(proxy_server, "_schedule_background_health_check_db_save", lambda *a, **kw: None) + monkeypatch.setattr(proxy_server, "_write_health_state_to_router_cache", lambda *a, **kw: None) monkeypatch.setattr( proxy_server, "health_check_filter_kwargs_from_general_settings", @@ -711,12 +744,8 @@ async def test_run_background_health_check_probes_only_listed_model_groups(monke "_run_direct_health_check_with_instrumentation", _fake_direct, ) - monkeypatch.setattr( - proxy_server, "_schedule_background_health_check_db_save", lambda *a, **kw: None - ) - monkeypatch.setattr( - proxy_server, "_write_health_state_to_router_cache", lambda *a, **kw: None - ) + monkeypatch.setattr(proxy_server, "_schedule_background_health_check_db_save", lambda *a, **kw: None) + monkeypatch.setattr(proxy_server, "_write_health_state_to_router_cache", lambda *a, **kw: None) monkeypatch.setattr( proxy_server, "health_check_filter_kwargs_from_general_settings", diff --git a/tests/test_litellm/proxy/test_health_check_functions.py b/tests/test_litellm/proxy/test_health_check_functions.py index 9adaddb032d..5aa80213134 100644 --- a/tests/test_litellm/proxy/test_health_check_functions.py +++ b/tests/test_litellm/proxy/test_health_check_functions.py @@ -24,12 +24,8 @@ from litellm.proxy.utils import PrismaClient def mock_prisma(): """Simplified mock PrismaClient with bound methods""" client = MagicMock() - client.db.litellm_healthchecktable.create = AsyncMock( - return_value={"id": "test-id"} - ) - client.db.litellm_healthchecktable.find_many = AsyncMock( - return_value=[{"id": "1", "model_name": "test"}] - ) + client.db.litellm_healthchecktable.create = AsyncMock(return_value={"id": "test-id"}) + client.db.litellm_healthchecktable.find_many = AsyncMock(return_value=[{"id": "1", "model_name": "test"}]) # Bind actual methods import types @@ -55,14 +51,10 @@ def mock_prisma(): ("healthy", 1, 0, False), # Database error case ], ) -async def test_save_health_check_result( - mock_prisma, status, healthy, unhealthy, should_succeed -): +async def test_save_health_check_result(mock_prisma, status, healthy, unhealthy, should_succeed): """Test health check result saving with various scenarios""" if not should_succeed: - mock_prisma.db.litellm_healthchecktable.create.side_effect = Exception( - "DB Error" - ) + mock_prisma.db.litellm_healthchecktable.create.side_effect = Exception("DB Error") result = await mock_prisma.save_health_check_result( model_name="test-model", @@ -190,9 +182,7 @@ def test_aggregate_health_check_results(): {"model": "gpt-4", "error": "Rate limit exceeded"}, ] - result = _aggregate_health_check_results( - model_param_to_info, healthy_endpoints, unhealthy_endpoints - ) + result = _aggregate_health_check_results(model_param_to_info, healthy_endpoints, unhealthy_endpoints) # Check gpt-3.5-turbo is healthy gpt35_key = ("model-123", "gpt-3.5-turbo") @@ -223,9 +213,7 @@ def test_aggregate_health_check_results_multiple_endpoints(): ] unhealthy_endpoints = [] - result = _aggregate_health_check_results( - model_param_to_info, healthy_endpoints, unhealthy_endpoints - ) + result = _aggregate_health_check_results(model_param_to_info, healthy_endpoints, unhealthy_endpoints) key = ("model-123", "gpt-3.5-turbo") assert result[key]["healthy_count"] == 2 @@ -401,7 +389,7 @@ async def test_save_background_health_checks_to_db(): start_time = 1234567890.0 - await _save_background_health_checks_to_db( + persisted = await _save_background_health_checks_to_db( mock_prisma, model_list, healthy_endpoints, @@ -410,7 +398,8 @@ async def test_save_background_health_checks_to_db(): "background_health_check", ) - # Should call get_all_latest_health_checks and save_health_check_result + # Should call get_all_latest_health_checks and save_health_check_result, and report completion + assert persisted is True mock_prisma.get_all_latest_health_checks.assert_called_once() mock_prisma.save_health_check_result.assert_called_once() @@ -421,22 +410,112 @@ async def test_save_background_health_checks_to_db(): assert call_kwargs["checked_by"] == "background_health_check" +def _two_model_results(): + return { + ("model-1", "gpt-4"): { + "model_name": "gpt-4", + "model_id": "model-1", + "healthy_count": 1, + "unhealthy_count": 0, + "error_message": None, + }, + ("model-2", "gpt-4o"): { + "model_name": "gpt-4o", + "model_id": "model-2", + "healthy_count": 0, + "unhealthy_count": 1, + "error_message": "boom", + }, + } + + +@pytest.mark.asyncio +async def test_save_health_check_results_if_changed_awaits_every_write_and_reports_success(): + """Writes are awaited, not detached, so the caller can tell the cycle's persistence completed.""" + mock_prisma = MagicMock() + mock_prisma.save_health_check_result = AsyncMock(return_value={"id": "row"}) + + persisted = await _save_health_check_results_if_changed( + mock_prisma, _two_model_results(), {}, 1234567890.0, "background_health_check" + ) + + assert (persisted, mock_prisma.save_health_check_result.await_count) == (True, 2) + + +@pytest.mark.asyncio +async def test_save_health_check_results_if_changed_reports_failure_when_a_write_returns_none(): + """save_health_check_result swallows DB errors and returns None; that must surface as False.""" + mock_prisma = MagicMock() + mock_prisma.save_health_check_result = AsyncMock(side_effect=[{"id": "row"}, None]) + + persisted = await _save_health_check_results_if_changed( + mock_prisma, _two_model_results(), {}, 1234567890.0, "background_health_check" + ) + + assert (persisted, mock_prisma.save_health_check_result.await_count) == (False, 2) + + +@pytest.mark.asyncio +async def test_save_health_check_results_if_changed_reports_success_when_nothing_needed_writing(): + mock_prisma = MagicMock() + mock_prisma.save_health_check_result = AsyncMock() + model_results = { + ("model-1", "gpt-4"): { + "model_name": "gpt-4", + "model_id": "model-1", + "healthy_count": 1, + "unhealthy_count": 0, + "error_message": None, + }, + } + latest_checks_map = { + "model-1": MagicMock(status="healthy", checked_at=datetime.now(timezone.utc) - timedelta(minutes=5)), + } + + persisted = await _save_health_check_results_if_changed( + mock_prisma, model_results, latest_checks_map, 1234567890.0, "background_health_check" + ) + + assert (persisted, mock_prisma.save_health_check_result.await_count) == (True, 0) + + +def _one_model_setup(): + model_list = [ + { + "model_name": "gpt-3.5-turbo", + "model_info": {"id": "model-123"}, + "litellm_params": {"model": "gpt-3.5-turbo"}, + }, + ] + return model_list, [{"model": "gpt-3.5-turbo"}], [] + + +@pytest.mark.asyncio +async def test_save_background_health_checks_to_db_returns_false_when_a_write_fails(): + mock_prisma = MagicMock() + mock_prisma.get_all_latest_health_checks = AsyncMock(return_value=[]) + mock_prisma.save_health_check_result = AsyncMock(return_value=None) + model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() + + persisted = await _save_background_health_checks_to_db( + mock_prisma, model_list, healthy_endpoints, unhealthy_endpoints, 1234567890.0, "background_health_check" + ) + + assert (persisted, mock_prisma.save_health_check_result.await_count) == (False, 1) + + @pytest.mark.asyncio async def test_save_background_health_checks_to_db_no_prisma(): """Test graceful handling when no prisma client""" - result = await _save_background_health_checks_to_db( - None, [], [], [], 0.0, "background_health_check" - ) - assert result is None + result = await _save_background_health_checks_to_db(None, [], [], [], 0.0, "background_health_check") + assert result is False @pytest.mark.asyncio async def test_save_background_health_checks_to_db_exception_handling(): """Test exception handling in background health check save""" mock_prisma = MagicMock() - mock_prisma.get_all_latest_health_checks = AsyncMock( - side_effect=Exception("DB Error") - ) + mock_prisma.get_all_latest_health_checks = AsyncMock(side_effect=Exception("DB Error")) model_list = [ { @@ -446,12 +525,13 @@ async def test_save_background_health_checks_to_db_exception_handling(): }, ] - # Should not raise exception, should handle gracefully - await _save_background_health_checks_to_db( + # Must not raise (the health check loop has to survive a DB outage) but must report + # the failure, so the window lock can be released for another pod to retry + persisted = await _save_background_health_checks_to_db( mock_prisma, model_list, [], [], 0.0, "background_health_check" ) - # Function should complete without raising + assert persisted is False def _raw_latest_row(model_name: str, model_id, checked_at: datetime) -> dict: @@ -660,12 +740,7 @@ def test_parse_background_health_check_model_groups_unset_returns_none(): assert parse_background_health_check_model_groups(None) is None assert parse_background_health_check_model_groups({}) is None - assert ( - parse_background_health_check_model_groups( - {"background_health_check_model_groups": None} - ) - is None - ) + assert parse_background_health_check_model_groups({"background_health_check_model_groups": None}) is None def test_parse_background_health_check_model_groups_list_returns_frozenset(): @@ -682,9 +757,7 @@ def test_parse_background_health_check_model_groups_malformed_raises(bad_value): from litellm.proxy.health_check import parse_background_health_check_model_groups with pytest.raises(ValueError, match="must be a list of model group names"): - parse_background_health_check_model_groups( - {"background_health_check_model_groups": bad_value} - ) + parse_background_health_check_model_groups({"background_health_check_model_groups": bad_value}) def test_filter_deployments_to_model_groups(): @@ -697,9 +770,7 @@ def test_filter_deployments_to_model_groups(): ] assert filter_deployments_to_model_groups(model_list, None) == tuple(model_list) - assert filter_deployments_to_model_groups( - model_list, frozenset({"prod-openai"}) - ) == (model_list[0], model_list[2]) + assert filter_deployments_to_model_groups(model_list, frozenset({"prod-openai"})) == (model_list[0], model_list[2]) assert filter_deployments_to_model_groups(model_list, frozenset()) == () From 7c1d060b7af3150320a2c88428871703d71f7a38 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Wed, 9 Sep 2026 18:04:36 -0700 Subject: [PATCH 28/60] fix(mcp): enforce OAuth identity binding across credential lifetime --- litellm/proxy/_experimental/mcp_server/db.py | 34 +++- .../mcp_server/discoverable_endpoints.py | 94 +++++++--- .../mcp_server/oauth_identity_binding.py | 170 ++++++++++++++---- .../authz_code_refresher.py | 37 +++- .../per_user_oauth_store.py | 13 ++ .../types/mcp_server/mcp_server_manager.py | 13 +- .../test_authz_code_refresher.py | 34 ++++ .../test_per_user_oauth_store.py | 25 +++ .../mcp_server/test_db_credentials.py | 24 ++- .../mcp_server/test_discoverable_endpoints.py | 60 ++++++- .../mcp_server/test_oauth_identity_binding.py | 111 +++++++++--- .../test_mcp_management_endpoints.py | 3 +- 12 files changed, 526 insertions(+), 92 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 7379126983a..08f83aeb756 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -6,10 +6,17 @@ from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict, cast +from typing_extensions import ReadOnly + from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.constants import MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( + RefreshTokenPresented, + credential_binding_matches, + enforce_oauth_identity_binding, +) from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request from litellm.proxy._types import ( LiteLLM_MCPServerTable, @@ -117,6 +124,7 @@ class _OAuthCredentialAccessToken(TypedDict): class OAuthCredentialPayload(_OAuthCredentialAccessToken, total=False): + identity_binding_proof: ReadOnly[str] type: str refresh_token: str expires_at: str @@ -1393,6 +1401,7 @@ async def store_user_oauth_credential( expires_in: int | None = None, scopes: list[str] | None = None, skip_byok_guard: bool = False, + identity_binding_proof: str | None = None, ) -> None: """Persist an OAuth2 access token for a user+server pair. @@ -1409,6 +1418,7 @@ async def store_user_oauth_credential( "type": "oauth2", "access_token": access_token, "connected_at": datetime.now(timezone.utc).isoformat(), + **({"identity_binding_proof": identity_binding_proof} if identity_binding_proof else {}), } if refresh_token: payload["refresh_token"] = refresh_token @@ -1628,6 +1638,11 @@ async def refresh_user_oauth_token( warning and returns ``None`` — the caller is responsible for clearing the stale credential and triggering re-authentication. """ + binding: Final = server.oauth_identity_binding + if binding is not None and binding.mode == "enforce": + if not await credential_binding_matches(binding, user_id, server.server_id, cred): + return None + refresh_token: Final[str | None] = cred.get("refresh_token") token_url: Final[str | None] = getattr(server, "effective_token_url", None) or getattr(server, "token_url", None) server_id: Final[str] = getattr(server, "server_id", "") @@ -1677,6 +1692,14 @@ async def refresh_user_oauth_token( ) return None + binding_proof: Final = await enforce_oauth_identity_binding( + server=server, + token_response=body, + litellm_user_id=user_id, + grant_type="refresh_token", + refresh_ownership=RefreshTokenPresented(refresh_token), + ) + access_token: Final[str | None] = body.get("access_token") if not access_token: verbose_proxy_logger.warning( @@ -1709,6 +1732,7 @@ async def refresh_user_oauth_token( refresh_token=new_refresh_token, expires_in=expires_in, scopes=scopes, + identity_binding_proof=binding_proof, skip_byok_guard=True, # Row is already OAuth2; skip the extra find_unique check ) @@ -1742,6 +1766,10 @@ async def resolve_valid_user_oauth_token( grant: Final = oauth_grant_state(cred) if cred is None or grant == "absent": return None + binding: Final = server.oauth_identity_binding + if binding is not None and binding.mode == "enforce": + if not await credential_binding_matches(binding, user_id, server.server_id, cred): + return None if grant == "valid": return cred if prisma_client is None: @@ -1782,7 +1810,11 @@ async def resolve_user_oauth_access_token( mcp_per_user_token_cache, ) - if prefetched_creds is None: + binding: Final = server.oauth_identity_binding + enforce_binding: Final = binding is not None and binding.mode == "enforce" + if enforce_binding: + await mcp_per_user_token_cache.delete(user_id, server_id) + if prefetched_creds is None and not enforce_binding: cached_token: Final = await mcp_per_user_token_cache.get(user_id, server_id) if cached_token is not None: return cached_token diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 9f6bc1860e1..c8e02be308b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -144,6 +144,7 @@ def encode_state_with_base_url( dcr_client_id: str | None = None, dcr_client_secret: str | None = None, dcr_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None, + oauth_nonce: str | None = None, ) -> str: """ Encode the base_url, original state, and PKCE parameters using encryption. @@ -154,9 +155,8 @@ def encode_state_with_base_url( code_challenge: PKCE code challenge from client code_challenge_method: PKCE code challenge method from client client_redirect_uri: Original redirect_uri from client - litellm_user_id: The SSO-authenticated litellm user captured at the bridge authorize - (interactive dcr_bridge oauth_delegate only); the callback seals it into the gateway - authorization code so the token mint can bind the envelope to this user + litellm_user_id: The authenticated user captured for bridge or identity-bound per-user OAuth; + the callback seals this credential owner into the authorization code mcp_server_id: The server the flow targets, sealed alongside litellm_user_id (bridge) or dcr_client_id (ephemeral mint) so the gateway code cannot be replayed against another server @@ -174,6 +174,7 @@ def encode_state_with_base_url( An encrypted string that encodes all values """ state_data: Final = { + "oauth_nonce": oauth_nonce, "base_url": base_url, "original_state": original_state, "code_challenge": code_challenge, @@ -215,10 +216,10 @@ _BRIDGE_AUTH_CODE_PREFIX: Final = "llm_bcode_" class _BridgeAuthorizationCode(BaseModel): - """The identity and upstream code the gateway seals into the authorization code it hands a DCR - client for an interactive dcr_bridge oauth_delegate sign-in, recovered at the token endpoint.""" + """Authenticated caller and upstream code sealed for bridge or identity-bound per-user OAuth.""" model_config = ConfigDict(frozen=True) + oauth_nonce: str | None = None upstream_code: str = Field(min_length=1) litellm_user_id: str = Field(min_length=1) mcp_server_id: str = Field(min_length=1) @@ -230,7 +231,12 @@ def is_bridge_authorization_code(code: str) -> bool: return code.startswith(_BRIDGE_AUTH_CODE_PREFIX) -def seal_bridge_authorization_code(upstream_code: str, litellm_user_id: str, mcp_server_id: str) -> str: +def seal_bridge_authorization_code( + upstream_code: str, + litellm_user_id: str, + mcp_server_id: str, + oauth_nonce: str | None = None, +) -> str: """Seal the upstream authorization code and the SSO-captured litellm user into a gateway authorization code. The DCR client only echoes this opaque value back at the token endpoint; the gateway decrypts it there to recover the user (to bind the envelope) and the upstream code (to @@ -239,7 +245,12 @@ def seal_bridge_authorization_code(upstream_code: str, litellm_user_id: str, mcp authenticated symmetric helper (the same family the OAuth state uses), so the client can neither read nor forge it.""" payload: Final = json.dumps( - {"upstream_code": upstream_code, "litellm_user_id": litellm_user_id, "mcp_server_id": mcp_server_id}, + { + "upstream_code": upstream_code, + "litellm_user_id": litellm_user_id, + "mcp_server_id": mcp_server_id, + "oauth_nonce": oauth_nonce, + }, sort_keys=True, ) return _BRIDGE_AUTH_CODE_PREFIX + encrypt_value_helper(payload) @@ -552,6 +563,7 @@ async def _store_per_user_token_server_side( server: MCPServer, user_id: str, token_response: dict[str, Any], + identity_binding_proof: str | None = None, ) -> None: """Persist the OAuth token server-side and warm the Redis cache. @@ -593,6 +605,7 @@ async def _store_per_user_token_server_side( refresh_token=refresh_token, expires_in=expires_in, scopes=scopes, + identity_binding_proof=identity_binding_proof, ) verbose_logger.info( "_store_per_user_token_server_side: stored token for user=%s server=%s", @@ -859,6 +872,11 @@ async def authorize_with_server( ), ) + binding: Final = resolved_server.oauth_identity_binding + enforce_binding: Final = binding is not None and binding.mode == "enforce" + if enforce_binding: + _require_s256_pkce(code_challenge, code_challenge_method) + if resolved_server.is_dcr_bridge: # Enforce S256 PKCE on both bridge arms. The relay arm forwards the validated, # now-non-optional pair to the upstream authorize; the short-circuit arm keeps @@ -889,19 +907,16 @@ async def authorize_with_server( base_url: Final = urlunparse(parsed._replace(query="")) request_base_url: Final = get_request_base_url(request) - # Interactive dcr_bridge oauth_delegate sign-in: this arm runs the gateway /callback and /token in - # the loop, so the gateway can capture the litellm user here (from the browser's UI session) and - # carry it to the back-channel token mint. Seal the SSO user and the target server into the state; - # the callback reads them back to mint the gateway authorization code. A DCR client cannot present a - # litellm key, so the browser session is the only identity source; without one there is nothing to - # bind, so send the user through login first. Every other oauth2 server keeps the identity-less state. + # Seal the authenticated caller into state so the token exchange cannot select another credential owner. litellm_user_id: str | None = None - if resolved_server.is_dcr_bridge and resolved_server.is_oauth_delegate: + if enforce_binding or (resolved_server.is_dcr_bridge and resolved_server.is_oauth_delegate): from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # inline import avoids a module-load circular import _user_id_from_session_cookie, ) - litellm_user_id = _user_id_from_session_cookie(request) + litellm_user_id = ( + await _extract_user_id_from_request(request) if enforce_binding else None + ) or _user_id_from_session_cookie(request) if litellm_user_id is None: return _redirect_to_litellm_login(request) denial: Final = await _bridge_authorize_access_denial( @@ -913,9 +928,11 @@ async def authorize_with_server( if denial is not None: return denial + oauth_nonce: Final = secrets.token_urlsafe(32) if enforce_binding else None encoded_state: Final = encode_state_with_base_url( base_url=base_url, original_state=state, + oauth_nonce=oauth_nonce, code_challenge=code_challenge, code_challenge_method=code_challenge_method, client_redirect_uri=redirect_uri, @@ -935,11 +952,16 @@ async def authorize_with_server( "state": relay_state, "response_type": response_type or "code", } + if oauth_nonce: + params["nonce"] = oauth_nonce if scope: params["scope"] = scope elif resolved_server.scopes: params["scope"] = " ".join(resolved_server.scopes) + if enforce_binding and "openid" not in params.get("scope", "").split(): + params["scope"] = f"openid {params.get('scope', '')}".strip() + if code_challenge: params["code_challenge"] = code_challenge if code_challenge_method: @@ -1020,6 +1042,12 @@ async def exchange_token_with_server( except TokenEndpointAuthConfigError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc + request_user_id: Final = ( + await _extract_user_id_from_request(request) + if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None + else None + ) + bridge_identity: _BridgeAuthorizationCode | None = None bridge_mint_ready: _BridgeMintReady | None = None bridge_upstream_refresh: SecretStr | None = None @@ -1081,6 +1109,14 @@ async def exchange_token_with_server( detail="Authorization code was issued for a different MCP server", ) code = bridge_identity.upstream_code + binding: Final = resolved_server.oauth_identity_binding + if binding is not None and binding.mode == "enforce": + if bridge_identity is None or not bridge_identity.oauth_nonce: + raise HTTPException(status_code=403, detail={"error": "oauth_identity_binding_failed"}) + if request_user_id is not None and request_user_id != bridge_identity.litellm_user_id: + raise HTTPException(status_code=403, detail={"error": "oauth_principal_mismatch"}) + if not code_verifier: + raise HTTPException(status_code=403, detail={"error": "oauth_identity_binding_failed"}) bridge_token_relay: Final = _dcr_bridge_relays_client_registration(resolved_server) if bridge_token_relay and not redirect_uri: raise HTTPException( @@ -1108,6 +1144,16 @@ async def exchange_token_with_server( return _bridge_mint_error_response(prepared) bridge_mint_ready = prepared + refresh_binding: Final = resolved_server.oauth_identity_binding + if grant_type == "refresh_token" and refresh_binding is not None and refresh_binding.mode == "enforce": + await enforce_oauth_identity_binding( + server=resolved_server, + token_response={}, + litellm_user_id=request_user_id, + grant_type=grant_type, + refresh_ownership=refresh_ownership, + ) + async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) try: response: Final = await async_client.post( @@ -1150,19 +1196,19 @@ async def exchange_token_with_server( # Bind the exchanged token to the LiteLLM caller BEFORE it is returned, stored, or cached, so a # token minted for a different upstream principal never becomes usable under the caller's user_id. - resolved_user_id: Final = ( - await _extract_user_id_from_request(request) - if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None - else None - ) - if isinstance(token_response, dict): + resolved_user_id: Final = bridge_identity.litellm_user_id if bridge_identity else request_user_id + binding_proof: Final = ( await enforce_oauth_identity_binding( server=resolved_server, token_response=token_response, litellm_user_id=resolved_user_id, grant_type=grant_type, refresh_ownership=refresh_ownership, + expected_nonce=bridge_identity.oauth_nonce if bridge_identity else None, ) + if isinstance(token_response, dict) + else None + ) # Store server-side when the server is configured for per-user OAuth and # the calling client has provided a valid LiteLLM identity. @@ -1175,6 +1221,7 @@ async def exchange_token_with_server( server=resolved_server, user_id=user_id, token_response=token_response, + identity_binding_proof=binding_proof, ) except Exception as exc: verbose_logger.warning( @@ -2161,7 +2208,10 @@ async def callback( forwarded_code = code if isinstance(litellm_user_id, str) and litellm_user_id and isinstance(mcp_server_id, str) and mcp_server_id: forwarded_code = seal_bridge_authorization_code( - upstream_code=code, litellm_user_id=litellm_user_id, mcp_server_id=mcp_server_id + upstream_code=code, + litellm_user_id=litellm_user_id, + mcp_server_id=mcp_server_id, + oauth_nonce=state_data.get("oauth_nonce"), ) elif isinstance(dcr_client_id, str) and dcr_client_id and isinstance(mcp_server_id, str) and mcp_server_id: forwarded_code = seal_passthrough_authorization_code( diff --git a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py index 58ec3cfdb4f..7ca9e6a23df 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py @@ -8,9 +8,12 @@ claim is compared to the caller's trusted LiteLLM identity. Mismatches fail clos and are logged in audit mode. """ +import hashlib +import hmac +import json from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass -from typing import Final, Literal, TypeAlias +from typing import Final, Literal, Protocol, TypeAlias import jwt from fastapi import HTTPException @@ -45,9 +48,17 @@ CallerPrincipalLoader: TypeAlias = Callable[ [str, MCPOAuthIdentityBinding], # mutable-ok: Callable parameter syntax requires a list Awaitable[str | None], ] + + +@dataclass(frozen=True, slots=True) +class VerifiedRefreshToken: + refresh_token: str + binding_proof: str + + StoredRefreshTokenLoader: TypeAlias = Callable[ - [str, str], # mutable-ok: Callable parameter syntax requires a list - Awaitable[str | None], + [str, str, MCPOAuthIdentityBinding], # mutable-ok: Callable parameter syntax requires a list + Awaitable[VerifiedRefreshToken | None], ] _RejectionCode: TypeAlias = Literal["oauth_principal_mismatch", "oauth_identity_binding_failed"] @@ -72,6 +83,18 @@ class RefreshTokenPresented: RefreshOwnership: TypeAlias = RefreshOwnershipProven | RefreshTokenPresented | None +class BindingValidator(Protocol): + async def __call__( + self, + *, + server: MCPServer, + token_response: Mapping[str, object], + litellm_user_id: str | None, + grant_type: str, + refresh_ownership: RefreshOwnership, + ) -> str | None: ... + + async def _fetch_issuer_jwks(binding: MCPOAuthIdentityBinding) -> Sequence[Mapping[str, object]]: jwks_url: Final[str] = binding.jwks_url or await _discover_jwks_url(binding.issuer) cached: Final = await _jwks_cache.async_get_cache(jwks_url) @@ -118,7 +141,7 @@ def _decode_id_token( signing_key: "jwt.PyJWK", ) -> "Mapping[str, object] | _BindingRejection": try: - decode_options: Final[Options] = {"require": ("iss", "exp")} + decode_options: Final[Options] = {"require": ("iss", "exp", "aud", "sub", "iat")} return jwt.decode( id_token, signing_key.key, @@ -169,7 +192,9 @@ async def _load_caller_principal(litellm_user_id: str, binding: MCPOAuthIdentity return loaded.user_email -async def _load_stored_refresh_token(litellm_user_id: str, server_id: str) -> str | None: +async def _load_stored_refresh_token( + litellm_user_id: str, server_id: str, binding: MCPOAuthIdentityBinding +) -> VerifiedRefreshToken | None: try: from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # keep database imports lazy get_user_oauth_credential, @@ -184,17 +209,89 @@ async def _load_stored_refresh_token(litellm_user_id: str, server_id: str) -> st user_id=litellm_user_id, server_id=server_id, ) - return cred.get("refresh_token") if cred else None + if not cred or not await credential_binding_matches(binding, litellm_user_id, server_id, cred): + return None + refresh_token: Final = cred.get("refresh_token") + proof: Final = cred.get("identity_binding_proof") + return VerifiedRefreshToken(refresh_token, proof) if refresh_token and proof else None except Exception: # noqa: BLE001 # a credential lookup failure must fail closed return None +async def current_binding_proof( + binding: MCPOAuthIdentityBinding, + user_id: str, + server_id: str, + caller_principal_loader: CallerPrincipalLoader = _load_caller_principal, +) -> str | None: + principal: Final = await caller_principal_loader(user_id, binding) + if not principal: + return None + return _binding_proof(binding, user_id, server_id, principal) + + +def _binding_proof(binding: MCPOAuthIdentityBinding, user_id: str, server_id: str, principal: str) -> str: + payload: Final = json.dumps( + ("oidc-nonce-v1", server_id, user_id, principal, binding.model_dump(mode="json")), + sort_keys=True, + ) + return hashlib.sha256(payload.encode()).hexdigest() + + +async def credential_binding_matches( + binding: MCPOAuthIdentityBinding, + user_id: str, + server_id: str, + credential: Mapping[str, object], + caller_principal_loader: CallerPrincipalLoader = _load_caller_principal, +) -> bool: + stored: Final = credential.get("identity_binding_proof") + if not isinstance(stored, str) or not stored: + return False + expected: Final = await current_binding_proof(binding, user_id, server_id, caller_principal_loader) + return expected is not None and hmac.compare_digest(stored, expected) + + def _principals_match(upstream: str, caller: str, binding: MCPOAuthIdentityBinding) -> bool: if binding.principal_claim == "email" or binding.caller_field == "user_email": return upstream.strip().casefold() == caller.strip().casefold() return upstream == caller +async def _evaluate_refresh_ownership( + binding: MCPOAuthIdentityBinding, + litellm_user_id: str | None, + server_id: str, + refresh_ownership: RefreshOwnership, + stored_refresh_token_loader: StoredRefreshTokenLoader, +) -> _BindingRejection | str: + match refresh_ownership: + case RefreshOwnershipProven(): + return _BindingRejection( + code="oauth_identity_binding_failed", + description="an identity envelope alone does not prove upstream principal binding", + ) + case None: + return _BindingRejection( + code="oauth_identity_binding_failed", + description="refresh_token grant without an id_token carries no refresh token to prove ownership of", + ) + case RefreshTokenPresented(refresh_token): + if not litellm_user_id: + return _BindingRejection( + code="oauth_identity_binding_failed", + description="the request carries no resolvable LiteLLM user identity to bind the credential to", + ) + stored: Final = await stored_refresh_token_loader(litellm_user_id, server_id, binding) + if stored is None or not hmac.compare_digest(stored.refresh_token, refresh_token): + return _BindingRejection( + code="oauth_identity_binding_failed", + description="the presented refresh_token is not the caller's stored credential for this server", + ) + return stored.binding_proof + assert_never(refresh_ownership) # pragma: no cover + + async def _evaluate_binding( binding: MCPOAuthIdentityBinding, token_response: Mapping[str, object], @@ -205,7 +302,8 @@ async def _evaluate_binding( jwks_fetcher: JwksFetcher, caller_principal_loader: CallerPrincipalLoader, stored_refresh_token_loader: StoredRefreshTokenLoader, -) -> _BindingRejection | None: + expected_nonce: str | None, +) -> _BindingRejection | str: id_token: Final = token_response.get("id_token") if not isinstance(id_token, str) or not id_token: if grant_type != "refresh_token": @@ -213,28 +311,13 @@ async def _evaluate_binding( code="oauth_identity_binding_failed", description="the upstream token response carries no id_token to bind the credential to a principal", ) - match refresh_ownership: - case RefreshOwnershipProven(): - return None - case None: - return _BindingRejection( - code="oauth_identity_binding_failed", - description="refresh_token grant without an id_token carries no refresh token to prove ownership of", - ) - case RefreshTokenPresented(refresh_token): - if not litellm_user_id: - return _BindingRejection( - code="oauth_identity_binding_failed", - description="the request carries no resolvable LiteLLM user identity to bind the credential to", - ) - stored: Final = await stored_refresh_token_loader(litellm_user_id, server_id) - if stored is None or stored != refresh_token: - return _BindingRejection( - code="oauth_identity_binding_failed", - description="the presented refresh_token is not the caller's stored credential for this server", - ) - return None - assert_never(refresh_ownership) # pragma: no cover + return await _evaluate_refresh_ownership( + binding, + litellm_user_id, + server_id, + refresh_ownership, + stored_refresh_token_loader, + ) if not litellm_user_id: return _BindingRejection( code="oauth_identity_binding_failed", @@ -247,12 +330,25 @@ async def _evaluate_binding( code="oauth_identity_binding_failed", description=f"could not fetch the issuer's JWKS: {exc}", ) - signing_key: Final = _select_signing_key(id_token, keys) + try: + signing_key: Final = _select_signing_key(id_token, keys) + except (jwt.PyJWTError, ValueError, TypeError): + return _BindingRejection( + code="oauth_identity_binding_failed", + description="invalid id_token header or issuer signing key", + ) if isinstance(signing_key, _BindingRejection): return signing_key claims: Final = _decode_id_token(id_token, binding, signing_key) if isinstance(claims, _BindingRejection): return claims + if grant_type == "authorization_code": + nonce: Final = claims.get("nonce") + if not expected_nonce or not isinstance(nonce, str) or not hmac.compare_digest(nonce, expected_nonce): + return _BindingRejection( + code="oauth_identity_binding_failed", + description="id_token nonce does not match the authenticated authorization transaction", + ) upstream: Final = _upstream_principal(claims, binding) if isinstance(upstream, _BindingRejection): return upstream @@ -267,7 +363,7 @@ async def _evaluate_binding( code="oauth_principal_mismatch", description="The browser account does not match the selected credential owner.", ) - return None + return _binding_proof(binding, litellm_user_id, server_id, caller) async def enforce_oauth_identity_binding( @@ -279,16 +375,17 @@ async def enforce_oauth_identity_binding( jwks_fetcher: JwksFetcher = _fetch_issuer_jwks, caller_principal_loader: CallerPrincipalLoader = _load_caller_principal, stored_refresh_token_loader: StoredRefreshTokenLoader = _load_stored_refresh_token, -) -> None: + expected_nonce: str | None = None, +) -> str | None: """Validate the exchanged token's upstream principal against the LiteLLM caller. No-op when the server has no binding or it is disabled. In enforce mode a failure raises 403 before the caller returns, stores, or caches the token; in audit mode failures are logged only. A refresh_token grant without an id_token is allowed only when the presented refresh token matches - the caller's stored credential or a sealed bridge envelope already proved ownership. + the caller's stored credential and that credential was previously identity-validated. """ binding: Final = server.oauth_identity_binding - if binding is None or binding.mode == "disabled": + if binding is None or binding.mode not in ("audit", "enforce"): return rejection: Final = await _evaluate_binding( binding=binding, @@ -300,9 +397,10 @@ async def enforce_oauth_identity_binding( jwks_fetcher=jwks_fetcher, caller_principal_loader=caller_principal_loader, stored_refresh_token_loader=stored_refresh_token_loader, + expected_nonce=expected_nonce, ) - if rejection is None: - return + if isinstance(rejection, str): + return rejection if binding.mode == "enforce" else None if binding.mode == "audit": verbose_logger.warning( "oauth_identity_binding audit: server=%s user=%s grant=%s rejected=%s (%s)", diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py index 92bd30694af..824459fd9b9 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py @@ -18,6 +18,11 @@ from litellm._logging import verbose_logger from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( TokenEndpointAuthConfigError, ) +from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( + BindingValidator, + RefreshTokenPresented, + enforce_oauth_identity_binding, +) from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( OAuthToken, @@ -39,6 +44,7 @@ class CredentialPersist(Protocol): refresh_token: str | None, expires_in: int | None, scopes: tuple[str, ...] | None, + identity_binding_proof: str | None = None, ) -> None: ... @@ -78,11 +84,13 @@ class AuthorizationCodeRefresher: persist: CredentialPersist, *, clock: Callable[[], float] = time.time, + identity_validator: BindingValidator = enforce_oauth_identity_binding, ) -> None: self._server_lookup = server_lookup self._token_endpoint = token_endpoint self._persist = persist self._clock = clock + self._identity_validator = identity_validator async def refresh(self, user_id: str, server_id: str, token: OAuthToken) -> OAuthToken | None: if token.refresh_token is None: @@ -104,6 +112,15 @@ class AuthorizationCodeRefresher: except TokenEndpointAuthConfigError as exc: verbose_logger.warning("MCP OAuth refresh misconfigured for server %s: %s", server_id, exc) return None + binding: Final = server.oauth_identity_binding + if binding is not None and binding.mode == "enforce": + await self._identity_validator( + server=server, + token_response={}, + litellm_user_id=user_id, + grant_type="refresh_token", + refresh_ownership=RefreshTokenPresented(token.refresh_token), + ) form: Final = { "grant_type": "refresh_token", "refresh_token": token.refresh_token, @@ -116,12 +133,30 @@ class AuthorizationCodeRefresher: if not isinstance(access_token, str) or not access_token: return None + binding_proof: Final = await self._identity_validator( + server=server, + token_response=body, + litellm_user_id=user_id, + grant_type="refresh_token", + refresh_ownership=RefreshTokenPresented(token.refresh_token), + ) rotated: Final = body.get("refresh_token") new_refresh: Final = rotated if isinstance(rotated, str) and rotated else token.refresh_token expires_in: Final = _parse_expires_in(body.get("expires_in")) scopes: Final = _parse_scopes(body.get("scope")) or token.scopes - await self._persist(user_id, server_id, access_token, new_refresh, expires_in, scopes or None) + if binding_proof is not None: + await self._persist( + user_id, + server_id, + access_token, + new_refresh, + expires_in, + scopes or None, + identity_binding_proof=binding_proof, + ) + else: + await self._persist(user_id, server_id, access_token, new_refresh, expires_in, scopes or None) return OAuthToken( access_token=access_token, expires_at=self._clock() + expires_in if expires_in is not None else None, diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py index 5e28396dcfb..efcc613c539 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py @@ -15,6 +15,7 @@ from collections.abc import Callable, Mapping from typing import TYPE_CHECKING, Final from litellm._logging import verbose_logger +from litellm.proxy._experimental.mcp_server.oauth_identity_binding import credential_binding_matches from litellm.proxy._experimental.mcp_server.outbound_credentials.authz_code_refresher import ( AuthorizationCodeRefresher, ) @@ -38,6 +39,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_cod OAuthTokenCacheCodec, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.v2_token_store import ( + CredentialReader, V2PerUserTokenStore, ) @@ -69,6 +71,7 @@ async def _persist_credential( refresh_token: str | None, expires_in: int | None, scopes: tuple[str, ...] | None, + identity_binding_proof: str | None = None, ) -> None: from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 store_user_oauth_credential, @@ -86,6 +89,7 @@ async def _persist_credential( expires_in=expires_in, scopes=list(scopes) if scopes else None, skip_byok_guard=True, + identity_binding_proof=identity_binding_proof, ) @@ -172,8 +176,10 @@ class LazyPerUserOAuthTokenStore: *, store_builder: StoreBuilder = _build_per_user_oauth_token_store, redis_available: Callable[[], bool] = _redis_cache_is_available, + credential_reader: CredentialReader = _read_credential, ) -> None: self._server_lookup = server_lookup + self._credential_reader = credential_reader self._store_builder = store_builder self._redis_available = redis_available self._store: InvalidatableOAuthTokenStore | None = None @@ -182,6 +188,13 @@ class LazyPerUserOAuthTokenStore: self._local_fetches = 0 async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: + server: Final = self._server_lookup(server_id) + binding: Final = server.oauth_identity_binding if server else None + if binding is not None and binding.mode == "enforce": + credential: Final = await self._credential_reader(user_id, server_id) + await self.invalidate(user_id, server_id) + if credential is None or not await credential_binding_matches(binding, user_id, server_id, credential): + return None if self._uses_redis: store = self._store if store is not None: diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 117d95a09de..f6b5441a626 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -1,7 +1,8 @@ from datetime import datetime from typing import Any, Final, Literal -from pydantic import BaseModel, ConfigDict, Field, field_validator +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator +from typing_extensions import Self from litellm.types.mcp import ( DEFAULT_SUBJECT_TOKEN_TYPE, @@ -42,7 +43,7 @@ class MCPOAuthIdentityBinding(BaseModel): """Per-server policy binding stored per-user OAuth credentials to the authenticated LiteLLM caller. When enabled for an interactive oauth2 server, the token relay validates the upstream OIDC - ``id_token`` (signature via the pinned issuer's JWKS, issuer, audience, expiry) and compares its + ``id_token`` (signature via the pinned issuer's JWKS, issuer, audience, expiry, nonce) and compares its principal claim to the LiteLLM caller's trusted identity before the token is returned, stored, or cached. ``audit`` logs mismatches without changing behavior; ``enforce`` fails closed with 403 ``oauth_principal_mismatch`` and disables the direct ``oauth-user-credential`` POST, which @@ -246,6 +247,14 @@ class MCPServer(BaseModel): """ return self.oauth2_flow == "client_credentials" + @model_validator(mode="after") + def validate_identity_binding_mode(self) -> Self: + binding: Final = self.oauth_identity_binding + if binding is not None and binding.mode != "disabled": + if not self.needs_user_oauth_token or self.delegate_auth_to_upstream: + raise ValueError("oauth_identity_binding requires gateway-managed per-user OAuth2 credentials") + return self + @property def needs_user_oauth_token(self) -> bool: """True if this is an OAuth2 server that relies on per-user tokens (no client_credentials).""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py index bb2f2ff8b02..16707f6ace8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py @@ -22,6 +22,7 @@ class _Server: server_id="srv", configured_token_url=None, ): + self.oauth_identity_binding = None self.token_url = token_url self.configured_token_url = configured_token_url self.client_id = client_id @@ -287,3 +288,36 @@ async def test_refresh_uses_admin_entered_token_url_when_issuer_yield_empties_re assert token is not None assert token.access_token == "new-at" assert posted[0][0] == "https://idp.example.com/token" + + +@pytest.mark.asyncio +async def test_identity_rejection_never_persists_or_returns_refreshed_token(): + from unittest.mock import AsyncMock + from fastapi import HTTPException + + validator = AsyncMock(side_effect=HTTPException(status_code=403, detail="oauth_principal_mismatch")) + persist = AsyncMock() + refresher = AuthorizationCodeRefresher( + _lookup(_Server()), _endpoint({"access_token": "foreign-token", "id_token": "foreign-identity"}), + persist, identity_validator=validator, + ) + with pytest.raises(HTTPException) as error: + await refresher.refresh("alice", "srv", OAuthToken(access_token="old", refresh_token="old-rt")) + assert error.value.status_code == 403 + persist.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_verified_refresh_preserves_binding_proof_in_storage(): + from unittest.mock import AsyncMock + + validator = AsyncMock(return_value="verified-binding") + persist = AsyncMock() + refresher = AuthorizationCodeRefresher( + _lookup(_Server()), _endpoint({"access_token": "new", "refresh_token": "rotated"}), + persist, identity_validator=validator, + ) + token = await refresher.refresh("alice", "srv", OAuthToken(access_token="old", refresh_token="old-rt")) + assert token.access_token == "new" + assert token.refresh_token == "rotated" + assert persist.await_args.kwargs["identity_binding_proof"] == "verified-binding" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py index ca32cf2bb8d..bd190b09c82 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py @@ -246,3 +246,28 @@ async def test_lazy_store_invalidate_works_after_redis_chain_is_built() -> None: assert build_calls == 1 assert redis_store.invalidations == [("u", "s")] + + +@pytest.mark.asyncio +async def test_enforcement_invalidates_cached_legacy_credentials_before_use(): + from litellm.types.mcp import MCPAuth, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + oauth_identity_binding=MCPOAuthIdentityBinding( + mode="enforce", issuer="https://idp.example.com", audiences=["client"], + ), + ) + cached = _RecordingStore("belongs-to-bob") + + async def read_legacy(user_id, server_id): + return {"access_token": "belongs-to-bob", "refresh_token": "bobs-refresh"} + + store = LazyPerUserOAuthTokenStore( + lambda server_id: server, store_builder=lambda lookup: (cached, False), + redis_available=lambda: False, credential_reader=read_legacy, + ) + assert await store.fetch("alice", "srv") is None + assert cached.calls == [] + assert cached.invalidations == [("alice", "srv")] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py index 355b3bfd30e..8e2adbeb06e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py @@ -628,6 +628,7 @@ async def test_oauth_round_trip_returns_payload(): access_token, refresh_token="rfr-xyz", scopes=["a", "b"], + identity_binding_proof="verified-proof", ) stored = _stored_value(prisma) @@ -642,6 +643,7 @@ async def test_oauth_round_trip_returns_payload(): assert result["access_token"] == access_token assert result["refresh_token"] == "rfr-xyz" assert result["scopes"] == ["a", "b"] + assert result["identity_binding_proof"] == "verified-proof" @pytest.mark.asyncio @@ -1184,6 +1186,7 @@ class _RefreshResponse: def _refresh_server(**overrides): base = dict( + oauth_identity_binding=None, token_url="https://idp.example.com/token", server_id="srv-1", client_id="cid", @@ -1309,7 +1312,7 @@ async def test_refresh_user_oauth_token_uses_client_secret_basic(monkeypatch): sends HTTP Basic and keeps the secret out of the body.""" import litellm.proxy._experimental.mcp_server.db as db_mod - server = MagicMock() + server = MagicMock(oauth_identity_binding=None) server.token_url = "https://idp.example.com/oauth2/token" server.server_id = "srv" server.client_id = "cid" @@ -1348,7 +1351,7 @@ async def test_refresh_user_oauth_token_defaults_to_client_secret_post(monkeypat the body (client_secret_post) and sends no Authorization header.""" import litellm.proxy._experimental.mcp_server.db as db_mod - server = MagicMock() + server = MagicMock(oauth_identity_binding=None) server.token_url = "https://idp.example.com/oauth2/token" server.server_id = "srv" server.client_id = "cid" @@ -1609,3 +1612,20 @@ def test_partial_update_defers_omitted_eligibility_fields_to_the_stored_row(): request = UpdateMCPServerRequest(server_id="relay-update", per_server_oauth_discovery=True) assert request.per_server_oauth_discovery is True + + +@pytest.mark.asyncio +async def test_enforcement_rejects_preexisting_unverified_credential(): + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + + server = MCPServer( + server_id="srv-1", name="srv-1", url="https://mcp.example.com", transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + oauth_identity_binding=MCPOAuthIdentityBinding( + mode="enforce", issuer="https://idp.example.com", audiences=["client"], + ), + ) + result = await resolve_valid_user_oauth_token( + user_id="alice", server=server, + cred={"access_token": "belongs-to-bob", "refresh_token": "bobs-refresh-token"}, + ) + assert result is None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 590d6606154..42187cb1771 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -8194,15 +8194,17 @@ async def test_token_exchange_refresh_passes_presented_refresh_ownership(): @pytest.mark.asyncio -async def test_token_exchange_authorization_code_passes_no_refresh_ownership(): +async def test_token_exchange_authorization_code_passes_no_refresh_ownership(monkeypatch): from fastapi import Request from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( exchange_token_with_server, + seal_bridge_authorization_code, ) from litellm.proxy._types import MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + monkeypatch.setenv("LITELLM_SALT_KEY", "identity-binding-test-salt") server = MCPServer( server_id="srv-1", name="srv-1", @@ -8244,16 +8246,17 @@ async def test_token_exchange_authorization_code_passes_no_refresh_ownership(): request=request, mcp_server=server, grant_type="authorization_code", - code="auth-code", + code=seal_bridge_authorization_code("auth-code", "user-a", "srv-1", "login-nonce"), redirect_uri="https://litellm.example.com/callback", client_id="cid", client_secret=None, - code_verifier=None, + code_verifier="test-verifier", ) assert result.status_code == 200 assert json.loads(result.body)["access_token"] == "at" assert enforce.await_args.kwargs["refresh_ownership"] is None + assert enforce.await_args.kwargs["expected_nonce"] == "login-nonce" def _upstream_token_response(status_code: int, *, json_body: object = None, text_body: str = "") -> "httpx.Response": @@ -11049,3 +11052,54 @@ def test_introspect_route_answers_for_authenticated_caller(monkeypatch): assert active.status_code == 200 assert active.json()["active"] is True assert active.json()["sub"] == "u1" + + +@pytest.mark.asyncio +async def test_identity_bound_authorization_carries_nonce_and_caller_through_callback(monkeypatch): + from http.cookies import SimpleCookie + from urllib.parse import parse_qs, urlparse + from fastapi import Request + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _oauth_state_cookie_name, authorize_with_server, callback, open_bridge_authorization_code, + ) + from litellm.types.mcp import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + + monkeypatch.setenv("LITELLM_SALT_KEY", "identity-binding-test-salt") + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + client_id="client", authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + oauth_identity_binding=MCPOAuthIdentityBinding( + mode="enforce", issuer="https://idp.example.com", audiences=["client"], + ), + ) + request = Request({"type": "http", "scheme": "https", "server": ("proxy.example.com", 443), + "path": "/authorize", "query_string": b"", "headers": []}) + with ( + patch( # test-quality-ok: isolate authenticated request resolution from the real encrypted OAuth round trip + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + new=AsyncMock(return_value="alice")), + patch( # test-quality-ok: isolate user access lookup while testing nonce and caller preservation + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._bridge_authorize_access_denial", + new=AsyncMock(return_value=None)), + ): + authorized = await authorize_with_server( + request, server, "client", "http://127.0.0.1:6274/callback", state="client-state", + code_challenge="pkce-challenge", code_challenge_method="S256", + ) + query = parse_qs(urlparse(authorized.headers["location"]).query) + assert len(query["nonce"][0]) >= 32 + cookies = SimpleCookie() + cookies.load(authorized.headers["set-cookie"]) + name = _oauth_state_cookie_name(query["state"][0]) + callback_request = Request({**request.scope, "path": "/callback", + "headers": [(b"cookie", f"{name}={cookies[name].value}".encode())]}) + completed = await callback(callback_request, code="upstream-code", state=query["state"][0]) + returned = parse_qs(urlparse(completed.headers["location"]).query) + sealed = open_bridge_authorization_code(returned["code"][0]) + assert sealed.litellm_user_id == "alice" + assert sealed.mcp_server_id == "srv" + assert sealed.upstream_code == "upstream-code" + assert sealed.oauth_nonce == query["nonce"][0] + assert returned["state"] == ["client-state"] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py index 36364c32ae5..6f291865e20 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py @@ -13,7 +13,9 @@ from pydantic import ValidationError from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( RefreshOwnershipProven, + VerifiedRefreshToken, RefreshTokenPresented, + current_binding_proof, _discover_jwks_url, _fetch_issuer_jwks, _load_caller_principal, @@ -21,7 +23,7 @@ from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( _select_signing_key, enforce_oauth_identity_binding, ) -from litellm.types.mcp import MCPTransport +from litellm.types.mcp import MCPAuth, MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer ISSUER: Final = "https://idp.example.com" @@ -48,6 +50,8 @@ def _sign_id_token(claims: Mapping[str, object]) -> str: "aud": AUDIENCE, "exp": int(time.time()) + 300, "iat": int(time.time()), + "nonce": "test-nonce", + "sub": "upstream-user", **claims, } return jwt.encode(payload, _PRIVATE_PEM, algorithm="RS256", headers={"kid": KID}) @@ -65,8 +69,8 @@ def _caller_loader(email: str | None): def _stored_refresh_token_loader(refresh_token: str | None): - async def load(_user_id: str, _server_id: str) -> str | None: - return refresh_token + async def load(_user_id: str, _server_id: str, _binding: MCPOAuthIdentityBinding) -> VerifiedRefreshToken | None: + return VerifiedRefreshToken(refresh_token, "verified-binding") if refresh_token else None return load @@ -77,6 +81,7 @@ def _server(mode: str = "enforce", **binding_overrides: object) -> MCPServer: name="srv-1", url="https://mcp.example.com", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth_identity_binding=MCPOAuthIdentityBinding( mode=mode, issuer=ISSUER, @@ -181,7 +186,9 @@ async def test_load_caller_principal_supports_user_id_and_database_email(): @pytest.mark.asyncio async def test_load_stored_refresh_token_returns_credential_and_fails_closed(): - get_credential: Final = AsyncMock(return_value={"refresh_token": "rt-1"}) + binding: Final = _server(caller_field="user_id", principal_claim="sub").oauth_identity_binding + proof: Final = await current_binding_proof(binding, "user-a", "srv-1") + get_credential: Final = AsyncMock(return_value={"refresh_token": "rt-1", "identity_binding_proof": proof}) with ( patch( # test-quality-ok: stored-token loading is a lazy database boundary without injection "litellm.proxy._experimental.mcp_server.db.get_user_oauth_credential", @@ -192,13 +199,15 @@ async def test_load_stored_refresh_token_returns_credential_and_fails_closed(): return_value="prisma", ), ): - assert await _load_stored_refresh_token("user-a", "srv-1") == "rt-1" + assert await _load_stored_refresh_token("user-a", "srv-1", binding) == VerifiedRefreshToken("rt-1", proof) + get_credential.return_value = {"refresh_token": "rt-1"} + assert await _load_stored_refresh_token("user-a", "srv-1", binding) is None with patch( # test-quality-ok: stored-token loading is a lazy database boundary without injection "litellm.proxy.utils.get_prisma_client_or_throw", side_effect=RuntimeError("database unavailable"), ): - assert await _load_stored_refresh_token("user-a", "srv-1") is None + assert await _load_stored_refresh_token("user-a", "srv-1", binding) is None @pytest.mark.asyncio @@ -209,11 +218,12 @@ async def test_matching_principal_passes(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) - assert result is None + assert result is not None @pytest.mark.asyncio @@ -225,6 +235,7 @@ async def test_mismatched_principal_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), @@ -243,6 +254,7 @@ async def test_missing_upstream_principal_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), @@ -263,6 +275,7 @@ async def test_jwks_fetch_failure_is_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=fail, caller_principal_loader=_caller_loader("alice@example.com"), @@ -284,6 +297,7 @@ async def test_missing_signing_key_is_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=no_keys, caller_principal_loader=_caller_loader("alice@example.com"), @@ -301,6 +315,7 @@ async def test_missing_caller_principal_is_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader(None), @@ -334,11 +349,12 @@ async def test_user_id_principal_matching_uses_exact_comparison(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("user-a"), ) - assert result is None + assert result is not None @pytest.mark.asyncio @@ -349,6 +365,7 @@ async def test_missing_id_token_rejected_on_authorization_code(): token_response={"access_token": "at"}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), @@ -369,7 +386,7 @@ async def test_refresh_without_id_token_allowed_when_presented_token_matches_sto caller_principal_loader=_caller_loader("alice@example.com"), stored_refresh_token_loader=_stored_refresh_token_loader("rt-1"), ) - assert result is None + assert result is not None @pytest.mark.asyncio @@ -407,21 +424,18 @@ async def test_refresh_without_id_token_rejects_missing_stored_token(): @pytest.mark.asyncio -async def test_refresh_without_id_token_passes_when_bridge_proves_ownership(): - async def fail_if_called(_user_id: str, _server_id: str) -> str | None: - raise AssertionError("stored refresh token loader should not be called") - - result: Final = await enforce_oauth_identity_binding( - server=_server(), - token_response={"access_token": "at"}, - litellm_user_id="user-a", - grant_type="refresh_token", - refresh_ownership=RefreshOwnershipProven(), - jwks_fetcher=_jwks_fetcher, - caller_principal_loader=_caller_loader("alice@example.com"), - stored_refresh_token_loader=fail_if_called, - ) - assert result is None +async def test_identity_envelope_does_not_prove_upstream_binding(): + with pytest.raises(HTTPException) as error: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at"}, + litellm_user_id="user-a", + grant_type="refresh_token", + refresh_ownership=RefreshOwnershipProven(), + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert error.value.status_code == 403 @pytest.mark.asyncio @@ -464,6 +478,7 @@ async def test_audit_mode_logs_but_does_not_reject(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), @@ -480,6 +495,7 @@ async def test_unverified_email_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), @@ -503,6 +519,7 @@ async def test_wrong_issuer_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), @@ -519,6 +536,7 @@ async def test_no_litellm_identity_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id=None, grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), @@ -533,6 +551,7 @@ async def test_disabled_binding_is_noop(): token_response={"access_token": "at"}, litellm_user_id=None, grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader(None), @@ -553,6 +572,7 @@ async def test_no_binding_is_noop(): token_response={"access_token": "at"}, litellm_user_id=None, grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader(None), @@ -576,9 +596,52 @@ async def test_wrong_audience_rejected(): token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", + expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) assert exc_info.value.status_code == 403 assert exc_info.value.detail["error"] == "oauth_identity_binding_failed" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("nonce", [None, "another-login"]) +async def test_authorization_code_rejects_missing_or_foreign_nonce(nonce): + token = _sign_id_token({"email": "alice@example.com", "email_verified": True, "nonce": nonce}) + with pytest.raises(HTTPException) as error: + await enforce_oauth_identity_binding( + server=_server(), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + refresh_ownership=None, + expected_nonce="this-login", + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert error.value.status_code == 403 + assert error.value.detail["error"] == "oauth_identity_binding_failed" + + +@pytest.mark.asyncio +async def test_binding_proof_rejects_changed_user_or_policy(): + from litellm.proxy._experimental.mcp_server.oauth_identity_binding import credential_binding_matches + + binding = _server(caller_field="user_id", principal_claim="sub").oauth_identity_binding + proof = await current_binding_proof(binding, "alice", "srv-1") + credential = {"identity_binding_proof": proof} + assert await credential_binding_matches(binding, "alice", "srv-1", credential) + assert not await credential_binding_matches(binding, "bob", "srv-1", credential) + assert not await credential_binding_matches(binding, "alice", "other-server", credential) + changed = binding.model_copy(update={"audiences": ["different-client"]}) + assert not await credential_binding_matches(changed, "alice", "srv-1", credential) + + +@pytest.mark.parametrize("auth_type", [MCPAuth.oauth_delegate, MCPAuth.true_passthrough]) +def test_identity_binding_rejects_modes_without_gateway_credential_custody(auth_type): + with pytest.raises(ValidationError, match="gateway-managed per-user"): + MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, auth_type=auth_type, + oauth_identity_binding=_server().oauth_identity_binding, + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 6a20688c09d..5e00e7d75be 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -4807,7 +4807,7 @@ async def test_store_mcp_oauth_user_credential_blocked_when_identity_binding_enf """The direct opaque-token POST must be closed for enforce-mode identity-bound servers, otherwise it bypasses the token-relay principal check.""" from litellm.proxy._types import MCPOAuthUserCredentialRequest - from litellm.types.mcp import MCPTransport + from litellm.types.mcp import MCPAuth, MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer if not mgmt_endpoints.MCP_AVAILABLE: @@ -4824,6 +4824,7 @@ async def test_store_mcp_oauth_user_credential_blocked_when_identity_binding_enf name=server_id, url="https://mcp.example.com", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth_identity_binding=MCPOAuthIdentityBinding( mode="enforce", issuer="https://idp.example.com", From 1273e46b8bf4821643eb4b948517c95de34bf632 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 15:45:25 -0400 Subject: [PATCH 29/60] feat(redis): add ElastiCache IAM authentication --- litellm/_redis.py | 57 ++++- litellm/_redis_credential_provider.py | 85 +++++++- litellm/proxy/_types.py | 4 + .../cache_settings_endpoints.py | 36 ++++ .../coordination_redis_endpoints.py | 29 +++ tests/test_litellm/test_redis.py | 196 +++++++++++++++++- .../test_redis_credential_provider.py | 122 +++++++++++ 7 files changed, 520 insertions(+), 9 deletions(-) create mode 100644 tests/test_litellm/test_redis_credential_provider.py diff --git a/litellm/_redis.py b/litellm/_redis.py index 3e68d50cf16..e289cc3b6c3 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -24,6 +24,7 @@ from redis.credentials import CredentialProvider from litellm import get_secret, get_secret_str from litellm._redis_credential_provider import ( AzureADCredentialProvider, + ElastiCacheIAMCredentialProvider, GCPIAMCredentialProvider, _generate_gcp_iam_access_token, ) @@ -75,6 +76,10 @@ def _get_redis_kwargs(): "azure_client_id", "azure_tenant_id", "azure_client_secret", + "aws_iam_auth", + "aws_iam_user_name", + "aws_iam_cache_name", + "aws_iam_region", } available_args: Final = {x for x in _unwrapped_init_args(redis.Redis) if x not in exclude_args} | include_args @@ -270,6 +275,32 @@ def _redis_kwargs_from_environment(): return return_dict +def _is_true(value: object | None) -> bool: + return value is True or (isinstance(value, str) and value.lower() == "true") + + +def _build_elasticache_iam_provider(redis_kwargs: dict) -> ElastiCacheIAMCredentialProvider | None: + if not _is_true(redis_kwargs.get("aws_iam_auth")): + return None + + required_settings: Final = { + "aws_iam_user_name": redis_kwargs.get("aws_iam_user_name"), + "aws_iam_cache_name": redis_kwargs.get("aws_iam_cache_name"), + "aws_iam_region": redis_kwargs.get("aws_iam_region") + or get_secret_str("AWS_REGION") + or get_secret_str("AWS_DEFAULT_REGION"), + } + missing_settings: Final = tuple(name for name, value in required_settings.items() if not value) + if missing_settings: + raise ValueError("AWS ElastiCache IAM Redis authentication requires: " + ", ".join(missing_settings)) + + return ElastiCacheIAMCredentialProvider( + user_name=str(required_settings["aws_iam_user_name"]), + cache_name=str(required_settings["aws_iam_cache_name"]), + region=str(required_settings["aws_iam_region"]), + ) + + def create_gcp_iam_redis_connect_func( service_account: str, ssl_ca_certs: str | None = None, @@ -540,6 +571,7 @@ def _get_redis_client_logic(**env_overrides): _azure_redis_ad_token: Final = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN") _azure_ad_enabled: Final = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true" + _aws_iam_enabled: Final = _is_true(redis_kwargs.get("aws_iam_auth")) if _azure_ad_enabled and _gcp_service_account is not None: verbose_logger.warning( @@ -560,13 +592,24 @@ def _get_redis_client_logic(**env_overrides): azure_tenant_id=_azure_tenant_id, azure_client_secret=_azure_client_secret, ) - # Marker for async paths to detect Azure AD auth. The live credential - # object is attached separately as `_azure_credential` by - # `create_azure_ad_redis_connect_func`; the raw client_id/tenant_id/secret - # are intentionally NOT exposed on the function to avoid leaking - # credentials via inspection or logging. redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True + if _aws_iam_enabled and _gcp_service_account is not None: + verbose_logger.warning( + "Both GCP IAM (gcp_service_account) and AWS ElastiCache IAM (aws_iam_auth) are configured " + "for Redis. Using GCP IAM. Remove one to avoid misconfiguration." + ) + elif _aws_iam_enabled and _azure_ad_enabled: + verbose_logger.warning( + "Both Azure AD (azure_redis_ad_token) and AWS ElastiCache IAM (aws_iam_auth) are configured " + "for Redis. Using Azure AD. Remove one to avoid misconfiguration." + ) + elif _aws_iam_enabled: + aws_provider: Final = _build_elasticache_iam_provider(redis_kwargs) + if aws_provider is not None: + verbose_logger.debug("Setting up AWS ElastiCache IAM authentication for Redis.") + redis_kwargs["credential_provider"] = aws_provider + redis_kwargs.pop("gcp_service_account", None) redis_kwargs.pop("gcp_ssl_ca_certs", None) @@ -575,6 +618,10 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("azure_client_id", None) redis_kwargs.pop("azure_tenant_id", None) redis_kwargs.pop("azure_client_secret", None) + redis_kwargs.pop("aws_iam_auth", None) + redis_kwargs.pop("aws_iam_user_name", None) + redis_kwargs.pop("aws_iam_cache_name", None) + redis_kwargs.pop("aws_iam_region", None) if redis_kwargs.get("credential_provider") is not None: redis_kwargs.pop("redis_connect_func", None) diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index ba0398789a6..a20bbfc530c 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -1,7 +1,8 @@ import asyncio import threading import time -from typing import Final, Protocol +from typing import Any, Final, Protocol +from urllib.parse import quote from redis.credentials import CredentialProvider @@ -117,6 +118,88 @@ class GCPIAMCredentialProvider(CredentialProvider): return (token,) +_ELASTICACHE_SERVICE_NAME: Final = "elasticache" +_ELASTICACHE_TOKEN_TTL_SECONDS: Final = 900 + + +class _FrozenBotocoreCredentials(Protocol): + access_key: str + secret_key: str + token: str | None + + +class _BotocoreCredentials(Protocol): + def get_frozen_credentials(self) -> _FrozenBotocoreCredentials: ... + + +class _BotocoreCredentialsResolver(Protocol): + def __call__(self) -> _BotocoreCredentials | None: ... + + +class ElastiCacheIAMCredentialProvider(CredentialProvider): + def __init__( + self, + user_name: str, + cache_name: str, + region: str, + credentials_resolver: _BotocoreCredentialsResolver | None = None, + token_lifetime_seconds: int = _ELASTICACHE_TOKEN_TTL_SECONDS, + ) -> None: + self._user_name = user_name + self._cache_name = cache_name + self._region = region + self._credentials_resolver = credentials_resolver or self._resolve_credentials + self._credentials: _BotocoreCredentials | None = None + self._token_lifetime_seconds = token_lifetime_seconds + + @staticmethod + def _resolve_credentials() -> Any: + try: + import botocore.session + except ImportError as e: + raise ImportError( + "botocore is required for ElastiCache IAM Redis authentication. Install it with: pip install boto3" + ) from e + + return botocore.session.get_session().get_credentials() + + def _get_credentials(self) -> tuple[str, str]: + credentials: Final = self._credentials if self._credentials is not None else self._credentials_resolver() + if credentials is None: + raise RuntimeError("Unable to resolve AWS credentials for ElastiCache IAM Redis authentication") + self._credentials = credentials + + frozen_credentials: Final = credentials.get_frozen_credentials() + if frozen_credentials is None: + raise RuntimeError("Unable to resolve AWS credentials for ElastiCache IAM Redis authentication") + + try: + from botocore.auth import SigV4QueryAuth + from botocore.awsrequest import AWSRequest + except ImportError as e: + raise ImportError( + "botocore is required for ElastiCache IAM Redis authentication. Install it with: pip install boto3" + ) from e + + request: Final = AWSRequest( + method="GET", + url=(f"https://{self._cache_name}/?Action=connect&User={quote(self._user_name, safe='')}"), + ) + SigV4QueryAuth( + frozen_credentials, + _ELASTICACHE_SERVICE_NAME, + self._region, + expires=self._token_lifetime_seconds, + ).add_auth(request) + return self._user_name, request.url.removeprefix("https://") + + def get_credentials(self) -> tuple[str, str]: + return self._get_credentials() + + async def get_credentials_async(self) -> tuple[str, str]: + return await asyncio.to_thread(self._get_credentials) + + class AzureADCredentialProvider(CredentialProvider): """ redis.credentials.CredentialProvider implementation that supplies Azure AD diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index c22bb76629d..99103ff3706 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2427,6 +2427,10 @@ class CoordinationRedisParams(LiteLLMPydanticObjectBase): ) sentinel_password: str | None = Field(None, description="password for the sentinel nodes") service_name: str | None = Field(None, description="sentinel service name") + aws_iam_auth: bool | str | None = Field(None, description="enable AWS ElastiCache IAM authentication") + aws_iam_user_name: str | None = Field(None, description="AWS ElastiCache IAM user name") + aws_iam_cache_name: str | None = Field(None, description="AWS ElastiCache cache name") + aws_iam_region: str | None = Field(None, description="AWS region for ElastiCache IAM authentication") def has_connection_target(self) -> bool: return any(value is not None for value in (self.host, self.url, self.startup_nodes, self.sentinel_nodes)) diff --git a/litellm/types/management_endpoints/cache_settings_endpoints.py b/litellm/types/management_endpoints/cache_settings_endpoints.py index 32b32991449..10aece6efa4 100644 --- a/litellm/types/management_endpoints/cache_settings_endpoints.py +++ b/litellm/types/management_endpoints/cache_settings_endpoints.py @@ -237,4 +237,40 @@ CACHE_SETTINGS_FIELDS: Final[list[CacheSettingsField]] = [ ui_field_name="SSL Check Hostname", redis_type=None, ), + CacheSettingsField( + field_name="aws_iam_auth", + field_type="Boolean", + field_value=None, + field_description="Enable AWS ElastiCache IAM authentication", + field_default=False, + ui_field_name="AWS IAM Authentication", + redis_type=None, + ), + CacheSettingsField( + field_name="aws_iam_user_name", + field_type="String", + field_value=None, + field_description="AWS ElastiCache IAM user name", + field_default=None, + ui_field_name="AWS IAM User Name", + redis_type=None, + ), + CacheSettingsField( + field_name="aws_iam_cache_name", + field_type="String", + field_value=None, + field_description="AWS ElastiCache cache name", + field_default=None, + ui_field_name="AWS IAM Cache Name", + redis_type=None, + ), + CacheSettingsField( + field_name="aws_iam_region", + field_type="String", + field_value=None, + field_description="AWS region for ElastiCache IAM authentication", + field_default=None, + ui_field_name="AWS IAM Region", + redis_type=None, + ), ] diff --git a/litellm/types/management_endpoints/coordination_redis_endpoints.py b/litellm/types/management_endpoints/coordination_redis_endpoints.py index 30033346ed7..86b32752cbb 100644 --- a/litellm/types/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/types/management_endpoints/coordination_redis_endpoints.py @@ -102,4 +102,33 @@ COORDINATION_REDIS_SETTINGS_FIELDS: Final[list[CoordinationRedisSettingsField]] ui_field_name="Service Name", section="sentinel", ), + CoordinationRedisSettingsField( + field_name="aws_iam_auth", + field_type="Boolean", + field_description="Enable AWS ElastiCache IAM authentication", + field_default=False, + ui_field_name="AWS IAM Authentication", + section="connection", + ), + CoordinationRedisSettingsField( + field_name="aws_iam_user_name", + field_type="String", + field_description="AWS ElastiCache IAM user name", + ui_field_name="AWS IAM User Name", + section="connection", + ), + CoordinationRedisSettingsField( + field_name="aws_iam_cache_name", + field_type="String", + field_description="AWS ElastiCache cache name", + ui_field_name="AWS IAM Cache Name", + section="connection", + ), + CoordinationRedisSettingsField( + field_name="aws_iam_region", + field_type="String", + field_description="AWS region for ElastiCache IAM authentication", + ui_field_name="AWS IAM Region", + section="connection", + ), ] diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index a96e8541e06..c99a3a074fb 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -24,6 +24,7 @@ from litellm._redis import ( ) from litellm._redis_credential_provider import ( AzureADCredentialProvider, + ElastiCacheIAMCredentialProvider, GCPIAMCredentialProvider, _token_cache, ) @@ -78,6 +79,8 @@ def clean_redis_environment(monkeypatch): "REDIS_URL", "REDIS_CLUSTER_NODES", "REDIS_SENTINEL_NODES", + "AWS_REGION", + "AWS_DEFAULT_REGION", *_get_redis_env_kwarg_mapping(), ): monkeypatch.delenv(var, raising=False) @@ -110,6 +113,17 @@ def test_credential_provider_is_not_environment_derived(): assert "credential_provider" not in mapping.values() +def test_aws_iam_settings_are_environment_derived(): + allowed = _get_redis_kwargs() + mapping = _get_redis_env_kwarg_mapping() + + assert {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} <= allowed + assert mapping["REDIS_AWS_IAM_AUTH"] == "aws_iam_auth" + assert mapping["REDIS_AWS_IAM_USER_NAME"] == "aws_iam_user_name" + assert mapping["REDIS_AWS_IAM_CACHE_NAME"] == "aws_iam_cache_name" + assert mapping["REDIS_AWS_IAM_REGION"] == "aws_iam_region" + + def test_sync_direct_preserves_credential_provider_identity(clean_redis_environment): provider = _StubCredentialProvider() @@ -300,6 +314,180 @@ def test_gcp_kwargs_never_survive_client_logic(clean_redis_environment, override assert "gcp_ssl_ca_certs" not in redis_kwargs +def test_aws_iam_environment_settings_install_provider(clean_redis_environment, monkeypatch): + monkeypatch.setenv("REDIS_AWS_IAM_AUTH", "true") + monkeypatch.setenv("REDIS_AWS_IAM_USER_NAME", "iam-user") + monkeypatch.setenv("REDIS_AWS_IAM_CACHE_NAME", "cache.example.com") + monkeypatch.setenv("REDIS_AWS_IAM_REGION", "us-east-1") + + redis_kwargs = _get_redis_client_logic(host="cache.example.com", port=6379) + + assert isinstance(redis_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) + assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() + + +def test_aws_iam_settings_are_removed_for_url_and_static_credentials(clean_redis_environment): + redis_kwargs = _get_redis_client_logic( + url="rediss://url-user:url-pass@cache.example.com:6380", + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + username="static-user", + password="static-password", + ) + + assert isinstance(redis_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) + assert redis_kwargs["url"] == "rediss://cache.example.com:6380" + assert "username" not in redis_kwargs + assert "password" not in redis_kwargs + assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() + + +@pytest.mark.parametrize("missing", ["aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"]) +def test_aws_iam_missing_setting_fails_closed(clean_redis_environment, missing): + settings = { + "aws_iam_auth": True, + "aws_iam_user_name": "iam-user", + "aws_iam_cache_name": "cache.example.com", + "aws_iam_region": "us-east-1", + } + settings[missing] = None + + with pytest.raises(ValueError, match=missing): + _get_redis_client_logic(host="cache.example.com", port=6379, **settings) + + +@pytest.mark.parametrize("region_var", ["AWS_REGION", "AWS_DEFAULT_REGION"]) +def test_aws_iam_region_falls_back_to_environment(clean_redis_environment, monkeypatch, region_var): + monkeypatch.setenv(region_var, "sa-east-1") + + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + ) + + provider = redis_kwargs["credential_provider"] + assert isinstance(provider, ElastiCacheIAMCredentialProvider) + assert provider._region == "sa-east-1" + + +def test_aws_iam_region_prefers_aws_region_over_default_region(clean_redis_environment, monkeypatch): + monkeypatch.setenv("AWS_REGION", "sa-east-1") + monkeypatch.setenv("AWS_DEFAULT_REGION", "eu-west-1") + + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + ) + + assert redis_kwargs["credential_provider"]._region == "sa-east-1" + + +def test_aws_iam_region_prefers_explicit_over_environment(clean_redis_environment, monkeypatch): + monkeypatch.setenv("AWS_REGION", "sa-east-1") + + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="explicit-region", + ) + + assert redis_kwargs["credential_provider"]._region == "explicit-region" + + +def test_aws_iam_settings_map_to_distinct_provider_fields(clean_redis_environment): + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + aws_iam_auth=True, + aws_iam_user_name="iam-user-value", + aws_iam_cache_name="iam-cache-value", + aws_iam_region="iam-region-value", + ) + + provider = redis_kwargs["credential_provider"] + assert isinstance(provider, ElastiCacheIAMCredentialProvider) + assert provider._user_name == "iam-user-value" + assert provider._cache_name == "iam-cache-value" + assert provider._region == "iam-region-value" + + +@pytest.mark.parametrize("aws_iam_auth", [False, "false"]) +def test_aws_iam_auth_disabled_does_not_install_provider(clean_redis_environment, aws_iam_auth): + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + aws_iam_auth=aws_iam_auth, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert "credential_provider" not in redis_kwargs + assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() + + +def test_explicit_provider_wins_over_aws_iam(clean_redis_environment): + provider = _StubCredentialProvider() + + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + credential_provider=provider, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert redis_kwargs["credential_provider"] is provider + assert "aws_iam_auth" not in redis_kwargs + + +def test_gcp_wins_over_aws_iam(clean_redis_environment): + with patch("litellm._redis.create_gcp_iam_redis_connect_func") as mock_gcp: + mock_gcp.return_value = _gcp_marker_callback() + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + gcp_service_account="sa@example.com", + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert "credential_provider" not in redis_kwargs + assert redis_kwargs["redis_connect_func"] is mock_gcp.return_value + + +def test_azure_wins_over_aws_iam(clean_redis_environment): + with patch("litellm._redis.create_azure_ad_redis_connect_func") as mock_azure: + mock_azure.return_value = MagicMock() + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + azure_redis_ad_token="true", + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert "credential_provider" not in redis_kwargs + assert redis_kwargs["redis_connect_func"] is mock_azure.return_value + + def test_provider_keeps_the_rest_of_the_url_intact(clean_redis_environment): provider = _StubCredentialProvider() @@ -1530,9 +1718,11 @@ def test_async_sentinel_keeps_the_credential_provider_off_the_monitors(markers, "redis_connect_func": SimpleNamespace(**markers), } - with patch("litellm._redis.async_redis.Sentinel") as mock_sentinel_cls: - with patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs): - get_redis_async_client() + with ( + patch("litellm._redis.async_redis.Sentinel") as mock_sentinel_cls, + patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs), + ): + get_redis_async_client() sentinel_kwargs = mock_sentinel_cls.call_args[1]["sentinel_kwargs"] assert sentinel_kwargs["password"] == sentinel_password diff --git a/tests/test_litellm/test_redis_credential_provider.py b/tests/test_litellm/test_redis_credential_provider.py new file mode 100644 index 00000000000..6299d4d42e0 --- /dev/null +++ b/tests/test_litellm/test_redis_credential_provider.py @@ -0,0 +1,122 @@ +import asyncio +from types import SimpleNamespace +from urllib.parse import parse_qs, urlsplit + +import pytest + +from litellm._redis_credential_provider import ( + ElastiCacheIAMCredentialProvider, + _BotocoreCredentials, +) + + +class _FakeCredentials: + def __init__(self, access_key: str) -> None: + self.access_key = access_key + self.secret_key = "synthetic-secret" + self.token = "synthetic-session-token" + + def get_frozen_credentials(self): + return self + + +class _FakeResolver: + def __init__(self, credentials: _BotocoreCredentials | None) -> None: + self.credentials = credentials + self.calls = 0 + + def __call__(self): + self.calls += 1 + return self.credentials + + +class _RotatingFakeCredentials: + def __init__(self) -> None: + self.calls = 0 + + def __bool__(self) -> bool: + return False + + def get_frozen_credentials(self): + self.calls += 1 + return SimpleNamespace( + access_key=f"AKIA-SYNTHETIC-{self.calls}", + secret_key="synthetic-secret", + token="synthetic-session-token", + ) + + +def test_elasticache_provider_signs_expected_query(): + resolver = _FakeResolver(_FakeCredentials("AKIA-SYNTHETIC")) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=resolver, + ) + + user_name, token = provider.get_credentials() + parsed = urlsplit("https://" + token) + query = parse_qs(parsed.query) + + assert user_name == "iam-user" + assert parsed.netloc == "cache.example.com" + assert query["Action"] == ["connect"] + assert query["User"] == ["iam-user"] + assert query["X-Amz-Expires"] == ["900"] + assert "elasticache" in query["X-Amz-Credential"][0] + assert query["X-Amz-Credential"][0].split("/")[2] == "us-east-1" + assert not token.startswith("https://") + + +def test_elasticache_provider_resolves_credentials_once_but_refreshes_signature(): + rotating_credentials = _RotatingFakeCredentials() + resolver = _FakeResolver(rotating_credentials) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=resolver, + ) + + first = provider.get_credentials() + second = provider.get_credentials() + async_result = asyncio.run(provider.get_credentials_async()) + + assert first[0] == second[0] == async_result[0] == "iam-user" + assert first[1] != second[1] + assert async_result[1] != second[1] + assert resolver.calls == 1 + assert rotating_credentials.calls == 3 + + +def test_elasticache_provider_reports_missing_credentials(): + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=_FakeResolver(None), + ) + + with pytest.raises(RuntimeError, match="Unable to resolve AWS credentials"): + provider.get_credentials() + + +def test_elasticache_provider_recovers_after_a_failed_resolution(): + resolver = _FakeResolver(None) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=resolver, + ) + + with pytest.raises(RuntimeError, match="Unable to resolve AWS credentials"): + provider.get_credentials() + + resolver.credentials = _FakeCredentials("AKIA-SYNTHETIC") + user_name, token = provider.get_credentials() + + assert user_name == "iam-user" + assert token + assert resolver.calls == 2 From 604b9edc533ec236eaa3563e61e5fade4bfd6279 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 15:58:02 -0400 Subject: [PATCH 30/60] test(redis): cover AWS IAM provider install on async cluster nodes No existing test combined aws_iam_auth with startup_nodes, the shape ElastiCache/Valkey Serverless deployments actually use. --- tests/test_litellm/test_redis.py | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index c99a3a074fb..21c23ba7866 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -488,6 +488,20 @@ def test_azure_wins_over_aws_iam(clean_redis_environment): assert redis_kwargs["redis_connect_func"] is mock_azure.return_value +def test_async_cluster_installs_aws_iam_provider(clean_redis_environment): + startup_nodes = [{"host": "cluster-node", "port": 6379}] + + client = get_redis_async_client( + startup_nodes=startup_nodes, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert isinstance(client.connection_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) + + def test_provider_keeps_the_rest_of_the_url_intact(clean_redis_environment): provider = _StubCredentialProvider() From f66891d80f4c7ce4d47a5ddb4ce499e5df21065a Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 19:19:01 -0400 Subject: [PATCH 31/60] fix(redis): type ElastiCache IAM configuration Generated with AI Co-Authored-By: Claude Code --- litellm/_redis.py | 41 ++++++++++--------- litellm/_redis_credential_provider.py | 30 ++++++-------- .../test_redis_credential_provider.py | 25 +++++------ ui/litellm-dashboard/src/lib/http/schema.d.ts | 20 +++++++++ 4 files changed, 65 insertions(+), 51 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index e289cc3b6c3..2fd8a7100e3 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -279,25 +279,24 @@ def _is_true(value: object | None) -> bool: return value is True or (isinstance(value, str) and value.lower() == "true") -def _build_elasticache_iam_provider(redis_kwargs: dict) -> ElastiCacheIAMCredentialProvider | None: - if not _is_true(redis_kwargs.get("aws_iam_auth")): - return None - - required_settings: Final = { - "aws_iam_user_name": redis_kwargs.get("aws_iam_user_name"), - "aws_iam_cache_name": redis_kwargs.get("aws_iam_cache_name"), - "aws_iam_region": redis_kwargs.get("aws_iam_region") - or get_secret_str("AWS_REGION") - or get_secret_str("AWS_DEFAULT_REGION"), - } - missing_settings: Final = tuple(name for name, value in required_settings.items() if not value) +def _build_elasticache_iam_provider( + user_name: object | None, + cache_name: object | None, + region: object | None, +) -> ElastiCacheIAMCredentialProvider: + required_settings: Final = ( + ("aws_iam_user_name", user_name), + ("aws_iam_cache_name", cache_name), + ("aws_iam_region", region), + ) + missing_settings: Final = tuple(name for name, value in required_settings if not value) if missing_settings: raise ValueError("AWS ElastiCache IAM Redis authentication requires: " + ", ".join(missing_settings)) return ElastiCacheIAMCredentialProvider( - user_name=str(required_settings["aws_iam_user_name"]), - cache_name=str(required_settings["aws_iam_cache_name"]), - region=str(required_settings["aws_iam_region"]), + user_name=str(user_name), + cache_name=str(cache_name), + region=str(region), ) @@ -605,10 +604,14 @@ def _get_redis_client_logic(**env_overrides): "for Redis. Using Azure AD. Remove one to avoid misconfiguration." ) elif _aws_iam_enabled: - aws_provider: Final = _build_elasticache_iam_provider(redis_kwargs) - if aws_provider is not None: - verbose_logger.debug("Setting up AWS ElastiCache IAM authentication for Redis.") - redis_kwargs["credential_provider"] = aws_provider + verbose_logger.debug("Setting up AWS ElastiCache IAM authentication for Redis.") + redis_kwargs["credential_provider"] = _build_elasticache_iam_provider( + user_name=redis_kwargs.get("aws_iam_user_name"), + cache_name=redis_kwargs.get("aws_iam_cache_name"), + region=redis_kwargs.get("aws_iam_region") + or get_secret_str("AWS_REGION") + or get_secret_str("AWS_DEFAULT_REGION"), + ) redis_kwargs.pop("gcp_service_account", None) redis_kwargs.pop("gcp_ssl_ca_certs", None) diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index a20bbfc530c..85759af2a88 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -1,11 +1,19 @@ +from __future__ import annotations + import asyncio import threading import time -from typing import Any, Final, Protocol +from collections.abc import Callable +from typing import TYPE_CHECKING, Any, Final, Protocol from urllib.parse import quote from redis.credentials import CredentialProvider +if TYPE_CHECKING: + from botocore.credentials import Credentials +else: + Credentials = Any # rebind-ok: runtime alias for the type-checking-only botocore import + # Azure AD scope for Redis Cache for Azure. AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default" @@ -122,38 +130,24 @@ _ELASTICACHE_SERVICE_NAME: Final = "elasticache" _ELASTICACHE_TOKEN_TTL_SECONDS: Final = 900 -class _FrozenBotocoreCredentials(Protocol): - access_key: str - secret_key: str - token: str | None - - -class _BotocoreCredentials(Protocol): - def get_frozen_credentials(self) -> _FrozenBotocoreCredentials: ... - - -class _BotocoreCredentialsResolver(Protocol): - def __call__(self) -> _BotocoreCredentials | None: ... - - class ElastiCacheIAMCredentialProvider(CredentialProvider): def __init__( self, user_name: str, cache_name: str, region: str, - credentials_resolver: _BotocoreCredentialsResolver | None = None, + credentials_resolver: Callable[[], Credentials | None] | None = None, token_lifetime_seconds: int = _ELASTICACHE_TOKEN_TTL_SECONDS, ) -> None: self._user_name = user_name self._cache_name = cache_name self._region = region self._credentials_resolver = credentials_resolver or self._resolve_credentials - self._credentials: _BotocoreCredentials | None = None + self._credentials: Credentials | None = None self._token_lifetime_seconds = token_lifetime_seconds @staticmethod - def _resolve_credentials() -> Any: + def _resolve_credentials() -> Credentials | None: try: import botocore.session except ImportError as e: diff --git a/tests/test_litellm/test_redis_credential_provider.py b/tests/test_litellm/test_redis_credential_provider.py index 6299d4d42e0..de4afb7c313 100644 --- a/tests/test_litellm/test_redis_credential_provider.py +++ b/tests/test_litellm/test_redis_credential_provider.py @@ -4,10 +4,7 @@ from urllib.parse import parse_qs, urlsplit import pytest -from litellm._redis_credential_provider import ( - ElastiCacheIAMCredentialProvider, - _BotocoreCredentials, -) +from litellm._redis_credential_provider import ElastiCacheIAMCredentialProvider class _FakeCredentials: @@ -20,16 +17,6 @@ class _FakeCredentials: return self -class _FakeResolver: - def __init__(self, credentials: _BotocoreCredentials | None) -> None: - self.credentials = credentials - self.calls = 0 - - def __call__(self): - self.calls += 1 - return self.credentials - - class _RotatingFakeCredentials: def __init__(self) -> None: self.calls = 0 @@ -46,6 +33,16 @@ class _RotatingFakeCredentials: ) +class _FakeResolver: + def __init__(self, credentials: _FakeCredentials | _RotatingFakeCredentials | None) -> None: + self.credentials = credentials + self.calls = 0 + + def __call__(self): + self.calls += 1 + return self.credentials + + def test_elasticache_provider_signs_expected_query(): resolver = _FakeResolver(_FakeCredentials("AKIA-SYNTHETIC")) provider = ElastiCacheIAMCredentialProvider( diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 6d62ce2b675..9ca0df396b7 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26357,6 +26357,26 @@ export interface components { * independently of the response-cache backend in `litellm_settings.cache_params`. */ CoordinationRedisParams: { + /** + * Aws Iam Auth + * @description enable AWS ElastiCache IAM authentication + */ + aws_iam_auth?: boolean | string | null; + /** + * Aws Iam Cache Name + * @description AWS ElastiCache cache name + */ + aws_iam_cache_name?: string | null; + /** + * Aws Iam Region + * @description AWS region for ElastiCache IAM authentication + */ + aws_iam_region?: string | null; + /** + * Aws Iam User Name + * @description AWS ElastiCache IAM user name + */ + aws_iam_user_name?: string | null; /** * Host * @description Redis hostname From bb043fe29a7f680f8a5aee1ba5998f3fdb26e561 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 19:48:03 -0400 Subject: [PATCH 32/60] test(redis): cover ElastiCache IAM failures Generated with AI Co-Authored-By: Claude Code --- tests/test_litellm/test_redis.py | 52 ++++++-------- .../test_redis_credential_provider.py | 70 +++++++++++++++++++ 2 files changed, 93 insertions(+), 29 deletions(-) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 21c23ba7866..43fb7563233 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -455,37 +455,33 @@ def test_explicit_provider_wins_over_aws_iam(clean_redis_environment): def test_gcp_wins_over_aws_iam(clean_redis_environment): - with patch("litellm._redis.create_gcp_iam_redis_connect_func") as mock_gcp: - mock_gcp.return_value = _gcp_marker_callback() - redis_kwargs = _get_redis_client_logic( - host="cache.example.com", - port=6379, - gcp_service_account="sa@example.com", - aws_iam_auth=True, - aws_iam_user_name="iam-user", - aws_iam_cache_name="cache.example.com", - aws_iam_region="us-east-1", - ) + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + gcp_service_account="sa@example.com", + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) assert "credential_provider" not in redis_kwargs - assert redis_kwargs["redis_connect_func"] is mock_gcp.return_value + assert redis_kwargs["redis_connect_func"]._gcp_service_account == "sa@example.com" def test_azure_wins_over_aws_iam(clean_redis_environment): - with patch("litellm._redis.create_azure_ad_redis_connect_func") as mock_azure: - mock_azure.return_value = MagicMock() - redis_kwargs = _get_redis_client_logic( - host="cache.example.com", - port=6379, - azure_redis_ad_token="true", - aws_iam_auth=True, - aws_iam_user_name="iam-user", - aws_iam_cache_name="cache.example.com", - aws_iam_region="us-east-1", - ) + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + azure_redis_ad_token="true", + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) assert "credential_provider" not in redis_kwargs - assert redis_kwargs["redis_connect_func"] is mock_azure.return_value + assert redis_kwargs["redis_connect_func"]._azure_redis_ad_token is True def test_async_cluster_installs_aws_iam_provider(clean_redis_environment): @@ -1732,11 +1728,9 @@ def test_async_sentinel_keeps_the_credential_provider_off_the_monitors(markers, "redis_connect_func": SimpleNamespace(**markers), } - with ( - patch("litellm._redis.async_redis.Sentinel") as mock_sentinel_cls, - patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs), - ): - get_redis_async_client() + with patch("litellm._redis.async_redis.Sentinel") as mock_sentinel_cls: + with patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs): + get_redis_async_client() sentinel_kwargs = mock_sentinel_cls.call_args[1]["sentinel_kwargs"] assert sentinel_kwargs["password"] == sentinel_password diff --git a/tests/test_litellm/test_redis_credential_provider.py b/tests/test_litellm/test_redis_credential_provider.py index de4afb7c313..336e30e12c6 100644 --- a/tests/test_litellm/test_redis_credential_provider.py +++ b/tests/test_litellm/test_redis_credential_provider.py @@ -1,4 +1,6 @@ import asyncio +import builtins +import sys from types import SimpleNamespace from urllib.parse import parse_qs, urlsplit @@ -87,6 +89,41 @@ def test_elasticache_provider_resolves_credentials_once_but_refreshes_signature( assert rotating_credentials.calls == 3 +def test_elasticache_provider_uses_botocore_session_credentials(monkeypatch): + credentials = _FakeCredentials("AKIA-SYNTHETIC") + monkeypatch.setattr("botocore.session.get_session", lambda: SimpleNamespace(get_credentials=lambda: credentials)) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + ) + + user_name, token = provider.get_credentials() + + assert user_name == "iam-user" + assert "AKIA-SYNTHETIC" in token + + +def test_elasticache_provider_reports_missing_botocore(monkeypatch): + original_import = builtins.__import__ + + def import_without_botocore(name, *args, **kwargs): + if name == "botocore.session": + raise ImportError("synthetic missing dependency") + return original_import(name, *args, **kwargs) + + monkeypatch.delitem(sys.modules, "botocore.session", raising=False) + monkeypatch.setattr(builtins, "__import__", import_without_botocore) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + ) + + with pytest.raises(ImportError, match="pip install boto3"): + provider.get_credentials() + + def test_elasticache_provider_reports_missing_credentials(): provider = ElastiCacheIAMCredentialProvider( user_name="iam-user", @@ -99,6 +136,39 @@ def test_elasticache_provider_reports_missing_credentials(): provider.get_credentials() +def test_elasticache_provider_reports_missing_frozen_credentials(): + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=_FakeResolver(SimpleNamespace(get_frozen_credentials=lambda: None)), + ) + + with pytest.raises(RuntimeError, match="Unable to resolve AWS credentials"): + provider.get_credentials() + + +def test_elasticache_provider_reports_missing_signing_dependency(monkeypatch): + original_import = builtins.__import__ + + def import_without_botocore_auth(name, *args, **kwargs): + if name == "botocore.auth": + raise ImportError("synthetic missing dependency") + return original_import(name, *args, **kwargs) + + monkeypatch.delitem(sys.modules, "botocore.auth", raising=False) + monkeypatch.setattr(builtins, "__import__", import_without_botocore_auth) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=_FakeResolver(_FakeCredentials("AKIA-SYNTHETIC")), + ) + + with pytest.raises(ImportError, match="pip install boto3"): + provider.get_credentials() + + def test_elasticache_provider_recovers_after_a_failed_resolution(): resolver = _FakeResolver(None) provider = ElastiCacheIAMCredentialProvider( From 921c263876600776be8dd69980f683b2c9ea3940 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 20:21:03 -0400 Subject: [PATCH 33/60] fix(proxy): validate coordination Redis mappings Generated with AI Co-Authored-By: Claude Code --- .../management_endpoints/coordination_redis_endpoints.py | 2 +- litellm/proxy/proxy_server.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py index 86ce336c7a3..88dc09ab001 100644 --- a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py @@ -140,7 +140,7 @@ def _merge_over_saved( def _validated_params(settings: Mapping[str, object]) -> CoordinationRedisParams: """Validate settings the way startup does: resolve env refs, then require a connection target.""" try: - params: Final = CoordinationRedisParams(**_resolve_env_refs(settings)) + params: Final = CoordinationRedisParams.model_validate(_resolve_env_refs(settings)) except ValidationError as e: invalid_fields: Final = sorted({str(error["loc"][0]) for error in e.errors() if error["loc"]}) raise HTTPException( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9269fd48e6c..c21df8ce931 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4906,7 +4906,7 @@ class ProxyConfig: if not isinstance(raw_params, dict): raise ValueError("general_settings.coordination_redis must be a mapping of Redis connection params") - coordination_params: Final = CoordinationRedisParams(**_resolve_coordination_redis_env_refs(raw_params)) + coordination_params: Final = CoordinationRedisParams.model_validate(_resolve_coordination_redis_env_refs(raw_params)) if not coordination_params.has_connection_target(): raise ValueError( "general_settings.coordination_redis needs a connection target: " @@ -9234,7 +9234,7 @@ class ProxyStartupEvent: if persisted is None: return None - coordination_params: Final = CoordinationRedisParams(**_resolve_coordination_redis_env_refs(persisted)) + coordination_params: Final = CoordinationRedisParams.model_validate(_resolve_coordination_redis_env_refs(persisted)) if not coordination_params.has_connection_target(): verbose_proxy_logger.warning( "coordination_redis saved in the database names no connection target; ignoring it." From 9763bd80ae442533d234e733cf334e3e1565681e Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 20:26:26 -0400 Subject: [PATCH 34/60] style(proxy): format coordination Redis validation Generated with AI Co-Authored-By: Claude Code --- litellm/proxy/proxy_server.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c21df8ce931..de10fde45c0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4906,7 +4906,9 @@ class ProxyConfig: if not isinstance(raw_params, dict): raise ValueError("general_settings.coordination_redis must be a mapping of Redis connection params") - coordination_params: Final = CoordinationRedisParams.model_validate(_resolve_coordination_redis_env_refs(raw_params)) + coordination_params: Final = CoordinationRedisParams.model_validate( + _resolve_coordination_redis_env_refs(raw_params) + ) if not coordination_params.has_connection_target(): raise ValueError( "general_settings.coordination_redis needs a connection target: " @@ -9234,7 +9236,9 @@ class ProxyStartupEvent: if persisted is None: return None - coordination_params: Final = CoordinationRedisParams.model_validate(_resolve_coordination_redis_env_refs(persisted)) + coordination_params: Final = CoordinationRedisParams.model_validate( + _resolve_coordination_redis_env_refs(persisted) + ) if not coordination_params.has_connection_target(): verbose_proxy_logger.warning( "coordination_redis saved in the database names no connection target; ignoring it." From a7ad6023d6a52ed035b638c4b3ab867c27a3b65b Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 21:14:07 -0400 Subject: [PATCH 35/60] ci: retry CodSpeed result upload Generated with AI Co-Authored-By: Claude Code From 0217a137fe3f8793163140fdb4d2f1fa4d132bbf Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Thu, 27 Aug 2026 13:49:12 -0400 Subject: [PATCH 36/60] fix(redis): require TLS for ElastiCache IAM auth --- litellm/_redis.py | 9 ++++ tests/test_litellm/test_redis.py | 78 ++++++++++++++++++++++++++++++++ 2 files changed, 87 insertions(+) diff --git a/litellm/_redis.py b/litellm/_redis.py index 2fd8a7100e3..50900b2977b 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -279,6 +279,13 @@ def _is_true(value: object | None) -> bool: return value is True or (isinstance(value, str) and value.lower() == "true") +def _uses_tls(redis_kwargs: Mapping[str, object]) -> bool: + if redis_kwargs.get("startup_nodes") is not None: + return _is_true(redis_kwargs.get("ssl")) + url: Final = redis_kwargs.get("url") + return urlsplit(url).scheme.lower() == "rediss" if isinstance(url, str) else _is_true(redis_kwargs.get("ssl")) + + def _build_elasticache_iam_provider( user_name: object | None, cache_name: object | None, @@ -604,6 +611,8 @@ def _get_redis_client_logic(**env_overrides): "for Redis. Using Azure AD. Remove one to avoid misconfiguration." ) elif _aws_iam_enabled: + if not _uses_tls(redis_kwargs): + raise ValueError("AWS ElastiCache IAM Redis authentication requires TLS") verbose_logger.debug("Setting up AWS ElastiCache IAM authentication for Redis.") redis_kwargs["credential_provider"] = _build_elasticache_iam_provider( user_name=redis_kwargs.get("aws_iam_user_name"), diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 43fb7563233..f2eda9a5753 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -319,6 +319,7 @@ def test_aws_iam_environment_settings_install_provider(clean_redis_environment, monkeypatch.setenv("REDIS_AWS_IAM_USER_NAME", "iam-user") monkeypatch.setenv("REDIS_AWS_IAM_CACHE_NAME", "cache.example.com") monkeypatch.setenv("REDIS_AWS_IAM_REGION", "us-east-1") + monkeypatch.setenv("REDIS_SSL", "true") redis_kwargs = _get_redis_client_logic(host="cache.example.com", port=6379) @@ -326,6 +327,77 @@ def test_aws_iam_environment_settings_install_provider(clean_redis_environment, assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() +@pytest.mark.parametrize( + "transport", + [ + pytest.param({"host": "cache.example.com", "port": 6379}, id="host_without_ssl"), + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": False}, id="host_ssl_false"), + pytest.param({"url": "redis://cache.example.com:6379", "ssl": True}, id="plaintext_url"), + pytest.param( + {"startup_nodes": [{"host": "cache.example.com", "port": 6379}]}, + id="cluster_without_ssl", + ), + pytest.param( + { + "url": "rediss://cache.example.com:6379", + "startup_nodes": [{"host": "cache.example.com", "port": 6379}], + }, + id="cluster_without_ssl_ignores_url_scheme", + ), + pytest.param( + { + "sentinel_nodes": [("sentinel.example.com", 26379)], + "service_name": "cache", + }, + id="sentinel_without_ssl", + ), + ], +) +def test_aws_iam_auth_rejects_non_tls_connections(clean_redis_environment, transport): + with pytest.raises(ValueError, match="requires TLS"): + _get_redis_client_logic( + **transport, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + +@pytest.mark.parametrize( + "transport", + [ + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": True}, id="host"), + pytest.param({"url": "rediss://cache.example.com:6379"}, id="url"), + pytest.param( + { + "startup_nodes": [{"host": "cache.example.com", "port": 6379}], + "ssl": True, + }, + id="cluster", + ), + pytest.param( + { + "sentinel_nodes": [("sentinel.example.com", 26379)], + "service_name": "cache", + "ssl": True, + }, + id="sentinel", + ), + ], +) +def test_aws_iam_auth_accepts_tls_connections(clean_redis_environment, transport): + redis_kwargs = _get_redis_client_logic( + **transport, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert isinstance(redis_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) + + def test_aws_iam_settings_are_removed_for_url_and_static_credentials(clean_redis_environment): redis_kwargs = _get_redis_client_logic( url="rediss://url-user:url-pass@cache.example.com:6380", @@ -351,6 +423,7 @@ def test_aws_iam_missing_setting_fails_closed(clean_redis_environment, missing): "aws_iam_user_name": "iam-user", "aws_iam_cache_name": "cache.example.com", "aws_iam_region": "us-east-1", + "ssl": True, } settings[missing] = None @@ -365,6 +438,7 @@ def test_aws_iam_region_falls_back_to_environment(clean_redis_environment, monke redis_kwargs = _get_redis_client_logic( host="cache.example.com", port=6379, + ssl=True, aws_iam_auth=True, aws_iam_user_name="iam-user", aws_iam_cache_name="cache.example.com", @@ -382,6 +456,7 @@ def test_aws_iam_region_prefers_aws_region_over_default_region(clean_redis_envir redis_kwargs = _get_redis_client_logic( host="cache.example.com", port=6379, + ssl=True, aws_iam_auth=True, aws_iam_user_name="iam-user", aws_iam_cache_name="cache.example.com", @@ -396,6 +471,7 @@ def test_aws_iam_region_prefers_explicit_over_environment(clean_redis_environmen redis_kwargs = _get_redis_client_logic( host="cache.example.com", port=6379, + ssl=True, aws_iam_auth=True, aws_iam_user_name="iam-user", aws_iam_cache_name="cache.example.com", @@ -409,6 +485,7 @@ def test_aws_iam_settings_map_to_distinct_provider_fields(clean_redis_environmen redis_kwargs = _get_redis_client_logic( host="cache.example.com", port=6379, + ssl=True, aws_iam_auth=True, aws_iam_user_name="iam-user-value", aws_iam_cache_name="iam-cache-value", @@ -489,6 +566,7 @@ def test_async_cluster_installs_aws_iam_provider(clean_redis_environment): client = get_redis_async_client( startup_nodes=startup_nodes, + ssl=True, aws_iam_auth=True, aws_iam_user_name="iam-user", aws_iam_cache_name="cache.example.com", From f1cd2b03caac2c544866be984089039aa1228956 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Mon, 31 Aug 2026 13:49:55 -0400 Subject: [PATCH 37/60] fix(redis): type signed ElastiCache IAM URL --- litellm/_redis_credential_provider.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index 85759af2a88..4d728ef21b6 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -4,7 +4,7 @@ import asyncio import threading import time from collections.abc import Callable -from typing import TYPE_CHECKING, Any, Final, Protocol +from typing import TYPE_CHECKING, Any, Final, Protocol, cast from urllib.parse import quote from redis.credentials import CredentialProvider @@ -185,7 +185,7 @@ class ElastiCacheIAMCredentialProvider(CredentialProvider): self._region, expires=self._token_lifetime_seconds, ).add_auth(request) - return self._user_name, request.url.removeprefix("https://") + return self._user_name, cast(str, request.url).removeprefix("https://") def get_credentials(self) -> tuple[str, str]: return self._get_credentials() From aca81c263ccd6660bffcf05115e4971eb92412c8 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Mon, 31 Aug 2026 14:07:04 -0400 Subject: [PATCH 38/60] fix(redis): validate signed ElastiCache IAM URL --- litellm/_redis_credential_provider.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index 4d728ef21b6..4441d318373 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -4,7 +4,7 @@ import asyncio import threading import time from collections.abc import Callable -from typing import TYPE_CHECKING, Any, Final, Protocol, cast +from typing import TYPE_CHECKING, Any, Final, Protocol from urllib.parse import quote from redis.credentials import CredentialProvider @@ -185,7 +185,10 @@ class ElastiCacheIAMCredentialProvider(CredentialProvider): self._region, expires=self._token_lifetime_seconds, ).add_auth(request) - return self._user_name, cast(str, request.url).removeprefix("https://") + signed_url: Final = request.url + if signed_url is None: + raise RuntimeError("Unable to generate AWS ElastiCache IAM credentials") + return self._user_name, signed_url.removeprefix("https://") def get_credentials(self) -> tuple[str, str]: return self._get_credentials() From cd4073895ff9884d6fa77337fe6ac4c62dc029ec Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 9 Sep 2026 10:03:26 -0400 Subject: [PATCH 39/60] refactor(redis): type botocore credentials without Any Deferred annotation evaluation keeps the type-checking-only botocore import off the runtime path, so the alias only reintroduced typing.Any, which the strict ruff budget now bans --- litellm/_redis_credential_provider.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index 4441d318373..15f625dc8b4 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -4,15 +4,13 @@ import asyncio import threading import time from collections.abc import Callable -from typing import TYPE_CHECKING, Any, Final, Protocol +from typing import TYPE_CHECKING, Final, Protocol from urllib.parse import quote from redis.credentials import CredentialProvider if TYPE_CHECKING: from botocore.credentials import Credentials -else: - Credentials = Any # rebind-ok: runtime alias for the type-checking-only botocore import # Azure AD scope for Redis Cache for Azure. AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default" From 2b0711e4f518dc796395972896aa5c3e6608b965 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 9 Sep 2026 10:06:22 -0400 Subject: [PATCH 40/60] style(redis): restore the Azure AD marker comment The comment documents that the raw Azure client id, tenant id and secret are deliberately kept off the connect function, so this branch should never have dropped it --- litellm/_redis.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/litellm/_redis.py b/litellm/_redis.py index 50900b2977b..d77043a6dfd 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -598,6 +598,11 @@ def _get_redis_client_logic(**env_overrides): azure_tenant_id=_azure_tenant_id, azure_client_secret=_azure_client_secret, ) + # Marker for async paths to detect Azure AD auth. The live credential + # object is attached separately as `_azure_credential` by + # `create_azure_ad_redis_connect_func`; the raw client_id/tenant_id/secret + # are intentionally NOT exposed on the function to avoid leaking + # credentials via inspection or logging. redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True if _aws_iam_enabled and _gcp_service_account is not None: From bf73c49d737b0ee1d4274a36a1846a8628de3478 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 9 Sep 2026 10:17:53 -0400 Subject: [PATCH 41/60] refactor(redis): drop unreachable frozen credentials guard --- litellm/_redis_credential_provider.py | 2 -- tests/test_litellm/test_redis_credential_provider.py | 12 ------------ 2 files changed, 14 deletions(-) diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index 15f625dc8b4..dd49a6793ec 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -162,8 +162,6 @@ class ElastiCacheIAMCredentialProvider(CredentialProvider): self._credentials = credentials frozen_credentials: Final = credentials.get_frozen_credentials() - if frozen_credentials is None: - raise RuntimeError("Unable to resolve AWS credentials for ElastiCache IAM Redis authentication") try: from botocore.auth import SigV4QueryAuth diff --git a/tests/test_litellm/test_redis_credential_provider.py b/tests/test_litellm/test_redis_credential_provider.py index 336e30e12c6..9a58722b04e 100644 --- a/tests/test_litellm/test_redis_credential_provider.py +++ b/tests/test_litellm/test_redis_credential_provider.py @@ -136,18 +136,6 @@ def test_elasticache_provider_reports_missing_credentials(): provider.get_credentials() -def test_elasticache_provider_reports_missing_frozen_credentials(): - provider = ElastiCacheIAMCredentialProvider( - user_name="iam-user", - cache_name="cache.example.com", - region="us-east-1", - credentials_resolver=_FakeResolver(SimpleNamespace(get_frozen_credentials=lambda: None)), - ) - - with pytest.raises(RuntimeError, match="Unable to resolve AWS credentials"): - provider.get_credentials() - - def test_elasticache_provider_reports_missing_signing_dependency(monkeypatch): original_import = builtins.__import__ From c4fc20bcf933606d0ffc91a47036dc16df7059ea Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 9 Sep 2026 14:43:58 -0400 Subject: [PATCH 42/60] fix(redis): accept every truthy flag and sign serverless ElastiCache caches The ElastiCache IAM gate read `aws_iam_auth` and `ssl` with a helper that only accepted the literal string "true", while the kwarg coercion that runs later accepts "true", "1" and "yes". Type coercion happens after the gate, so `REDIS_AWS_IAM_AUTH=1` silently skipped IAM auth and `REDIS_SSL=1` made the "requires TLS" check fail closed on a connection that was in fact TLS. Both helpers now share `_str_to_bool`. AWS signs serverless cache tokens with an extra `ResourceType=ServerlessCache` query parameter, so tokens minted for a serverless cache were rejected. Adds an `aws_iam_serverless` setting (`REDIS_AWS_IAM_SERVERLESS`) that puts the parameter into the signed URL, and lowercases the cache name because ElastiCache lowercases it at creation time. --- litellm/_redis.py | 51 ++++++------ litellm/_redis_credential_provider.py | 17 ++-- litellm/proxy/_types.py | 3 + .../cache_settings_endpoints.py | 9 ++ .../coordination_redis_endpoints.py | 8 ++ tests/test_litellm/test_redis.py | 82 +++++++++++++++++-- .../test_redis_credential_provider.py | 54 ++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 ++ 8 files changed, 192 insertions(+), 37 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index d77043a6dfd..c5acdcb038b 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -39,6 +39,14 @@ from ._logging import verbose_logger AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default" +_AWS_IAM_KWARG_NAMES: Final = ( + "aws_iam_auth", + "aws_iam_user_name", + "aws_iam_cache_name", + "aws_iam_region", + "aws_iam_serverless", +) + def _unwrapped_init_args(cls: type) -> frozenset[str]: """Every parameter on a single class's own ``__init__``, decorator-unwrapped. @@ -76,10 +84,7 @@ def _get_redis_kwargs(): "azure_client_id", "azure_tenant_id", "azure_client_secret", - "aws_iam_auth", - "aws_iam_user_name", - "aws_iam_cache_name", - "aws_iam_region", + *_AWS_IAM_KWARG_NAMES, } available_args: Final = {x for x in _unwrapped_init_args(redis.Redis) if x not in exclude_args} | include_args @@ -275,22 +280,25 @@ def _redis_kwargs_from_environment(): return return_dict -def _is_true(value: object | None) -> bool: - return value is True or (isinstance(value, str) and value.lower() == "true") +def _coerces_to_true(value: object | None) -> bool: + return _str_to_bool(value) if isinstance(value, str) else bool(value) def _uses_tls(redis_kwargs: Mapping[str, object]) -> bool: if redis_kwargs.get("startup_nodes") is not None: - return _is_true(redis_kwargs.get("ssl")) + return _coerces_to_true(redis_kwargs.get("ssl")) url: Final = redis_kwargs.get("url") - return urlsplit(url).scheme.lower() == "rediss" if isinstance(url, str) else _is_true(redis_kwargs.get("ssl")) + if isinstance(url, str): + return urlsplit(url).scheme.lower() == "rediss" + return _coerces_to_true(redis_kwargs.get("ssl")) -def _build_elasticache_iam_provider( - user_name: object | None, - cache_name: object | None, - region: object | None, -) -> ElastiCacheIAMCredentialProvider: +def _build_elasticache_iam_provider(redis_kwargs: Mapping[str, object]) -> ElastiCacheIAMCredentialProvider: + user_name: Final = redis_kwargs.get("aws_iam_user_name") + cache_name: Final = redis_kwargs.get("aws_iam_cache_name") + region: Final = ( + redis_kwargs.get("aws_iam_region") or get_secret_str("AWS_REGION") or get_secret_str("AWS_DEFAULT_REGION") + ) required_settings: Final = ( ("aws_iam_user_name", user_name), ("aws_iam_cache_name", cache_name), @@ -304,6 +312,7 @@ def _build_elasticache_iam_provider( user_name=str(user_name), cache_name=str(cache_name), region=str(region), + is_serverless=_coerces_to_true(redis_kwargs.get("aws_iam_serverless")), ) @@ -577,7 +586,7 @@ def _get_redis_client_logic(**env_overrides): _azure_redis_ad_token: Final = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN") _azure_ad_enabled: Final = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true" - _aws_iam_enabled: Final = _is_true(redis_kwargs.get("aws_iam_auth")) + _aws_iam_enabled: Final = _coerces_to_true(redis_kwargs.get("aws_iam_auth")) if _azure_ad_enabled and _gcp_service_account is not None: verbose_logger.warning( @@ -619,13 +628,7 @@ def _get_redis_client_logic(**env_overrides): if not _uses_tls(redis_kwargs): raise ValueError("AWS ElastiCache IAM Redis authentication requires TLS") verbose_logger.debug("Setting up AWS ElastiCache IAM authentication for Redis.") - redis_kwargs["credential_provider"] = _build_elasticache_iam_provider( - user_name=redis_kwargs.get("aws_iam_user_name"), - cache_name=redis_kwargs.get("aws_iam_cache_name"), - region=redis_kwargs.get("aws_iam_region") - or get_secret_str("AWS_REGION") - or get_secret_str("AWS_DEFAULT_REGION"), - ) + redis_kwargs["credential_provider"] = _build_elasticache_iam_provider(redis_kwargs) redis_kwargs.pop("gcp_service_account", None) redis_kwargs.pop("gcp_ssl_ca_certs", None) @@ -635,10 +638,8 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("azure_client_id", None) redis_kwargs.pop("azure_tenant_id", None) redis_kwargs.pop("azure_client_secret", None) - redis_kwargs.pop("aws_iam_auth", None) - redis_kwargs.pop("aws_iam_user_name", None) - redis_kwargs.pop("aws_iam_cache_name", None) - redis_kwargs.pop("aws_iam_region", None) + for aws_iam_key in _AWS_IAM_KWARG_NAMES: + redis_kwargs.pop(aws_iam_key, None) if redis_kwargs.get("credential_provider") is not None: redis_kwargs.pop("redis_connect_func", None) diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index dd49a6793ec..7d90f944657 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -5,7 +5,7 @@ import threading import time from collections.abc import Callable from typing import TYPE_CHECKING, Final, Protocol -from urllib.parse import quote +from urllib.parse import urlencode from redis.credentials import CredentialProvider @@ -126,6 +126,7 @@ class GCPIAMCredentialProvider(CredentialProvider): _ELASTICACHE_SERVICE_NAME: Final = "elasticache" _ELASTICACHE_TOKEN_TTL_SECONDS: Final = 900 +_ELASTICACHE_SERVERLESS_RESOURCE_TYPE: Final = "ServerlessCache" class ElastiCacheIAMCredentialProvider(CredentialProvider): @@ -134,12 +135,14 @@ class ElastiCacheIAMCredentialProvider(CredentialProvider): user_name: str, cache_name: str, region: str, + is_serverless: bool = False, credentials_resolver: Callable[[], Credentials | None] | None = None, token_lifetime_seconds: int = _ELASTICACHE_TOKEN_TTL_SECONDS, ) -> None: self._user_name = user_name - self._cache_name = cache_name + self._cache_name = cache_name.lower() self._region = region + self._is_serverless = is_serverless self._credentials_resolver = credentials_resolver or self._resolve_credentials self._credentials: Credentials | None = None self._token_lifetime_seconds = token_lifetime_seconds @@ -171,10 +174,14 @@ class ElastiCacheIAMCredentialProvider(CredentialProvider): "botocore is required for ElastiCache IAM Redis authentication. Install it with: pip install boto3" ) from e - request: Final = AWSRequest( - method="GET", - url=(f"https://{self._cache_name}/?Action=connect&User={quote(self._user_name, safe='')}"), + query: Final = urlencode( + ( + ("Action", "connect"), + ("User", self._user_name), + *((("ResourceType", _ELASTICACHE_SERVERLESS_RESOURCE_TYPE),) if self._is_serverless else ()), + ) ) + request: Final = AWSRequest(method="GET", url=f"https://{self._cache_name}/?{query}") SigV4QueryAuth( frozen_credentials, _ELASTICACHE_SERVICE_NAME, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 99103ff3706..07e4515059a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2431,6 +2431,9 @@ class CoordinationRedisParams(LiteLLMPydanticObjectBase): aws_iam_user_name: str | None = Field(None, description="AWS ElastiCache IAM user name") aws_iam_cache_name: str | None = Field(None, description="AWS ElastiCache cache name") aws_iam_region: str | None = Field(None, description="AWS region for ElastiCache IAM authentication") + aws_iam_serverless: bool | str | None = Field( + None, description="the ElastiCache cache is serverless rather than a self-designed cluster" + ) def has_connection_target(self) -> bool: return any(value is not None for value in (self.host, self.url, self.startup_nodes, self.sentinel_nodes)) diff --git a/litellm/types/management_endpoints/cache_settings_endpoints.py b/litellm/types/management_endpoints/cache_settings_endpoints.py index 10aece6efa4..cb05f3a50ac 100644 --- a/litellm/types/management_endpoints/cache_settings_endpoints.py +++ b/litellm/types/management_endpoints/cache_settings_endpoints.py @@ -273,4 +273,13 @@ CACHE_SETTINGS_FIELDS: Final[list[CacheSettingsField]] = [ ui_field_name="AWS IAM Region", redis_type=None, ), + CacheSettingsField( + field_name="aws_iam_serverless", + field_type="Boolean", + field_value=None, + field_description="The ElastiCache cache is serverless rather than a self-designed cluster", + field_default=False, + ui_field_name="AWS IAM Serverless Cache", + redis_type=None, + ), ] diff --git a/litellm/types/management_endpoints/coordination_redis_endpoints.py b/litellm/types/management_endpoints/coordination_redis_endpoints.py index 86b32752cbb..d70a921a22f 100644 --- a/litellm/types/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/types/management_endpoints/coordination_redis_endpoints.py @@ -131,4 +131,12 @@ COORDINATION_REDIS_SETTINGS_FIELDS: Final[list[CoordinationRedisSettingsField]] ui_field_name="AWS IAM Region", section="connection", ), + CoordinationRedisSettingsField( + field_name="aws_iam_serverless", + field_type="Boolean", + field_description="The ElastiCache cache is serverless rather than a self-designed cluster", + field_default=False, + ui_field_name="AWS IAM Serverless Cache", + section="connection", + ), ] diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index f2eda9a5753..5337ed825a7 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -113,15 +113,25 @@ def test_credential_provider_is_not_environment_derived(): assert "credential_provider" not in mapping.values() +_AWS_IAM_SETTINGS = { + "aws_iam_auth", + "aws_iam_user_name", + "aws_iam_cache_name", + "aws_iam_region", + "aws_iam_serverless", +} + + def test_aws_iam_settings_are_environment_derived(): allowed = _get_redis_kwargs() mapping = _get_redis_env_kwarg_mapping() - assert {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} <= allowed + assert _AWS_IAM_SETTINGS <= allowed assert mapping["REDIS_AWS_IAM_AUTH"] == "aws_iam_auth" assert mapping["REDIS_AWS_IAM_USER_NAME"] == "aws_iam_user_name" assert mapping["REDIS_AWS_IAM_CACHE_NAME"] == "aws_iam_cache_name" assert mapping["REDIS_AWS_IAM_REGION"] == "aws_iam_region" + assert mapping["REDIS_AWS_IAM_SERVERLESS"] == "aws_iam_serverless" def test_sync_direct_preserves_credential_provider_identity(clean_redis_environment): @@ -319,12 +329,15 @@ def test_aws_iam_environment_settings_install_provider(clean_redis_environment, monkeypatch.setenv("REDIS_AWS_IAM_USER_NAME", "iam-user") monkeypatch.setenv("REDIS_AWS_IAM_CACHE_NAME", "cache.example.com") monkeypatch.setenv("REDIS_AWS_IAM_REGION", "us-east-1") - monkeypatch.setenv("REDIS_SSL", "true") + monkeypatch.setenv("REDIS_AWS_IAM_SERVERLESS", "1") + monkeypatch.setenv("REDIS_SSL", "1") redis_kwargs = _get_redis_client_logic(host="cache.example.com", port=6379) - assert isinstance(redis_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) - assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() + provider = redis_kwargs["credential_provider"] + assert isinstance(provider, ElastiCacheIAMCredentialProvider) + assert provider._is_serverless is True + assert not _AWS_IAM_SETTINGS & redis_kwargs.keys() @pytest.mark.parametrize( @@ -332,6 +345,9 @@ def test_aws_iam_environment_settings_install_provider(clean_redis_environment, [ pytest.param({"host": "cache.example.com", "port": 6379}, id="host_without_ssl"), pytest.param({"host": "cache.example.com", "port": 6379, "ssl": False}, id="host_ssl_false"), + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "false"}, id="host_ssl_false_string"), + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "0"}, id="host_ssl_zero_string"), + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "no"}, id="host_ssl_no_string"), pytest.param({"url": "redis://cache.example.com:6379", "ssl": True}, id="plaintext_url"), pytest.param( {"startup_nodes": [{"host": "cache.example.com", "port": 6379}]}, @@ -368,6 +384,13 @@ def test_aws_iam_auth_rejects_non_tls_connections(clean_redis_environment, trans "transport", [ pytest.param({"host": "cache.example.com", "port": 6379, "ssl": True}, id="host"), + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "true"}, id="host_ssl_true_string"), + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "1"}, id="host_ssl_one_string"), + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "yes"}, id="host_ssl_yes_string"), + pytest.param( + {"startup_nodes": [{"host": "cache.example.com", "port": 6379}], "ssl": "1"}, + id="cluster_ssl_one_string", + ), pytest.param({"url": "rediss://cache.example.com:6379"}, id="url"), pytest.param( { @@ -413,7 +436,7 @@ def test_aws_iam_settings_are_removed_for_url_and_static_credentials(clean_redis assert redis_kwargs["url"] == "rediss://cache.example.com:6380" assert "username" not in redis_kwargs assert "password" not in redis_kwargs - assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() + assert not _AWS_IAM_SETTINGS & redis_kwargs.keys() @pytest.mark.parametrize("missing", ["aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"]) @@ -499,7 +522,7 @@ def test_aws_iam_settings_map_to_distinct_provider_fields(clean_redis_environmen assert provider._region == "iam-region-value" -@pytest.mark.parametrize("aws_iam_auth", [False, "false"]) +@pytest.mark.parametrize("aws_iam_auth", [None, False, "", "false", "0", "no"]) def test_aws_iam_auth_disabled_does_not_install_provider(clean_redis_environment, aws_iam_auth): redis_kwargs = _get_redis_client_logic( host="cache.example.com", @@ -511,7 +534,52 @@ def test_aws_iam_auth_disabled_does_not_install_provider(clean_redis_environment ) assert "credential_provider" not in redis_kwargs - assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() + assert not _AWS_IAM_SETTINGS & redis_kwargs.keys() + + +@pytest.mark.parametrize("aws_iam_auth", [True, "true", "True", "TRUE", "1", "yes"]) +def test_aws_iam_auth_enabled_by_any_truthy_flag(clean_redis_environment, aws_iam_auth): + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + ssl=True, + aws_iam_auth=aws_iam_auth, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache-name", + aws_iam_region="us-east-1", + ) + + assert isinstance(redis_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) + + +@pytest.mark.parametrize( + "aws_iam_serverless, expected", + [ + pytest.param(None, False, id="unset"), + pytest.param(False, False, id="bool_false"), + pytest.param("false", False, id="string_false"), + pytest.param("0", False, id="string_zero"), + pytest.param(True, True, id="bool_true"), + pytest.param("true", True, id="string_true"), + pytest.param("1", True, id="string_one"), + ], +) +def test_aws_iam_serverless_flag_reaches_the_provider(clean_redis_environment, aws_iam_serverless, expected): + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + ssl=True, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache-name", + aws_iam_region="us-east-1", + aws_iam_serverless=aws_iam_serverless, + ) + + provider = redis_kwargs["credential_provider"] + assert isinstance(provider, ElastiCacheIAMCredentialProvider) + assert provider._is_serverless is expected + assert "aws_iam_serverless" not in redis_kwargs def test_explicit_provider_wins_over_aws_iam(clean_redis_environment): diff --git a/tests/test_litellm/test_redis_credential_provider.py b/tests/test_litellm/test_redis_credential_provider.py index 9a58722b04e..e71fc5d530d 100644 --- a/tests/test_litellm/test_redis_credential_provider.py +++ b/tests/test_litellm/test_redis_credential_provider.py @@ -175,3 +175,57 @@ def test_elasticache_provider_recovers_after_a_failed_resolution(): assert user_name == "iam-user" assert token assert resolver.calls == 2 + + +@pytest.mark.parametrize( + "provider_kwargs, expected_resource_type", + [ + pytest.param({}, None, id="default_is_self_designed"), + pytest.param({"is_serverless": False}, None, id="self_designed"), + pytest.param({"is_serverless": True}, ["ServerlessCache"], id="serverless"), + ], +) +def test_elasticache_provider_signs_resource_type_only_for_serverless(provider_kwargs, expected_resource_type): + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache-name", + region="us-east-1", + credentials_resolver=_FakeResolver(_FakeCredentials("AKIA-SYNTHETIC")), + **provider_kwargs, + ) + + _, token = provider.get_credentials() + query = parse_qs(urlsplit("https://" + token).query) + + assert query.get("ResourceType") == expected_resource_type + assert query["X-Amz-Signature"] + + +def test_elasticache_provider_lowercases_the_cache_name(): + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="Mixed-Case-Cache", + region="us-east-1", + credentials_resolver=_FakeResolver(_FakeCredentials("AKIA-SYNTHETIC")), + ) + + _, token = provider.get_credentials() + + assert urlsplit("https://" + token).netloc == "mixed-case-cache" + + +def test_elasticache_provider_encodes_reserved_characters_in_the_user_name(): + user_name = "iam user/with+reserved&chars" + provider = ElastiCacheIAMCredentialProvider( + user_name=user_name, + cache_name="cache-name", + region="us-east-1", + credentials_resolver=_FakeResolver(_FakeCredentials("AKIA-SYNTHETIC")), + ) + + returned_user_name, token = provider.get_credentials() + query = parse_qs(urlsplit("https://" + token).query) + + assert returned_user_name == user_name + assert query["User"] == [user_name] + assert query["Action"] == ["connect"] diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 9ca0df396b7..0448f8373c1 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26372,6 +26372,11 @@ export interface components { * @description AWS region for ElastiCache IAM authentication */ aws_iam_region?: string | null; + /** + * Aws Iam Serverless + * @description the ElastiCache cache is serverless rather than a self-designed cluster + */ + aws_iam_serverless?: boolean | string | null; /** * Aws Iam User Name * @description AWS ElastiCache IAM user name From f0e3e031c3fd79acf04db3bca41c4f2f5d1f12ec Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 9 Sep 2026 15:01:06 -0400 Subject: [PATCH 43/60] test(redis): pin ElastiCache IAM signing and TLS coercion invariants Strengthens the serverless test to assert ResourceType is signed rather than merely present, ties _uses_tls to the redis-py kwarg coercion so the two cannot drift, locks the stripped-kwarg name tuple to the test's expectations, and adds "off" and "True" sentinel flag values. Renames the provider builder's parameter to redis_settings. --- tests/test_litellm/test_redis.py | 15 ++++++++++++++- .../test_redis_credential_provider.py | 19 ++++++++++++------- 2 files changed, 26 insertions(+), 8 deletions(-) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 5337ed825a7..0e8c86d26df 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -10,13 +10,16 @@ from redis.credentials import CredentialProvider import litellm from litellm._redis import ( + _AWS_IAM_KWARG_NAMES, _async_auth_kwargs, + _coerce_redis_kwargs_types, _get_redis_client_logic, _get_redis_cluster_kwargs, _get_redis_env_kwarg_mapping, _get_redis_kwargs, _get_redis_url_kwargs, _pretty_print_redis_config, + _uses_tls, get_redis_async_client, get_redis_client, get_redis_connection_pool, @@ -31,6 +34,7 @@ from litellm._redis_credential_provider import ( from litellm.caching.redis_cache import RedisCache from litellm.caching.redis_cluster_cache import RedisClusterCache from litellm.constants import REDIS_CLUSTER_HEALTH_CHECK_INTERVAL +from litellm.proxy._types import CoordinationRedisParams class _StubCredentialProvider(CredentialProvider): @@ -127,6 +131,8 @@ def test_aws_iam_settings_are_environment_derived(): mapping = _get_redis_env_kwarg_mapping() assert _AWS_IAM_SETTINGS <= allowed + assert set(_AWS_IAM_KWARG_NAMES) == _AWS_IAM_SETTINGS + assert {f for f in CoordinationRedisParams.model_fields if f.startswith("aws_iam_")} == _AWS_IAM_SETTINGS assert mapping["REDIS_AWS_IAM_AUTH"] == "aws_iam_auth" assert mapping["REDIS_AWS_IAM_USER_NAME"] == "aws_iam_user_name" assert mapping["REDIS_AWS_IAM_CACHE_NAME"] == "aws_iam_cache_name" @@ -348,6 +354,7 @@ def test_aws_iam_environment_settings_install_provider(clean_redis_environment, pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "false"}, id="host_ssl_false_string"), pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "0"}, id="host_ssl_zero_string"), pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "no"}, id="host_ssl_no_string"), + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "off"}, id="host_ssl_off_string"), pytest.param({"url": "redis://cache.example.com:6379", "ssl": True}, id="plaintext_url"), pytest.param( {"startup_nodes": [{"host": "cache.example.com", "port": 6379}]}, @@ -385,6 +392,7 @@ def test_aws_iam_auth_rejects_non_tls_connections(clean_redis_environment, trans [ pytest.param({"host": "cache.example.com", "port": 6379, "ssl": True}, id="host"), pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "true"}, id="host_ssl_true_string"), + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "True"}, id="host_ssl_true_capitalized"), pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "1"}, id="host_ssl_one_string"), pytest.param({"host": "cache.example.com", "port": 6379, "ssl": "yes"}, id="host_ssl_yes_string"), pytest.param( @@ -421,6 +429,11 @@ def test_aws_iam_auth_accepts_tls_connections(clean_redis_environment, transport assert isinstance(redis_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) +@pytest.mark.parametrize("ssl", ["true", "True", "TRUE", "1", "yes", "YES", "false", "0", "no", "off", "", "maybe"]) +def test_tls_detection_agrees_with_the_ssl_kwarg_coercion(ssl): + assert _uses_tls({"ssl": ssl}) is _coerce_redis_kwargs_types({"ssl": ssl})["ssl"] + + def test_aws_iam_settings_are_removed_for_url_and_static_credentials(clean_redis_environment): redis_kwargs = _get_redis_client_logic( url="rediss://url-user:url-pass@cache.example.com:6380", @@ -522,7 +535,7 @@ def test_aws_iam_settings_map_to_distinct_provider_fields(clean_redis_environmen assert provider._region == "iam-region-value" -@pytest.mark.parametrize("aws_iam_auth", [None, False, "", "false", "0", "no"]) +@pytest.mark.parametrize("aws_iam_auth", [None, False, "", "false", "0", "no", "off"]) def test_aws_iam_auth_disabled_does_not_install_provider(clean_redis_environment, aws_iam_auth): redis_kwargs = _get_redis_client_logic( host="cache.example.com", diff --git a/tests/test_litellm/test_redis_credential_provider.py b/tests/test_litellm/test_redis_credential_provider.py index e71fc5d530d..96b1933b0e3 100644 --- a/tests/test_litellm/test_redis_credential_provider.py +++ b/tests/test_litellm/test_redis_credential_provider.py @@ -178,14 +178,14 @@ def test_elasticache_provider_recovers_after_a_failed_resolution(): @pytest.mark.parametrize( - "provider_kwargs, expected_resource_type", + "provider_kwargs, expected_operation_params", [ - pytest.param({}, None, id="default_is_self_designed"), - pytest.param({"is_serverless": False}, None, id="self_designed"), - pytest.param({"is_serverless": True}, ["ServerlessCache"], id="serverless"), + pytest.param({}, frozenset({"Action", "User"}), id="default_is_self_designed"), + pytest.param({"is_serverless": False}, frozenset({"Action", "User"}), id="self_designed"), + pytest.param({"is_serverless": True}, frozenset({"Action", "User", "ResourceType"}), id="serverless"), ], ) -def test_elasticache_provider_signs_resource_type_only_for_serverless(provider_kwargs, expected_resource_type): +def test_elasticache_provider_signs_resource_type_only_for_serverless(provider_kwargs, expected_operation_params): provider = ElastiCacheIAMCredentialProvider( user_name="iam-user", cache_name="cache-name", @@ -195,9 +195,14 @@ def test_elasticache_provider_signs_resource_type_only_for_serverless(provider_k ) _, token = provider.get_credentials() - query = parse_qs(urlsplit("https://" + token).query) + query_string = urlsplit("https://" + token).query + param_names = tuple(pair.split("=", 1)[0] for pair in query_string.split("&")) + first_auth_param = next(i for i, name in enumerate(param_names) if name.startswith("X-Amz-")) + query = parse_qs(query_string) - assert query.get("ResourceType") == expected_resource_type + assert frozenset(param_names[:first_auth_param]) == expected_operation_params + assert all(name.startswith("X-Amz-") for name in param_names[first_auth_param:]) + assert query.get("ResourceType") == (["ServerlessCache"] if "ResourceType" in expected_operation_params else None) assert query["X-Amz-Signature"] From 13ddd1ec6481756ac31cac4b7253e23eef3bf656 Mon Sep 17 00:00:00 2001 From: Wolfram Ravenwolf Date: Thu, 10 Sep 2026 20:47:13 +0200 Subject: [PATCH 44/60] fix(wandb): gate reasoning effort on model capabilities --- litellm/llms/wandb/chat/transformation.py | 6 +- ...odel_prices_and_context_window_backup.json | 40 +++++ model_prices_and_context_window.json | 40 +++++ .../wandb/test_wandb_chat_transformation.py | 170 ++++++++++++++---- 4 files changed, 219 insertions(+), 37 deletions(-) diff --git a/litellm/llms/wandb/chat/transformation.py b/litellm/llms/wandb/chat/transformation.py index 9c3ff818331..fdd6644f03d 100644 --- a/litellm/llms/wandb/chat/transformation.py +++ b/litellm/llms/wandb/chat/transformation.py @@ -6,12 +6,16 @@ This is OpenAI compatible - no translation needed / occurs from typing import Final +import litellm from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig class WandbConfig(OpenAIGPTConfig): def get_supported_openai_params(self, model: str) -> list[str]: # mutable-ok: inherited contract - return super().get_supported_openai_params(model) + ["reasoning_effort"] # mutable-ok: inherited contract + supported_params: Final = super().get_supported_openai_params(model) + if litellm.supports_reasoning(model=model, custom_llm_provider="wandb"): + return supported_params + ["reasoning_effort"] # mutable-ok: inherited contract + return supported_params def map_openai_params( self, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 55618d9f772..da47ef85eb1 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -45960,6 +45960,7 @@ "output_cost_per_token": 0.0 }, "wandb/openai/gpt-oss-120b": { + "supports_reasoning": true, "max_tokens": 131072, "max_input_tokens": 131072, "max_output_tokens": 131072, @@ -45970,6 +45971,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/openai/gpt-oss-20b": { + "supports_reasoning": true, "max_tokens": 131072, "max_input_tokens": 131072, "max_output_tokens": 131072, @@ -45980,6 +45982,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/zai-org/GLM-4.5": { + "supports_reasoning": true, "max_tokens": 131072, "max_input_tokens": 131072, "max_output_tokens": 131072, @@ -46008,6 +46011,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/Qwen/Qwen3-235B-A22B-Thinking-2507": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "max_output_tokens": 262144, @@ -46064,6 +46068,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/deepseek-ai/DeepSeek-V3.1": { + "supports_reasoning": true, "max_tokens": 128000, "max_input_tokens": 161000, "max_output_tokens": 128000, @@ -46074,6 +46079,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/deepseek-ai/DeepSeek-R1-0528": { + "supports_reasoning": true, "max_tokens": 161000, "max_input_tokens": 161000, "max_output_tokens": 161000, @@ -55631,6 +55637,7 @@ "supports_vision": false }, "wandb/deepseek-ai/DeepSeek-V4-Flash": { + "supports_reasoning": true, "max_tokens": 1048576, "max_input_tokens": 1048576, "input_cost_per_token": 1.4e-07, @@ -55643,6 +55650,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/deepseek-ai/DeepSeek-V4-Flash-0731": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 1.3e-07, @@ -55655,6 +55663,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/deepseek-ai/DeepSeek-V4-Pro": { + "supports_reasoning": true, "max_tokens": 1048576, "max_input_tokens": 1048576, "input_cost_per_token": 1.15e-06, @@ -55667,6 +55676,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/google/gemma-4-31B-it": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 1e-07, @@ -55707,6 +55717,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/MiniMaxAI/MiniMax-M3": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 2.3e-07, @@ -55719,6 +55730,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/moonshotai/Kimi-K2.7-Code": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 7.1e-07, @@ -55731,6 +55743,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/moonshotai/Kimi-K2.6": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 6.5e-07, @@ -55743,6 +55756,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/nvidia/NVIDIA-Nemotron-3.5-Lightning-30B-A3B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 1e-07, @@ -55755,6 +55769,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/nvidia/NVIDIA-Nemotron-3-Ultra-550B-A55B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 7.5e-07, @@ -55777,6 +55792,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/Qwen/Qwen3.8-27B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 4e-07, @@ -55789,6 +55805,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/Qwen/Qwen3.6-35B-A3B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 2.5e-07, @@ -55799,6 +55816,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/Qwen/Qwen3.6-27B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 6e-07, @@ -55811,6 +55829,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/Qwen/Qwen3.5-35B-A3B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 2.5e-07, @@ -55830,7 +55849,28 @@ "supports_vision": false, "source": "https://wandb.ai/site/pricing/tokens/" }, + "wandb/deepseek-ai/DeepSeek-V4-Pro-0813": { + "litellm_provider": "wandb", + "mode": "chat", + "supports_reasoning": true, + "input_cost_per_token": 0.00000131, + "output_cost_per_token": 0.00000396, + "cache_read_input_token_cost": 0.000000044, + "supports_prompt_caching": true, + "source": "https://wandb.ai/site/pricing/tokens/" + }, + "wandb/ibm-granite/granite-4.2-8b": { + "litellm_provider": "wandb", + "mode": "chat", + "supports_reasoning": true, + "input_cost_per_token": 0.0000001, + "output_cost_per_token": 0.00000015, + "cache_read_input_token_cost": 0.00000005, + "supports_prompt_caching": true, + "source": "https://wandb.ai/site/pricing/tokens/" + }, "wandb/zai-org/GLM-5.2": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 7.6e-07, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 55618d9f772..da47ef85eb1 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -45960,6 +45960,7 @@ "output_cost_per_token": 0.0 }, "wandb/openai/gpt-oss-120b": { + "supports_reasoning": true, "max_tokens": 131072, "max_input_tokens": 131072, "max_output_tokens": 131072, @@ -45970,6 +45971,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/openai/gpt-oss-20b": { + "supports_reasoning": true, "max_tokens": 131072, "max_input_tokens": 131072, "max_output_tokens": 131072, @@ -45980,6 +45982,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/zai-org/GLM-4.5": { + "supports_reasoning": true, "max_tokens": 131072, "max_input_tokens": 131072, "max_output_tokens": 131072, @@ -46008,6 +46011,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/Qwen/Qwen3-235B-A22B-Thinking-2507": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "max_output_tokens": 262144, @@ -46064,6 +46068,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/deepseek-ai/DeepSeek-V3.1": { + "supports_reasoning": true, "max_tokens": 128000, "max_input_tokens": 161000, "max_output_tokens": 128000, @@ -46074,6 +46079,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/deepseek-ai/DeepSeek-R1-0528": { + "supports_reasoning": true, "max_tokens": 161000, "max_input_tokens": 161000, "max_output_tokens": 161000, @@ -55631,6 +55637,7 @@ "supports_vision": false }, "wandb/deepseek-ai/DeepSeek-V4-Flash": { + "supports_reasoning": true, "max_tokens": 1048576, "max_input_tokens": 1048576, "input_cost_per_token": 1.4e-07, @@ -55643,6 +55650,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/deepseek-ai/DeepSeek-V4-Flash-0731": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 1.3e-07, @@ -55655,6 +55663,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/deepseek-ai/DeepSeek-V4-Pro": { + "supports_reasoning": true, "max_tokens": 1048576, "max_input_tokens": 1048576, "input_cost_per_token": 1.15e-06, @@ -55667,6 +55676,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/google/gemma-4-31B-it": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 1e-07, @@ -55707,6 +55717,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/MiniMaxAI/MiniMax-M3": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 2.3e-07, @@ -55719,6 +55730,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/moonshotai/Kimi-K2.7-Code": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 7.1e-07, @@ -55731,6 +55743,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/moonshotai/Kimi-K2.6": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 6.5e-07, @@ -55743,6 +55756,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/nvidia/NVIDIA-Nemotron-3.5-Lightning-30B-A3B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 1e-07, @@ -55755,6 +55769,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/nvidia/NVIDIA-Nemotron-3-Ultra-550B-A55B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 7.5e-07, @@ -55777,6 +55792,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/Qwen/Qwen3.8-27B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 4e-07, @@ -55789,6 +55805,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/Qwen/Qwen3.6-35B-A3B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 2.5e-07, @@ -55799,6 +55816,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/Qwen/Qwen3.6-27B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 6e-07, @@ -55811,6 +55829,7 @@ "source": "https://wandb.ai/site/pricing/tokens/" }, "wandb/Qwen/Qwen3.5-35B-A3B": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 2.5e-07, @@ -55830,7 +55849,28 @@ "supports_vision": false, "source": "https://wandb.ai/site/pricing/tokens/" }, + "wandb/deepseek-ai/DeepSeek-V4-Pro-0813": { + "litellm_provider": "wandb", + "mode": "chat", + "supports_reasoning": true, + "input_cost_per_token": 0.00000131, + "output_cost_per_token": 0.00000396, + "cache_read_input_token_cost": 0.000000044, + "supports_prompt_caching": true, + "source": "https://wandb.ai/site/pricing/tokens/" + }, + "wandb/ibm-granite/granite-4.2-8b": { + "litellm_provider": "wandb", + "mode": "chat", + "supports_reasoning": true, + "input_cost_per_token": 0.0000001, + "output_cost_per_token": 0.00000015, + "cache_read_input_token_cost": 0.00000005, + "supports_prompt_caching": true, + "source": "https://wandb.ai/site/pricing/tokens/" + }, "wandb/zai-org/GLM-5.2": { + "supports_reasoning": true, "max_tokens": 262144, "max_input_tokens": 262144, "input_cost_per_token": 7.6e-07, diff --git a/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py b/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py index 4fd66448009..481a52bebad 100644 --- a/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py +++ b/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py @@ -6,30 +6,90 @@ Nebius AI Studio is an OpenAI-compatible provider with minor customizations. """ import json +from typing import Final import pytest +import respx import litellm from litellm import completion from litellm.llms.wandb.chat.transformation import WandbConfig +WANDB_REASONING_MODELS: Final = ( + "deepseek-ai/DeepSeek-V4-Flash", + "deepseek-ai/DeepSeek-V4-Flash-0731", + "deepseek-ai/DeepSeek-V4-Pro", + "deepseek-ai/DeepSeek-V4-Pro-0813", + "deepseek-ai/DeepSeek-V3.1", + "google/gemma-4-31B-it", + "ibm-granite/granite-4.2-8b", + "MiniMaxAI/MiniMax-M3", + "moonshotai/Kimi-K2.7-Code", + "moonshotai/Kimi-K2.6", + "nvidia/NVIDIA-Nemotron-3.5-Lightning-30B-A3B", + "nvidia/NVIDIA-Nemotron-3-Ultra-550B-A55B", + "openai/gpt-oss-120b", + "openai/gpt-oss-20b", + "Qwen/Qwen3.8-27B", + "Qwen/Qwen3.6-35B-A3B", + "Qwen/Qwen3.6-27B", + "Qwen/Qwen3.5-35B-A3B", + "zai-org/GLM-5.2", + "moonshotai/Kimi-K2.5", + "MiniMaxAI/MiniMax-M2.5", + "zai-org/GLM-4.5", + "Qwen/Qwen3-235B-A22B-Thinking-2507", + "deepseek-ai/DeepSeek-R1-0528", +) + + +@pytest.fixture +def wandb_test_config(local_model_cost_map, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "telemetry", False) + monkeypatch.setattr(litellm, "drop_params", False) + + +@pytest.fixture +def wandb_request_mock(respx_mock: respx.MockRouter) -> respx.Route: + return respx_mock.post("https://api.inference.wandb.ai/v1/chat/completions").respond( + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "test-model", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Done"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + status_code=200, + ) + + class TestWandbConfig: """Test class for WandB Inference functionality""" - def test_map_openai_params_preserves_reasoning_effort(self): - supported_params = litellm.get_supported_openai_params(model="wandb/openai/gpt-oss-20b") + @pytest.mark.parametrize("model", WANDB_REASONING_MODELS) + def test_map_openai_params_preserves_reasoning_effort(self, wandb_test_config, model: str): + assert litellm.model_cost[f"wandb/{model}"].get("supports_reasoning") is True + supported_params = litellm.get_supported_openai_params(model=f"wandb/{model}") assert supported_params is not None assert "reasoning_effort" in supported_params result = WandbConfig().map_openai_params( - non_default_params={"reasoning_effort": "medium"}, + non_default_params={"reasoning_effort": "medium", "max_completion_tokens": 64}, optional_params={}, - model="openai/gpt-oss-20b", + model=model, drop_params=True, ) - assert result == {"reasoning_effort": "medium"} + assert result == {"reasoning_effort": "medium", "max_tokens": 64} def test_default_api_base(self): """Test that default API base is used when none is provided""" @@ -155,40 +215,78 @@ class TestWandbConfig: assert "Hey from LiteLLM" in content @pytest.mark.respx() - def test_wandb_completion_preserves_reasoning_effort_with_drop_params(self, respx_mock, monkeypatch): - monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - - api_base = "https://api.inference.wandb.ai/v1" - request_mock = respx_mock.post(f"{api_base}/chat/completions").respond( - json={ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "gpt-oss-20b", - "choices": [ - { - "index": 0, - "message": {"role": "assistant", "content": "Done"}, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 1, - "total_tokens": 2, - }, - }, - status_code=200, - ) - + @pytest.mark.parametrize( + "model,effort", + tuple((model, "medium") for model in WANDB_REASONING_MODELS) + + ( + ("Qwen/Qwen3.8-27B", "low"), + ("Qwen/Qwen3.8-27B", "xhigh"), + ), + ) + def test_wandb_completion_preserves_reasoning_effort_with_drop_params( + self, wandb_test_config, wandb_request_mock: respx.Route, model: str, effort: str + ): completion( - model="wandb/openai/gpt-oss-20b", + model=f"wandb/{model}", messages=[{"role": "user", "content": "Hello"}], api_key="fake-wandb-key", - api_base=api_base, - reasoning_effort="medium", + api_base="https://api.inference.wandb.ai/v1", + reasoning_effort=effort, + max_completion_tokens=64, drop_params=True, ) - request_body = json.loads(request_mock.calls[0].request.content) - assert request_body["reasoning_effort"] == "medium" + assert wandb_request_mock.call_count == 1 + request_body = json.loads(wandb_request_mock.calls[0].request.content) + assert request_body["model"] == model + assert request_body["reasoning_effort"] == effort + assert request_body["max_tokens"] == 64 + assert "max_completion_tokens" not in request_body + + @pytest.mark.respx(assert_all_called=False) + @pytest.mark.parametrize("drop_params", [True, False]) + @pytest.mark.parametrize( + "model,explicit_false", + [ + ("meta-llama/Llama-3.1-8B-Instruct", False), + ("unknown-model", False), + ("openai/gpt-oss-20b", True), + ], + ) + def test_wandb_completion_without_reasoning_support( + self, + wandb_test_config, + wandb_request_mock: respx.Route, + respx_mock: respx.MockRouter, + monkeypatch: pytest.MonkeyPatch, + model: str, + explicit_false: bool, + drop_params: bool, + ): + with monkeypatch.context() as context: + if explicit_false: + context.setitem(litellm.model_cost[f"wandb/{model}"], "supports_reasoning", False) + + kwargs = { + "model": f"wandb/{model}", + "messages": [{"role": "user", "content": "Hello"}], + "api_key": "fake-wandb-key", + "api_base": "https://api.inference.wandb.ai/v1", + "reasoning_effort": "medium", + "drop_params": drop_params, + } + if not drop_params: + with pytest.raises(litellm.UnsupportedParamsError, match="reasoning_effort"): + completion(**kwargs) + assert len(respx_mock.calls) == 0 + return + + completion(**kwargs) + assert wandb_request_mock.call_count == 1 + request_body = json.loads(wandb_request_mock.calls[0].request.content) + assert request_body["model"] == model + assert "reasoning_effort" not in request_body + + supported_params = litellm.get_supported_openai_params(model=f"wandb/{model}") + assert supported_params is not None + assert "reasoning_effort" not in supported_params From f4ebcef0a17c2952ab90100fdba702d863c728f9 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Thu, 10 Sep 2026 12:32:09 -0700 Subject: [PATCH 45/60] fix(router): classify encrypted delegated tasks with native Responses --- .../transformation.py | 3 + .../complexity_router/README.md | 10 + .../complexity_router/complexity_router.py | 119 ++++++++++-- .../router_strategy/test_complexity_router.py | 178 ++++++++++++++++++ 4 files changed, 293 insertions(+), 17 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index c3c18e3d009..5a6debc4af5 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -1201,6 +1201,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # Cast to Any to match the expected union type for tools list items tools.append(cast(Any, web_search_tool)) + def transform_response_format_to_text_format(self, response_format: object) -> "ResponseText | None": + return self._transform_response_format_to_text_format(response_format) + def _transform_response_format_to_text_format(self, response_format: object) -> "ResponseText | None": """ Transform Chat Completion response_format parameter to Responses API text.format parameter. diff --git a/litellm/router_strategy/complexity_router/README.md b/litellm/router_strategy/complexity_router/README.md index d605b43e42a..79119008aa2 100644 --- a/litellm/router_strategy/complexity_router/README.md +++ b/litellm/router_strategy/complexity_router/README.md @@ -361,6 +361,16 @@ model_list: keep the classifier deployment or provider default, or set a supported value such as `none` or `low` to override that call. +When the current ask is a Responses API `agent_message` containing `encrypted_content`, LLM +classification preserves the encrypted task and uses native Responses. This also bypasses the +local scoring shortcut in `heuristic_first` and `hybrid` modes. The configured classifier must use +a native OpenAI or Azure OpenAI Responses deployment with access to the encrypted content. The +provider handles the encrypted task, and the classifier still chooses the tier dynamically + +Unsupported classifier deployments and provider decryption errors use the existing +`classifier_fallback` policy. No fixed tier is introduced for encrypted tasks. Plaintext asks and +requests carrying only historical encrypted reasoning retain the existing classifier path + Classifier calls have a one-attempt hard deadline. After a timeout, the router opens a process-local circuit for that classifier and sends every session through `classifier_fallback` for `classifier_llm_config.circuit_breaker_cooldown_seconds` (30 seconds by default). When the cooldown diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index faafcea404a..3babe5b96ed 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -25,7 +25,7 @@ from threading import Lock from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast -from pydantic import BaseModel, create_model +from pydantic import BaseModel, TypeAdapter, create_model from litellm._logging import verbose_router_logger from litellm.constants import ( @@ -56,6 +56,7 @@ from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionImageObject, ChatCompletionTextObject, + ResponsesAPIResponse, ) from litellm.types.utils import ( AUTOROUTER_CLASSIFIER_CALL_ORIGIN, @@ -364,7 +365,7 @@ def _parent_session_kwargs(request_kwargs: Mapping[str, Any] | None) -> Mapping[ return {k: kwargs[k] for k in ("litellm_session_id", "litellm_trace_id") if kwargs.get(k) is not None} -def _response_cost_or_none(response: ModelResponse) -> float | None: +def _response_cost_or_none(response: ModelResponse | ResponsesAPIResponse) -> float | None: hidden_params: Final = response._hidden_params if not isinstance(hidden_params, dict): return None @@ -486,6 +487,36 @@ def _human_text(content: object, marker_pairs: tuple[tuple[str, str], ...] = _DE return _strip_reminder_blocks(_message_text(content), marker_pairs) +def _encrypted_classifier_task( + request_kwargs: Mapping[str, object] | None, + marker_pairs: tuple[tuple[str, str], ...], +) -> dict[str, object] | None: + from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages + + raw_input: Final = (request_kwargs or EMPTY_MAPPING).get("input") + if not isinstance(raw_input, list) or (request_kwargs or EMPTY_MAPPING).get("messages"): + return None + items: Final = TypeAdapter(tuple[dict[str, object], ...]).validate_python(raw_input) + current: Final = next( + ( + item + for item in reversed(items) + if (messages := resolve_structured_messages(messages=None, request_kwargs={"input": [item]})) + and any(_iter_human_asks_newest_first(messages, marker_pairs)) + ), + None, + ) + if current is None or current.get("type") != "agent_message" or not isinstance(current.get("content"), list): + return None + parts: Final = TypeAdapter(tuple[dict[str, object], ...]).validate_python(current["content"]) + if not any(part.get("type") == "encrypted_content" and part.get("encrypted_content") for part in parts): + return None + return { + **current, + "content": [part for part in parts if part.get("type") in ("input_text", "encrypted_content")], + } + + def _iter_human_asks_newest_first( messages: Sequence[Mapping[str, object]], marker_pairs: tuple[tuple[str, str], ...] = _DEFAULT_REMINDER_MARKERS, @@ -1630,6 +1661,10 @@ 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 in ("heuristic_first", "hybrid") and _encrypted_classifier_task( + request_kwargs, self._reminder_markers + ): + return await self._llm_classifier_outcome(prompt, system_prompt, request_kwargs, messages) if self.config.classifier_type == "heuristic_first" and self.config.classifier_llm_config is not None: return await self._classify_heuristic_first(prompt, system_prompt, request_kwargs, messages) if self.config.classifier_type == "hybrid" and self.config.classifier_llm_config is not None: @@ -1970,8 +2005,9 @@ class ComplexityRouter(CustomLogger): > 1 ) + encrypted_task: Final = _encrypted_classifier_task(request_kwargs, self._reminder_markers) user_payload: Final = self._build_classifier_user_payload( - prompt=prompt, + prompt="The delegated task in the following agent_message." if encrypted_task is not None else prompt, system_prompt=system_prompt, prior_turns=prior_turns, messages=messages, @@ -2004,34 +2040,37 @@ class ComplexityRouter(CustomLogger): if llm_config.reasoning_effort is not None: classifier_call_params = MappingProxyType({"reasoning_effort": llm_config.reasoning_effort}) - proxy_server_request: Final = { - "body": { - "model": llm_config.model, - "messages": messages_for_call, - "response_format": response_format, - **classifier_call_params, - } - } + payload: Final = ( + self._native_classifier_payload(llm_config.model, messages_for_call, response_format, encrypted_task) + if encrypted_task is not None + else {"messages": messages_for_call, "response_format": response_format, **classifier_call_params} + ) + proxy_server_request: Final = {"body": {"model": llm_config.model, **payload}} + classify: Final = ( + self.litellm_router_instance.aresponses + if encrypted_task is not None + else self.litellm_router_instance.acompletion + ) classifier_timeout_s: Final[float] = llm_config.timeout_ms / 1000 - response: Final[ModelResponse] = await asyncio.wait_for( - self.litellm_router_instance.acompletion( + response: Final[ModelResponse | ResponsesAPIResponse] = await asyncio.wait_for( + classify( model=llm_config.model, - messages=messages_for_call, stream=False, - response_format=response_format, timeout=classifier_timeout_s, num_retries=0, disable_fallbacks=True, metadata=metadata, proxy_server_request=proxy_server_request, turn_off_message_logging=turn_off_message_logging, - **classifier_call_params, + **payload, **_parent_session_kwargs(request_kwargs), ), timeout=classifier_timeout_s, ) - content: Final = response.choices[0].message.content + content: Final = ( + response.output_text if isinstance(response, ResponsesAPIResponse) else response.choices[0].message.content + ) if not content: raise ValueError("LLM classifier returned empty content") raw_tier: Final = _LabeledTierClassification.model_validate_json(content).tier @@ -2040,6 +2079,52 @@ class ComplexityRouter(CustomLogger): raise ValueError(f"LLM classifier returned an unrecognized tier: {raw_tier!r}") return tier, _response_cost_or_none(response) + def _native_classifier_payload( + self, + model: str, + messages: list[AllMessageValues], # mutable-ok: existing transformation accepts the SDK message list + response_format: Mapping[str, object], + encrypted_task: Mapping[str, object], + ) -> Mapping[str, object]: + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider, get_llm_provider + from litellm.types.router import LiteLLM_Params + + deployments: Final = self._group_deployments(model) + if not deployments: + raise ValueError("Encrypted task classification requires a native OpenAI Responses classifier deployment") + for params in (LiteLLM_Params.model_validate(deployment.get("litellm_params")) for deployment in deployments): + if declared_authenticating_provider(params.model, params.custom_llm_provider): + raise ValueError( + "Encrypted task classification requires a native OpenAI Responses classifier deployment" + ) + _, provider, _, _ = get_llm_provider(model=params.model, litellm_params=params) + if ( + provider not in ("openai", "azure") + or params.use_chat_completions_api + or params.model.startswith("openai/chat_completions/") + ): + raise ValueError( + "Encrypted task classification requires a native OpenAI Responses classifier deployment" + ) + transformation: Final = LiteLLMResponsesTransformationHandler() + input_items, instructions = transformation.convert_chat_completion_messages_to_responses_api(messages) + llm_config: Final = self.config.classifier_llm_config + reasoning: Final = ( + {"reasoning": {"effort": llm_config.reasoning_effort}} + if llm_config is not None and llm_config.reasoning_effort is not None + else {} + ) + return { + "input": [*input_items, encrypted_task], + "instructions": instructions, + "text": transformation.transform_response_format_to_text_format(dict(response_format)), + "store": False, + **reasoning, + } + @staticmethod def _build_classifier_user_payload( prompt: str, diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 51103297c58..87438ef5c69 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -5,6 +5,8 @@ Tests the rule-based complexity scoring and tier assignment logic. """ import asyncio +import copy +import json import logging import sys import time @@ -66,6 +68,7 @@ from litellm.types.router import ( LiteLLM_Params, TaggedPreRoutingStrategy, ) +from litellm.types.llms.openai import ResponsesAPIResponse requires_semantic_router = pytest.mark.skipif( @@ -2482,6 +2485,181 @@ class TestTierLabels: assert set(config.tier_boundaries) == {"simple_medium", "medium_complex", "complex_reasoning"} +def _encrypted_agent_task() -> dict[str, object]: + return { + "type": "agent_message", + "author": "/root", + "recipient": "/root/child", + "content": [ + {"type": "input_text", "text": "Message Type: NEW_TASK\nTask name: /root/child\nPayload:\nHello"}, + {"type": "encrypted_content", "encrypted_content": "opaque-provider-task"}, + ], + } + + +def _native_classifier_response(content: str) -> ResponsesAPIResponse: + response: Final = ResponsesAPIResponse( + id="resp_classifier", + created_at=0, + status="completed", + output=[{"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": content}]}], + ) + response._hidden_params = {"response_cost": 0.0001} + return response + + +def _native_classifier_router( + output: str = '{"tier":"REASONING"}', + classifier_type: str = "llm", + deployment_model: str = "openai/gpt-6-astra", + failure: Exception | None = None, +) -> tuple[ComplexityRouter, MagicMock]: + dependency: Final = MagicMock( + aresponses=AsyncMock(return_value=_native_classifier_response(output), side_effect=failure), + acompletion=AsyncMock(return_value=_llm_response('{"tier":"SIMPLE"}')), + get_model_list=MagicMock(return_value=[{"litellm_params": {"model": deployment_model}}]), + ) + return ( + ComplexityRouter( + model_name="encrypted-router", + litellm_router_instance=dependency, + complexity_router_config={ + "tiers": {"SIMPLE": "cheap-model", "REASONING": "deep-model"}, + "classifier_type": classifier_type, + "classifier_llm_config": {"model": "classifier", "timeout_ms": 100, "reasoning_effort": "low"}, + "heuristic_first_max_tier": "SIMPLE" if classifier_type == "heuristic_first" else None, + "hybrid_boundary_margin": 0.01 if classifier_type == "hybrid" else None, + "classifier_fallback": "default_model", + "default_model": "deep-model", + "session_affinity": False, + "deployment_affinity": False, + }, + ), + dependency, + ) + + +class TestEncryptedTaskClassifier: + @pytest.mark.asyncio + @pytest.mark.parametrize("classifier_type", ["llm", "heuristic_first", "hybrid"]) + @pytest.mark.parametrize("tier,model", [("SIMPLE", "cheap-model"), ("REASONING", "deep-model")]) + async def test_encrypted_task_routes_by_native_verdict(self, classifier_type: str, tier: str, model: str): + router, dependency = _native_classifier_router(json.dumps({"tier": tier}), classifier_type) + task: Final = _encrypted_agent_task() + request: Final = { + "input": [ + {"role": "user", "content": "Prior task context"}, + task, + {"type": "function_call_output", "call_id": "call_1", "output": "Tool output"}, + {"role": "user", "content": "Injected reminder"}, + ], + "instructions": "Caller constraints", + "tools": [{"type": "function", "name": "execute"}], + "previous_response_id": "resp_parent", + "litellm_session_id": "parent-session", + "litellm_trace_id": "parent-trace", + "turn_off_message_logging": True, + "litellm_metadata": {"user_api_key_hash": "caller-key-hash"}, + } + original: Final = copy.deepcopy(request) + + result: Final = await router.async_pre_routing_hook(model="encrypted-router", request_kwargs=request) + + assert result.model == model + assert result.routing_decision["tier"] == tier + assert result.routing_decision["cause"] == "llm_classifier" + assert result.routing_decision["classifier_cost"] == 0.0001 + assert result.messages is None + assert request == original + dependency.acompletion.assert_not_called() + call: Final = dependency.aresponses.call_args.kwargs + assert call["input"][-1] == task + assert "opaque-provider-task" not in json.dumps(call["input"][:-1]) + assert "Prior task context" in json.dumps(call["input"][:-1]) + assert "Caller constraints" in json.dumps(call["input"][:-1]) + assert "Caller constraints" not in call["instructions"] + assert "SIMPLE" in call["instructions"] and "REASONING" in call["instructions"] + assert call["text"]["format"]["schema"]["properties"]["tier"]["enum"] == [ + "SIMPLE", "MEDIUM", "COMPLEX", "REASONING" + ] + assert call["text"]["format"]["strict"] is True + assert call["reasoning"] == {"effort": "low"} + assert call["store"] is False + assert call["stream"] is False + assert "tools" not in call and "previous_response_id" not in call + assert "messages" not in call and "response_format" not in call + assert call["timeout"] == 0.1 and call["num_retries"] == 0 and call["disable_fallbacks"] is True + assert call["litellm_session_id"] == "parent-session" + assert call["litellm_trace_id"] == "parent-trace" + assert call["turn_off_message_logging"] is True + assert call["metadata"]["user_api_key_hash"] == "caller-key-hash" + assert call["proxy_server_request"]["body"]["input"] == call["input"] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "items", + [ + [{"type": "reasoning", "encrypted_content": "opaque-history", "summary": []}, {"role": "user", "content": "hi"}], + [_encrypted_agent_task(), {"role": "user", "content": "hi"}], + [{**_encrypted_agent_task(), "content": [{"type": "input_text", "text": "hi"}]}], + [{"role": "user", "content": "gAAAA is plain text"}], + [{"role": "user", "content": "hi"}, {"type": "function_call_output", "call_id": "call_1", "output": "opaque-provider-task"}], + ], + ids=["historical-reasoning", "older-encrypted-task", "plaintext-agent", "ciphertext-looking-text", "tool-output"], + ) + async def test_other_asks_keep_chat_classifier(self, items: list[dict[str, object]]): + router, dependency = _native_classifier_router() + + result: Final = await router.async_pre_routing_hook(model="encrypted-router", request_kwargs={"input": items}) + + assert result.model == "cheap-model" + assert result.routing_decision["cause"] == "llm_classifier" + dependency.aresponses.assert_not_called() + dependency.acompletion.assert_awaited_once() + + @pytest.mark.asyncio + @pytest.mark.parametrize("output", ["", "not-json", '{"tier":"UNKNOWN"}']) + async def test_invalid_native_verdict_uses_existing_fallback(self, output: str): + router, dependency = _native_classifier_router(output=output) + + result: Final = await router.async_pre_routing_hook( + model="encrypted-router", request_kwargs={"input": [_encrypted_agent_task()]} + ) + + assert result.model == "deep-model" + assert result.routing_decision["cause"] == "default_model_fallback" + dependency.aresponses.assert_awaited_once() + dependency.acompletion.assert_not_called() + + @pytest.mark.asyncio + @pytest.mark.parametrize("deployment_model", ["anthropic/test-classifier", "openai/chat_completions/gpt-6-astra"]) + async def test_incompatible_classifier_does_not_flatten_encryption(self, deployment_model: str): + router, dependency = _native_classifier_router(deployment_model=deployment_model) + + result: Final = await router.async_pre_routing_hook( + model="encrypted-router", request_kwargs={"input": [_encrypted_agent_task()]} + ) + + assert result.model == "deep-model" + assert result.routing_decision["cause"] == "default_model_fallback" + dependency.aresponses.assert_not_called() + dependency.acompletion.assert_not_called() + + @pytest.mark.asyncio + @pytest.mark.parametrize("failure", [ValueError("invalid_encrypted_content"), TimeoutError("classifier timed out")]) + async def test_native_provider_failure_uses_existing_fallback(self, failure: Exception): + router, dependency = _native_classifier_router(failure=failure) + + result: Final = await router.async_pre_routing_hook( + model="encrypted-router", request_kwargs={"input": [_encrypted_agent_task()]} + ) + + assert result.model == "deep-model" + assert result.routing_decision["cause"] == "default_model_fallback" + dependency.aresponses.assert_awaited_once() + dependency.acompletion.assert_not_called() + + class TestLLMClassifier: """Test the LLM-based classifier path (aclassify) and its fallback behavior.""" From 208d554c008da91b74062413ad8d5d454f7746b5 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Thu, 10 Sep 2026 13:04:36 -0700 Subject: [PATCH 46/60] fix(router): validate encrypted classifiers after deployment selection --- .../llms/base_llm/responses/transformation.py | 3 + .../llms/openai/responses/transformation.py | 3 + litellm/responses/main.py | 8 + .../complexity_router/README.md | 3 + .../complexity_router/complexity_router.py | 35 ++-- .../router_strategy/test_complexity_router.py | 151 ++++++++++++++++-- 6 files changed, 168 insertions(+), 35 deletions(-) diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py index 1365941fe2a..14f00aaaa21 100644 --- a/litellm/llms/base_llm/responses/transformation.py +++ b/litellm/llms/base_llm/responses/transformation.py @@ -62,6 +62,9 @@ class BaseResponsesAPIConfig(ABC): """ return False + def supports_encrypted_agent_messages(self) -> bool: + return False + def sign_request( self, headers: dict, diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 666e9b32011..833ae206024 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -110,6 +110,9 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): def supports_native_file_search(self) -> bool: return True + def supports_encrypted_agent_messages(self) -> bool: + return self.custom_llm_provider in (LlmProviders.OPENAI, LlmProviders.AZURE) + @staticmethod def _is_gpt_5_model(model: str) -> bool: """Return True only for actual OpenAI GPT-5 models. diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 7a210cdd970..e88cd618a8b 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -1078,6 +1078,7 @@ def responses( litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) _is_async: Final = kwargs.pop("aresponses", False) is True skip_mcp_handler: Final = kwargs.pop("_skip_mcp_handler", False) + require_encrypted_task_support: Final = kwargs.pop("_require_encrypted_task_support", False) is True use_chat_completions_api = _pop_use_chat_completions_api_kw(kwargs) client_headers: Final = kwargs.get("headers") @@ -1186,6 +1187,13 @@ def responses( model, custom_llm_provider, deployment_model_info ) + if require_encrypted_task_support and ( + _bridges_to_chat_completions(responses_api_provider_config, use_chat_completions_api) + or responses_api_provider_config is None + or not responses_api_provider_config.supports_encrypted_agent_messages() + ): + raise ValueError("Encrypted task classification requires a compatible native Responses deployment") + local_vars.update(kwargs) # Map reasoning_effort (from litellm_params/proxy config) to reasoning when not set if reasoning is None and "reasoning_effort" in local_vars: diff --git a/litellm/router_strategy/complexity_router/README.md b/litellm/router_strategy/complexity_router/README.md index 267def7cc5b..38dfd143cbc 100644 --- a/litellm/router_strategy/complexity_router/README.md +++ b/litellm/router_strategy/complexity_router/README.md @@ -367,6 +367,9 @@ local scoring shortcut in `heuristic_first` and `hybrid` modes. The configured c a native OpenAI or Azure OpenAI Responses deployment with access to the encrypted content. The provider handles the encrypted task, and the classifier still chooses the tier dynamically +Compatibility is checked after normal deployment selection. A paused incompatible member of the +classifier group does not prevent an eligible compatible deployment from classifying the task + Unsupported classifier deployments and provider decryption errors use the existing `classifier_fallback` policy. No fixed tier is introduced for encrypted tasks. Plaintext asks and requests carrying only historical encrypted reasoning retain the existing classifier path diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 77de4df3cd7..6df74b26f60 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -25,7 +25,7 @@ from threading import Lock from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast -from pydantic import BaseModel, TypeAdapter, create_model +from pydantic import BaseModel, TypeAdapter, ValidationError, create_model from litellm._logging import verbose_router_logger from litellm.constants import ( @@ -504,7 +504,10 @@ def _encrypted_classifier_task( raw_input: Final = (request_kwargs or EMPTY_MAPPING).get("input") if not isinstance(raw_input, list) or (request_kwargs or EMPTY_MAPPING).get("messages"): return None - items: Final = TypeAdapter(tuple[dict[str, object], ...]).validate_python(raw_input) + try: + items: Final = TypeAdapter(tuple[dict[str, object], ...]).validate_python(raw_input) + except ValidationError: + return None current: Final = next( ( item @@ -516,7 +519,10 @@ def _encrypted_classifier_task( ) if current is None or current.get("type") != "agent_message" or not isinstance(current.get("content"), list): return None - parts: Final = TypeAdapter(tuple[dict[str, object], ...]).validate_python(current["content"]) + try: + parts: Final = TypeAdapter(tuple[dict[str, object], ...]).validate_python(current["content"]) + except ValidationError: + return None if not any(part.get("type") == "encrypted_content" and part.get("encrypted_content") for part in parts): return None return { @@ -2058,7 +2064,7 @@ class ComplexityRouter(CustomLogger): classifier_call_params = MappingProxyType({"reasoning_effort": llm_config.reasoning_effort}) payload: Final = ( - self._native_classifier_payload(llm_config.model, messages_for_call, response_format, encrypted_task) + self._native_classifier_payload(messages_for_call, response_format, encrypted_task) if encrypted_task is not None else {"messages": messages_for_call, "response_format": response_format, **classifier_call_params} ) @@ -2098,7 +2104,6 @@ class ComplexityRouter(CustomLogger): def _native_classifier_payload( self, - model: str, messages: list[AllMessageValues], # mutable-ok: existing transformation accepts the SDK message list response_format: Mapping[str, object], encrypted_task: Mapping[str, object], @@ -2106,26 +2111,7 @@ class ComplexityRouter(CustomLogger): from litellm.completion_extras.litellm_responses_transformation.transformation import ( LiteLLMResponsesTransformationHandler, ) - from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider, get_llm_provider - from litellm.types.router import LiteLLM_Params - deployments: Final = self._group_deployments(model) - if not deployments: - raise ValueError("Encrypted task classification requires a native OpenAI Responses classifier deployment") - for params in (LiteLLM_Params.model_validate(deployment.get("litellm_params")) for deployment in deployments): - if declared_authenticating_provider(params.model, params.custom_llm_provider): - raise ValueError( - "Encrypted task classification requires a native OpenAI Responses classifier deployment" - ) - _, provider, _, _ = get_llm_provider(model=params.model, litellm_params=params) - if ( - provider not in ("openai", "azure") - or params.use_chat_completions_api - or params.model.startswith("openai/chat_completions/") - ): - raise ValueError( - "Encrypted task classification requires a native OpenAI Responses classifier deployment" - ) transformation: Final = LiteLLMResponsesTransformationHandler() input_items, instructions = transformation.convert_chat_completion_messages_to_responses_api(messages) llm_config: Final = self.config.classifier_llm_config @@ -2139,6 +2125,7 @@ class ComplexityRouter(CustomLogger): "instructions": instructions, "text": transformation.transform_response_format_to_text_format(dict(response_format)), "store": False, + "_require_encrypted_task_support": True, **reasoning, } diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 780465dad3b..2ea8a8b53fe 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -5,8 +5,10 @@ Tests the rule-based complexity scoring and tier assignment logic. """ import asyncio +from collections.abc import AsyncIterator import json from copy import deepcopy +from functools import partial import logging import sys import time @@ -14,6 +16,7 @@ from typing import Dict, Final, List from unittest.mock import AsyncMock, MagicMock, patch import pytest +import httpx from pydantic import ValidationError import litellm @@ -69,6 +72,7 @@ from litellm.types.router import ( TaggedPreRoutingStrategy, ) from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler requires_semantic_router = pytest.mark.skipif( @@ -2520,11 +2524,21 @@ def _native_classifier_router( classifier_type: str = "llm", deployment_model: str = "openai/gpt-6-astra", failure: Exception | None = None, + native_router: Router | None = None, + http_handler: AsyncHTTPHandler | None = None, ) -> tuple[ComplexityRouter, MagicMock]: dependency: Final = MagicMock( - aresponses=AsyncMock(return_value=_native_classifier_response(output), side_effect=failure), + aresponses=( + native_router.factory_function(partial(litellm.aresponses, client=http_handler), call_type="aresponses") + if native_router is not None + else AsyncMock(return_value=_native_classifier_response(output), side_effect=failure) + ), acompletion=AsyncMock(return_value=_llm_response('{"tier":"SIMPLE"}')), - get_model_list=MagicMock(return_value=[{"litellm_params": {"model": deployment_model}}]), + get_model_list=( + native_router.get_model_list + if native_router is not None + else MagicMock(return_value=[{"litellm_params": {"model": deployment_model}}]) + ), ) return ( ComplexityRouter( @@ -2533,7 +2547,11 @@ def _native_classifier_router( complexity_router_config={ "tiers": {"SIMPLE": "cheap-model", "REASONING": "deep-model"}, "classifier_type": classifier_type, - "classifier_llm_config": {"model": "classifier", "timeout_ms": 100, "reasoning_effort": "low"}, + "classifier_llm_config": { + "model": "classifier", + "timeout_ms": 5000 if native_router is not None else 100, + "reasoning_effort": "low", + }, "heuristic_first_max_tier": "SIMPLE" if classifier_type == "heuristic_first" else None, "hybrid_boundary_margin": 0.01 if classifier_type == "hybrid" else None, "classifier_fallback": "default_model", @@ -2546,6 +2564,18 @@ def _native_classifier_router( ) +@pytest.fixture +async def native_classifier_http() -> AsyncIterator[tuple[AsyncHTTPHandler, MagicMock]]: + respond: Final = MagicMock( + return_value=httpx.Response(200, json=_native_classifier_response('{"tier":"REASONING"}').model_dump()) + ) + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + handler: Final = AsyncHTTPHandler() + await handler.client.aclose() + handler.client = client + yield handler, respond + + class TestEncryptedTaskClassifier: @pytest.mark.asyncio @pytest.mark.parametrize("classifier_type", ["llm", "heuristic_first", "hybrid"]) @@ -2587,11 +2617,15 @@ class TestEncryptedTaskClassifier: assert "Caller constraints" not in call["instructions"] assert "SIMPLE" in call["instructions"] and "REASONING" in call["instructions"] assert call["text"]["format"]["schema"]["properties"]["tier"]["enum"] == [ - "SIMPLE", "MEDIUM", "COMPLEX", "REASONING" + "SIMPLE", + "MEDIUM", + "COMPLEX", + "REASONING", ] assert call["text"]["format"]["strict"] is True assert call["reasoning"] == {"effort": "low"} assert call["store"] is False + assert call["_require_encrypted_task_support"] is True assert call["stream"] is False assert "tools" not in call and "previous_response_id" not in call assert "messages" not in call and "response_format" not in call @@ -2606,13 +2640,25 @@ class TestEncryptedTaskClassifier: @pytest.mark.parametrize( "items", [ - [{"type": "reasoning", "encrypted_content": "opaque-history", "summary": []}, {"role": "user", "content": "hi"}], + [ + {"type": "reasoning", "encrypted_content": "opaque-history", "summary": []}, + {"role": "user", "content": "hi"}, + ], [_encrypted_agent_task(), {"role": "user", "content": "hi"}], [{**_encrypted_agent_task(), "content": [{"type": "input_text", "text": "hi"}]}], [{"role": "user", "content": "gAAAA is plain text"}], - [{"role": "user", "content": "hi"}, {"type": "function_call_output", "call_id": "call_1", "output": "opaque-provider-task"}], + [ + {"role": "user", "content": "hi"}, + {"type": "function_call_output", "call_id": "call_1", "output": "opaque-provider-task"}, + ], + ], + ids=[ + "historical-reasoning", + "older-encrypted-task", + "plaintext-agent", + "ciphertext-looking-text", + "tool-output", ], - ids=["historical-reasoning", "older-encrypted-task", "plaintext-agent", "ciphertext-looking-text", "tool-output"], ) async def test_other_asks_keep_chat_classifier(self, items: list[dict[str, object]]): router, dependency = _native_classifier_router() @@ -2639,9 +2685,28 @@ class TestEncryptedTaskClassifier: dependency.acompletion.assert_not_called() @pytest.mark.asyncio - @pytest.mark.parametrize("deployment_model", ["anthropic/test-classifier", "openai/chat_completions/gpt-6-astra"]) - async def test_incompatible_classifier_does_not_flatten_encryption(self, deployment_model: str): - router, dependency = _native_classifier_router(deployment_model=deployment_model) + @pytest.mark.parametrize( + "deployment_model", + ["anthropic/test-classifier", "openai/chat_completions/gpt-6-astra", "xai/test-classifier"], + ) + async def test_incompatible_classifier_does_not_flatten_encryption( + self, deployment_model: str, native_classifier_http: tuple[AsyncHTTPHandler, MagicMock] + ): + handler, respond = native_classifier_http + native: Final = Router( + model_list=[ + { + "model_name": "classifier", + "litellm_params": { + "model": deployment_model, + "api_key": "test-key", + "api_base": "https://classifier.test/v1", + }, + } + ], + num_retries=0, + ) + router, _ = _native_classifier_router(native_router=native, http_handler=handler) result: Final = await router.async_pre_routing_hook( model="encrypted-router", request_kwargs={"input": [_encrypted_agent_task()]} @@ -2649,8 +2714,72 @@ class TestEncryptedTaskClassifier: assert result.model == "deep-model" assert result.routing_decision["cause"] == "default_model_fallback" + respond.assert_not_called() + + @pytest.mark.asyncio + @pytest.mark.parametrize("blocked", [True, False]) + async def test_native_classifier_validates_selected_deployment( + self, blocked: bool, native_classifier_http: tuple[AsyncHTTPHandler, MagicMock] + ): + handler, respond = native_classifier_http + native: Final = Router( + model_list=[ + { + "model_name": "classifier", + "litellm_params": {"model": "anthropic/test-classifier", "api_key": "test-key", "order": 0}, + "model_info": {"id": "incompatible", "blocked": blocked}, + }, + { + "model_name": "classifier", + "litellm_params": { + "model": "openai/gpt-6-astra", + "api_key": "test-key", + "order": 1, + "api_base": "https://classifier.test/v1", + }, + "model_info": {"id": "compatible"}, + }, + ], + num_retries=0, + ) + router, _ = _native_classifier_router(native_router=native, http_handler=handler) + task: Final = _encrypted_agent_task() + + result: Final = await router.async_pre_routing_hook(model="encrypted-router", request_kwargs={"input": [task]}) + + assert result.model == "deep-model" + assert result.routing_decision["cause"] == ("llm_classifier" if blocked else "default_model_fallback") + if blocked: + respond.assert_called_once() + request: Final = respond.call_args.args[0] + assert request.url.path == "/v1/responses" + body: Final = json.loads(request.content) + assert body["input"][-1] == task + assert "_require_encrypted_task_support" not in body + else: + respond.assert_not_called() + + @pytest.mark.asyncio + @pytest.mark.parametrize("classifier_type", ["llm", "heuristic_first", "hybrid"]) + @pytest.mark.parametrize( + "input_items", + [ + ["unsupported-input-item"], + [{**_encrypted_agent_task(), "content": [{"type": "input_text", "text": "hi"}, None]}], + ], + ) + async def test_encrypted_detection_does_not_reject_other_input_shapes( + self, classifier_type: str, input_items: list[object] + ): + router, dependency = _native_classifier_router(classifier_type=classifier_type) + + result: Final = await router.aclassify("hi", request_kwargs={"input": input_items}) + + assert result.cause != "default_model_fallback" + assert result.tier == ComplexityTier.SIMPLE dependency.aresponses.assert_not_called() - dependency.acompletion.assert_not_called() + if classifier_type == "llm": + dependency.acompletion.assert_awaited_once() @pytest.mark.asyncio @pytest.mark.parametrize("failure", [ValueError("invalid_encrypted_content"), TimeoutError("classifier timed out")]) From 55ab5ee53d48ef165886de31144b1e172f8208e1 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Thu, 10 Sep 2026 13:08:46 -0700 Subject: [PATCH 47/60] fix(mcp): preserve identity checks across cached OAuth credentials --- AGENTS.md | 2 + litellm/proxy/_experimental/mcp_server/db.py | 34 +++++--- .../mcp_server/oauth2_token_cache.py | 25 ++++-- .../mcp_server/oauth_identity_binding.py | 2 +- .../authz_code_refresher.py | 11 +++ .../outbound_credentials/oauth_token_store.py | 1 + .../per_user_oauth_store.py | 33 ++++++-- .../outbound_credentials/token_cache_codec.py | 36 ++++++-- .../outbound_credentials/v2_token_store.py | 2 + .../test_authz_code_refresher.py | 18 ++-- .../test_per_user_oauth_store.py | 82 +++++++++++++++++-- .../test_token_cache_codec.py | 14 ++++ .../mcp_server/test_db_credentials.py | 75 +++++++++++++++++ .../mcp_server/test_oauth_identity_binding.py | 30 +++++-- 14 files changed, 311 insertions(+), 54 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 41921fdff4d..a1e8f6f618d 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1 +1,3 @@ Read @CLAUDE.md for coding guidelines + +Before requesting maintainer review, verify the current PR tip passes required CI and code coverage, meets Greptile confidence of at least 4/5, and has acceptable Veria and Bugbot reviews. Inspect warnings and findings, fix actionable issues, and rerun the affected checks and reviewers after changes. Record evidence for any false positive or unavailable review; never treat a pending or missing bot result as a pass. Do not lower coverage thresholds or lint budgets to satisfy a check diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 08f83aeb756..789b2ffaef4 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -6,6 +6,7 @@ from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict, cast +from fastapi import HTTPException from typing_extensions import ReadOnly from litellm._logging import verbose_proxy_logger @@ -1692,13 +1693,18 @@ async def refresh_user_oauth_token( ) return None - binding_proof: Final = await enforce_oauth_identity_binding( - server=server, - token_response=body, - litellm_user_id=user_id, - grant_type="refresh_token", - refresh_ownership=RefreshTokenPresented(refresh_token), - ) + try: + binding_proof: Final = await enforce_oauth_identity_binding( + server=server, + token_response=body, + litellm_user_id=user_id, + grant_type="refresh_token", + refresh_ownership=RefreshTokenPresented(refresh_token), + ) + except HTTPException as exc: + if exc.status_code != 403: + raise + return None access_token: Final[str | None] = body.get("access_token") if not access_token: @@ -1812,8 +1818,14 @@ async def resolve_user_oauth_access_token( binding: Final = server.oauth_identity_binding enforce_binding: Final = binding is not None and binding.mode == "enforce" - if enforce_binding: - await mcp_per_user_token_cache.delete(user_id, server_id) + if prefetched_creds is None and enforce_binding and binding is not None: + bound_token: Final = await mcp_per_user_token_cache.get_token(user_id, server_id) + if bound_token is not None: + if await credential_binding_matches( + binding, user_id, server_id, {"identity_binding_proof": bound_token.identity_binding_proof} + ): + return bound_token.access_token + await mcp_per_user_token_cache.delete(user_id, server_id) if prefetched_creds is None and not enforce_binding: cached_token: Final = await mcp_per_user_token_cache.get(user_id, server_id) if cached_token is not None: @@ -1848,7 +1860,9 @@ async def resolve_user_oauth_access_token( access_token: Final[str] = cred["access_token"] if prefetched_creds is None: ttl: Final = _compute_per_user_token_ttl(server, _remaining_token_seconds(cred.get("expires_at"))) - await mcp_per_user_token_cache.set(user_id, server_id, access_token, ttl) + await mcp_per_user_token_cache.set( + user_id, server_id, access_token, ttl, identity_binding_proof=cred.get("identity_binding_proof") + ) return access_token except Exception as e: verbose_proxy_logger.warning( diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py index a4ef970b87a..42edc2999ab 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py @@ -28,6 +28,8 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import ( build_upstream_oauth2_token_request, resolve_upstream_resource, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken +from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import OAuthTokenCacheCodec from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, @@ -233,8 +235,17 @@ class MCPPerUserTokenCache: def _cache_key(self, user_id: str, server_id: str) -> str: return f"{MCP_PER_USER_TOKEN_REDIS_KEY_PREFIX}:{user_id}:{server_id}" + def _codec(self) -> OAuthTokenCacheCodec: + return OAuthTokenCacheCodec( + encrypt_value_helper, + lambda blob: decrypt_value_helper(blob, key="mcp_per_user_token", exception_type="debug"), + ) + async def get(self, user_id: str, server_id: str) -> str | None: - """Return the plaintext access_token, or None on miss/error.""" + token: Final = await self.get_token(user_id, server_id) + return token.access_token if token is not None else None + + async def get_token(self, user_id: str, server_id: str) -> OAuthToken | None: try: from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 @@ -242,12 +253,7 @@ class MCPPerUserTokenCache: encrypted: Final = await user_api_key_cache.async_get_cache(key) if encrypted is None: return None - plaintext: Final = decrypt_value_helper( - encrypted, - key="mcp_per_user_token", - exception_type="debug", - ) - return plaintext or None + return self._codec().decode(encrypted) except Exception as exc: verbose_logger.debug( "MCPPerUserTokenCache.get failed for user=%s server=%s: %s", @@ -263,13 +269,16 @@ class MCPPerUserTokenCache: server_id: str, access_token: str, ttl: int, + identity_binding_proof: str | None = None, ) -> None: """Store NaCl-encrypted access_token in Redis with the given TTL.""" try: from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 key: Final = self._cache_key(user_id, server_id) - encrypted: Final = encrypt_value_helper(access_token) + encrypted: Final = self._codec().encode( + OAuthToken(access_token=access_token, identity_binding_proof=identity_binding_proof) + ) await user_api_key_cache.async_set_cache(key, encrypted, ttl=ttl) verbose_logger.debug( "MCPPerUserTokenCache.set: cached token for user=%s server=%s ttl=%ds", diff --git a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py index 7ca9e6a23df..f02e6c85d9b 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py @@ -342,7 +342,7 @@ async def _evaluate_binding( claims: Final = _decode_id_token(id_token, binding, signing_key) if isinstance(claims, _BindingRejection): return claims - if grant_type == "authorization_code": + if grant_type == "authorization_code" and (binding.mode == "enforce" or expected_nonce is not None): nonce: Final = claims.get("nonce") if not expected_nonce or not isinstance(nonce, str) or not hmac.compare_digest(nonce, expected_nonce): return _BindingRejection( diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py index 824459fd9b9..af1e82eab82 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py @@ -14,6 +14,8 @@ import time from collections.abc import Awaitable, Callable from typing import TYPE_CHECKING, Final, Protocol +from fastapi import HTTPException + from litellm._logging import verbose_logger from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( TokenEndpointAuthConfigError, @@ -93,6 +95,14 @@ class AuthorizationCodeRefresher: self._identity_validator = identity_validator async def refresh(self, user_id: str, server_id: str, token: OAuthToken) -> OAuthToken | None: + try: + return await self._refresh(user_id, server_id, token) + except HTTPException as exc: + if exc.status_code != 403: + raise + return None + + async def _refresh(self, user_id: str, server_id: str, token: OAuthToken) -> OAuthToken | None: if token.refresh_token is None: return None server: Final = self._server_lookup(server_id) @@ -162,4 +172,5 @@ class AuthorizationCodeRefresher: expires_at=self._clock() + expires_in if expires_in is not None else None, refresh_token=new_refresh, scopes=scopes, + identity_binding_proof=binding_proof, ) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py index e0cd9e8e5ed..0ac4296f498 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py @@ -41,6 +41,7 @@ class OAuthToken: expires_at: float | None = None refresh_token: str | None = None scopes: tuple[str, ...] = () + identity_binding_proof: str | None = None def __repr__(self) -> str: has_refresh: Final = self.refresh_token is not None diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py index efcc613c539..2897c3e8e4a 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py @@ -12,6 +12,7 @@ from __future__ import annotations import asyncio from collections.abc import Callable, Mapping +from functools import partial from typing import TYPE_CHECKING, Final from litellm._logging import verbose_logger @@ -39,7 +40,6 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_cod OAuthTokenCacheCodec, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.v2_token_store import ( - CredentialReader, V2PerUserTokenStore, ) @@ -146,12 +146,26 @@ def _runtime_backend_and_coordinator() -> tuple[TokenCacheBackend | None, Refres return backend, coordinator, True +async def _read_bound_credential( + server_lookup: ServerLookup, user_id: str, server_id: str +) -> Mapping[str, object] | None: + credential: Final = await _read_credential(user_id, server_id) + server: Final = server_lookup(server_id) + binding: Final = server.oauth_identity_binding if server else None + if credential is not None and binding is not None and binding.mode == "enforce": + if not await credential_binding_matches(binding, user_id, server_id, credential): + return None + return credential + + def _build_per_user_oauth_token_store( server_lookup: ServerLookup, ) -> tuple[CachedOAuthTokenStore, bool]: backend, coordinator, uses_redis = _runtime_backend_and_coordinator() refresher: Final = AuthorizationCodeRefresher(server_lookup, _post_token_endpoint, _persist_credential) - refreshing: Final = RefreshingTokenStore(V2PerUserTokenStore(_read_credential), refresher, coordinator=coordinator) + refreshing: Final = RefreshingTokenStore( + V2PerUserTokenStore(partial(_read_bound_credential, server_lookup)), refresher, coordinator=coordinator + ) return CachedOAuthTokenStore(refreshing, default_ttl_seconds=_DEFAULT_TTL_SECONDS, backend=backend), uses_redis @@ -176,10 +190,8 @@ class LazyPerUserOAuthTokenStore: *, store_builder: StoreBuilder = _build_per_user_oauth_token_store, redis_available: Callable[[], bool] = _redis_cache_is_available, - credential_reader: CredentialReader = _read_credential, ) -> None: self._server_lookup = server_lookup - self._credential_reader = credential_reader self._store_builder = store_builder self._redis_available = redis_available self._store: InvalidatableOAuthTokenStore | None = None @@ -188,13 +200,18 @@ class LazyPerUserOAuthTokenStore: self._local_fetches = 0 async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: + token: Final = await self._fetch_token(user_id, server_id) server: Final = self._server_lookup(server_id) binding: Final = server.oauth_identity_binding if server else None - if binding is not None and binding.mode == "enforce": - credential: Final = await self._credential_reader(user_id, server_id) - await self.invalidate(user_id, server_id) - if credential is None or not await credential_binding_matches(binding, user_id, server_id, credential): + if token is not None and binding is not None and binding.mode == "enforce": + if not await credential_binding_matches( + binding, user_id, server_id, {"identity_binding_proof": token.identity_binding_proof} + ): + await self.invalidate(user_id, server_id) return None + return token + + async def _fetch_token(self, user_id: str, server_id: str) -> OAuthToken | None: if self._uses_redis: store = self._store if store is not None: diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_cache_codec.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_cache_codec.py index 14cc309b685..56dc3f19cdd 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_cache_codec.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_cache_codec.py @@ -1,24 +1,26 @@ """Serialize + encrypt boundary for caching an OAuth token in a shared (Redis) cache. -A cross-replica cache must serialize the token, and a plaintext bearer in Redis is a leak, so this -encrypts the value (NaCl in production via the injected ``encrypt``, identity in tests). It caches -**only** the ``access_token``: the hot path needs just the bearer, expiry is carried by the cache -entry's TTL (set from the token's ``expires_at`` by the cache), and the long-lived refresh_token stays -in the DB - the refresh path is always a cache miss that re-reads it - so it never reaches Redis. A -decoded token therefore carries only the bearer (``expires_at`` and ``refresh_token`` both None); the -TTL, not the value, bounds its life. An empty/undecryptable blob (e.g. master-key rotation) is a miss. +Shared cache values contain an encrypted access token and optional identity-binding proof. +Refresh tokens remain in the database; cache TTL bounds the access token's lifetime. +Legacy bearer-only entries decode without proof and cannot satisfy identity enforcement. """ from __future__ import annotations +import json from collections.abc import Callable from dataclasses import dataclass from typing import Final +from pydantic import TypeAdapter, ValidationError + from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( OAuthToken, ) +_BOUND_PREFIX: Final = "litellm-bound-oauth-v1:" +_BOUND_PAYLOAD: Final = TypeAdapter(dict[str, str]) + @dataclass(frozen=True, slots=True) class OAuthTokenCacheCodec: @@ -26,10 +28,30 @@ class OAuthTokenCacheCodec: decrypt: Callable[[str], str | None] def encode(self, token: OAuthToken) -> str: + if token.identity_binding_proof is not None: + return self.encrypt( + _BOUND_PREFIX + + json.dumps( + { + "access_token": token.access_token, + "identity_binding_proof": token.identity_binding_proof, + } + ) + ) return self.encrypt(token.access_token) def decode(self, blob: str) -> OAuthToken | None: access_token: Final = self.decrypt(blob) if not access_token: return None + if access_token.startswith(_BOUND_PREFIX): + try: + payload: Final = _BOUND_PAYLOAD.validate_json(access_token[len(_BOUND_PREFIX) :]) + except ValidationError: + return None + bearer: Final = payload.get("access_token") + proof: Final = payload.get("identity_binding_proof") + if not bearer or not proof: + return None + return OAuthToken(access_token=bearer, identity_binding_proof=proof) return OAuthToken(access_token=access_token, refresh_token=None) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py index 4d88ec8b025..0f18931d118 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py @@ -46,11 +46,13 @@ def _to_oauth_token(payload: Mapping[str, object]) -> OAuthToken | None: return None refresh_token: Final = payload.get("refresh_token") expires_at: Final = payload.get("expires_at") + binding_proof: Final = payload.get("identity_binding_proof") return OAuthToken( access_token=access_token, expires_at=_iso_to_epoch(expires_at) if isinstance(expires_at, str) else None, refresh_token=refresh_token if isinstance(refresh_token, str) else None, scopes=_to_scopes(payload.get("scopes")), + identity_binding_proof=binding_proof if isinstance(binding_proof, str) else None, ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py index 16707f6ace8..df068b60338 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py @@ -293,17 +293,18 @@ async def test_refresh_uses_admin_entered_token_url_when_issuer_yield_empties_re @pytest.mark.asyncio async def test_identity_rejection_never_persists_or_returns_refreshed_token(): from unittest.mock import AsyncMock + from fastapi import HTTPException validator = AsyncMock(side_effect=HTTPException(status_code=403, detail="oauth_principal_mismatch")) persist = AsyncMock() refresher = AuthorizationCodeRefresher( - _lookup(_Server()), _endpoint({"access_token": "foreign-token", "id_token": "foreign-identity"}), - persist, identity_validator=validator, + _lookup(_Server()), + _endpoint({"access_token": "foreign-token", "id_token": "foreign-identity"}), + persist, + identity_validator=validator, ) - with pytest.raises(HTTPException) as error: - await refresher.refresh("alice", "srv", OAuthToken(access_token="old", refresh_token="old-rt")) - assert error.value.status_code == 403 + assert await refresher.refresh("alice", "srv", OAuthToken(access_token="old", refresh_token="old-rt")) is None persist.assert_not_awaited() @@ -314,10 +315,13 @@ async def test_verified_refresh_preserves_binding_proof_in_storage(): validator = AsyncMock(return_value="verified-binding") persist = AsyncMock() refresher = AuthorizationCodeRefresher( - _lookup(_Server()), _endpoint({"access_token": "new", "refresh_token": "rotated"}), - persist, identity_validator=validator, + _lookup(_Server()), + _endpoint({"access_token": "new", "refresh_token": "rotated"}), + persist, + identity_validator=validator, ) token = await refresher.refresh("alice", "srv", OAuthToken(access_token="old", refresh_token="old-rt")) assert token.access_token == "new" assert token.refresh_token == "rotated" + assert token.identity_binding_proof == "verified-binding" assert persist.await_args.kwargs["identity_binding_proof"] == "verified-binding" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py index bd190b09c82..e7308884060 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py @@ -254,20 +254,86 @@ async def test_enforcement_invalidates_cached_legacy_credentials_before_use(): from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer server = MCPServer( - server_id="srv", name="srv", transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + server_id="srv", + name="srv", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth_identity_binding=MCPOAuthIdentityBinding( - mode="enforce", issuer="https://idp.example.com", audiences=["client"], + mode="enforce", + issuer="https://idp.example.com", + audiences=["client"], ), ) cached = _RecordingStore("belongs-to-bob") - async def read_legacy(user_id, server_id): - return {"access_token": "belongs-to-bob", "refresh_token": "bobs-refresh"} - store = LazyPerUserOAuthTokenStore( - lambda server_id: server, store_builder=lambda lookup: (cached, False), - redis_available=lambda: False, credential_reader=read_legacy, + lambda server_id: server, + store_builder=lambda lookup: (cached, False), + redis_available=lambda: False, ) assert await store.fetch("alice", "srv") is None - assert cached.calls == [] + assert cached.calls == [("alice", "srv")] assert cached.invalidations == [("alice", "srv")] + + +@pytest.mark.asyncio +async def test_enforced_cache_hit_avoids_credential_read_and_rejects_changed_policy(monkeypatch): + from unittest.mock import AsyncMock + + from litellm.proxy._experimental.mcp_server.oauth_identity_binding import current_binding_proof + from litellm.proxy._experimental.mcp_server.outbound_credentials import per_user_oauth_store as module + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="srv", + name="srv", + transport="http", + auth_type="oauth2", + oauth_identity_binding={ + "mode": "enforce", + "issuer": "https://idp.example", + "audiences": ["client"], + "caller_field": "user_id", + "principal_claim": "sub", + }, + ) + proof = await current_binding_proof(server.oauth_identity_binding, "alice", "srv") + read = AsyncMock(return_value={"access_token": "alice-token", "identity_binding_proof": proof}) + monkeypatch.setattr(module, "_read_credential", read) + monkeypatch.setattr(module, "_runtime_backend_and_coordinator", lambda: (None, None, False)) + store = LazyPerUserOAuthTokenStore(lambda _: server, redis_available=lambda: False) + assert (await store.fetch("alice", "srv")).access_token == "alice-token" + assert (await store.fetch("alice", "srv")).access_token == "alice-token" + read.assert_awaited_once_with("alice", "srv") + server.oauth_identity_binding = server.oauth_identity_binding.model_copy(update={"audiences": ["changed"]}) + assert await store.fetch("alice", "srv") is None + read.assert_awaited_once() + assert await store.fetch("alice", "srv") is None + assert read.await_count == 2 + + +@pytest.mark.asyncio +async def test_expired_unverified_credential_never_reaches_refresh(monkeypatch): + from unittest.mock import AsyncMock + + from litellm.proxy._experimental.mcp_server.outbound_credentials import per_user_oauth_store as module + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="srv", + name="srv", + transport="http", + auth_type="oauth2", + token_url="https://idp.example/token", + oauth_identity_binding={"mode": "enforce", "issuer": "https://idp.example", "audiences": ["client"]}, + ) + read = AsyncMock( + return_value={"access_token": "bob", "refresh_token": "bob-refresh", "expires_at": "2000-01-01T00:00:00Z"} + ) + post = AsyncMock() + monkeypatch.setattr(module, "_read_credential", read) + monkeypatch.setattr(module, "_post_token_endpoint", post) + monkeypatch.setattr(module, "_runtime_backend_and_coordinator", lambda: (None, None, False)) + store = LazyPerUserOAuthTokenStore(lambda _: server, redis_available=lambda: False) + assert await store.fetch("alice", "srv") is None + post.assert_not_awaited() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_cache_codec.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_cache_codec.py index 17a7e13af03..12a04c4b4dc 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_cache_codec.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_cache_codec.py @@ -52,3 +52,17 @@ def test_undecryptable_blob_is_a_miss(): def test_empty_plaintext_is_a_miss(): codec = OAuthTokenCacheCodec(encrypt=lambda s: s, decrypt=lambda b: b) assert codec.decode("") is None + + +def test_bound_token_round_trip_preserves_proof_without_refresh_secret(): + codec = _wrapping_codec() + blob = codec.encode(OAuthToken("alice-token", refresh_token="private-refresh", identity_binding_proof="proof")) + assert "private-refresh" not in blob + decoded = codec.decode(blob) + assert decoded == OAuthToken("alice-token", identity_binding_proof="proof") + + +def test_malformed_bound_entries_fail_closed(): + codec = _wrapping_codec() + for payload in ("not-json", "{}", '{"access_token":"at"}', '{"access_token":1,"identity_binding_proof":"p"}'): + assert codec.decode("enc:litellm-bound-oauth-v1:" + payload) is None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py index 8e2adbeb06e..60a5e1a22bb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py @@ -1629,3 +1629,78 @@ async def test_enforcement_rejects_preexisting_unverified_credential(): cred={"access_token": "belongs-to-bob", "refresh_token": "bobs-refresh-token"}, ) assert result is None + + +@pytest.mark.asyncio +async def test_refresh_identity_rejection_returns_reauthentication_without_persisting(monkeypatch): + from fastapi import HTTPException + from litellm.proxy._experimental.mcp_server import db as module + + validator = AsyncMock(side_effect=HTTPException(status_code=403, detail="oauth_principal_mismatch")) + monkeypatch.setattr(module, "enforce_oauth_identity_binding", validator) + result, captured = await _run_refresh( + monkeypatch, _refresh_server(), {"access_token": "bob", "refresh_token": "rotated"} + ) + assert result is None + assert captured["data"]["grant_type"] == "refresh_token" + module.store_user_oauth_credential.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_verified_legacy_cache_reads_avoid_database_and_reject_policy_changes(monkeypatch): + from litellm.caching.caching import DualCache + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import db as module + from litellm.proxy._experimental.mcp_server.oauth2_token_cache import mcp_per_user_token_cache + from litellm.proxy._experimental.mcp_server.oauth_identity_binding import current_binding_proof + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="srv", + name="srv", + transport="http", + auth_type="oauth2", + oauth_identity_binding={ + "mode": "enforce", + "issuer": "https://idp.example", + "audiences": ["client"], + "caller_field": "user_id", + "principal_claim": "sub", + }, + ) + proof = await current_binding_proof(server.oauth_identity_binding, "alice", "srv") + monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache()) + read = AsyncMock(return_value=None) + monkeypatch.setattr(module, "get_user_oauth_credential", read) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + await mcp_per_user_token_cache.set("alice", "srv", "alice-token", 60, identity_binding_proof=proof) + assert await module.resolve_user_oauth_access_token("alice", server) == "alice-token" + assert await module.resolve_user_oauth_access_token("alice", server) == "alice-token" + read.assert_not_awaited() + server.oauth_identity_binding = server.oauth_identity_binding.model_copy(update={"audiences": ["changed"]}) + assert await module.resolve_user_oauth_access_token("alice", server) is None + read.assert_awaited_once() + assert await mcp_per_user_token_cache.get_token("alice", "srv") is None + + +@pytest.mark.asyncio +async def test_unverified_legacy_cache_cannot_bypass_enforcement(monkeypatch): + from litellm.caching.caching import DualCache + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import db as module + from litellm.proxy._experimental.mcp_server.oauth2_token_cache import mcp_per_user_token_cache + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="srv", + name="srv", + transport="http", + auth_type="oauth2", + oauth_identity_binding={"mode": "enforce", "issuer": "https://idp.example", "audiences": ["client"]}, + ) + monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache()) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + monkeypatch.setattr(module, "get_user_oauth_credential", AsyncMock(return_value={"access_token": "bob"})) + await mcp_per_user_token_cache.set("alice", "srv", "bob", 60) + assert await module.resolve_user_oauth_access_token("alice", server) is None + assert await mcp_per_user_token_cache.get("alice", "srv") is None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py index 6f291865e20..0036035f448 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py @@ -13,14 +13,14 @@ from pydantic import ValidationError from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( RefreshOwnershipProven, - VerifiedRefreshToken, RefreshTokenPresented, - current_binding_proof, + VerifiedRefreshToken, _discover_jwks_url, _fetch_issuer_jwks, _load_caller_principal, _load_stored_refresh_token, _select_signing_key, + current_binding_proof, enforce_oauth_identity_binding, ) from litellm.types.mcp import MCPAuth, MCPTransport @@ -471,19 +471,20 @@ async def test_refresh_with_mismatched_id_token_rejected(): @pytest.mark.asyncio -async def test_audit_mode_logs_but_does_not_reject(): +async def test_audit_mode_logs_but_does_not_reject(caplog): token: Final = _sign_id_token({"email": "mallory@example.com", "email_verified": True}) result: Final = await enforce_oauth_identity_binding( server=_server(mode="audit"), token_response={"access_token": "at", "id_token": token}, litellm_user_id="user-a", grant_type="authorization_code", - expected_nonce="test-nonce", refresh_ownership=None, jwks_fetcher=_jwks_fetcher, caller_principal_loader=_caller_loader("alice@example.com"), ) assert result is None + assert "oauth_principal_mismatch" in caplog.text + assert "nonce" not in caplog.text @pytest.mark.asyncio @@ -642,6 +643,25 @@ async def test_binding_proof_rejects_changed_user_or_policy(): def test_identity_binding_rejects_modes_without_gateway_credential_custody(auth_type): with pytest.raises(ValidationError, match="gateway-managed per-user"): MCPServer( - server_id="srv", name="srv", transport=MCPTransport.http, auth_type=auth_type, + server_id="srv", + name="srv", + transport=MCPTransport.http, + auth_type=auth_type, oauth_identity_binding=_server().oauth_identity_binding, ) + + +@pytest.mark.asyncio +async def test_audit_matching_login_without_nonce_does_not_report_failure(caplog): + token: Final = _sign_id_token({"email": "alice@example.com", "email_verified": True}) + result: Final = await enforce_oauth_identity_binding( + server=_server(mode="audit"), + token_response={"access_token": "at", "id_token": token}, + litellm_user_id="user-a", + grant_type="authorization_code", + refresh_ownership=None, + jwks_fetcher=_jwks_fetcher, + caller_principal_loader=_caller_loader("alice@example.com"), + ) + assert result is None + assert "oauth_identity_binding audit" not in caplog.text From b0d66a15b86a5e5d2f503d8c579b5e3a528ab3a0 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Thu, 10 Sep 2026 13:18:41 -0700 Subject: [PATCH 48/60] feat(ocr): add core foundation and Mistral adapter (#40530) * feat(ocr): add core foundation and transport primitives * fix(ocr): decline missing Mistral credentials * fix(rust): compile trace parity on Rust 1.98 * refactor(ocr): define native response capability * refactor(auth): generalize missing API key errors * refactor(core): keep URL helpers usage scoped * refactor(ocr): support native responses across adapters * refactor(ocr): preserve unmapped provider params * refactor(ocr): distinguish request preparation from payload transforms * refactor(ocr): trace payload transformation at codec boundary --- litellm-rust/Cargo.lock | 2 + litellm-rust/Cargo.toml | 1 + .../src/audio_transcription/hooks.rs | 2 +- .../crates/ai-gateway/src/ocr/hooks.rs | 2 +- .../ai-gateway/src/routes/messages/mod.rs | 3 +- .../crates/ai-gateway/src/trace_parity.rs | 8 +- litellm-rust/crates/core/Cargo.toml | 8 + litellm-rust/crates/core/src/constants.rs | 3 + litellm-rust/crates/core/src/error.rs | 107 +++++++++ litellm-rust/crates/core/src/http_utils.rs | 105 +++++++- litellm-rust/crates/core/src/lib.rs | 1 + .../crates/core/src/ocr/adapters/mistral.rs | 147 ++++++++++++ .../crates/core/src/ocr/adapters/mod.rs | 69 ++++++ litellm-rust/crates/core/src/ocr/client.rs | 75 ++++++ .../crates/core/src/ocr/codecs/mistral/mod.rs | 5 + .../src/ocr/codecs/mistral/transformation.rs | 102 ++++++++ .../core/src/ocr/codecs/mistral/types.rs | 53 ++++ .../crates/core/src/ocr/codecs/mod.rs | 1 + litellm-rust/crates/core/src/ocr/error.rs | 42 ++++ litellm-rust/crates/core/src/ocr/handler.rs | 72 ++++++ litellm-rust/crates/core/src/ocr/hooks.rs | 146 +++++++++++ litellm-rust/crates/core/src/ocr/mod.rs | 19 ++ litellm-rust/crates/core/src/ocr/prepare.rs | 142 +++++++++++ litellm-rust/crates/core/src/ocr/registry.rs | 53 ++++ .../crates/core/src/ocr/transformation.rs | 6 +- litellm-rust/crates/core/src/ocr/types.rs | 154 ++++++++++-- litellm-rust/crates/core/src/ocr/wire.rs | 154 ++++++++++++ .../providers/azure_ai/ocr/transformation.rs | 14 +- .../providers/mistral/ocr/transformation.rs | 11 +- .../providers/reducto/ocr/transformation.rs | 10 +- .../providers/vertex_ai/ocr/transformation.rs | 8 +- litellm-rust/crates/core/src/url_utils.rs | 95 ++++++++ litellm-rust/crates/core/tests/ocr.rs | 227 ++++++++++++++++++ litellm-rust/crates/core/tests/ocr/support.rs | 106 ++++++++ .../crates/python-bridge/src/errors.rs | 1 + 35 files changed, 1904 insertions(+), 50 deletions(-) create mode 100644 litellm-rust/crates/core/src/ocr/adapters/mistral.rs create mode 100644 litellm-rust/crates/core/src/ocr/adapters/mod.rs create mode 100644 litellm-rust/crates/core/src/ocr/client.rs create mode 100644 litellm-rust/crates/core/src/ocr/codecs/mistral/mod.rs create mode 100644 litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs create mode 100644 litellm-rust/crates/core/src/ocr/codecs/mistral/types.rs create mode 100644 litellm-rust/crates/core/src/ocr/codecs/mod.rs create mode 100644 litellm-rust/crates/core/src/ocr/error.rs create mode 100644 litellm-rust/crates/core/src/ocr/handler.rs create mode 100644 litellm-rust/crates/core/src/ocr/hooks.rs create mode 100644 litellm-rust/crates/core/src/ocr/prepare.rs create mode 100644 litellm-rust/crates/core/src/ocr/registry.rs create mode 100644 litellm-rust/crates/core/src/ocr/wire.rs create mode 100644 litellm-rust/crates/core/src/url_utils.rs create mode 100644 litellm-rust/crates/core/tests/ocr.rs create mode 100644 litellm-rust/crates/core/tests/ocr/support.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 43b9ec1aac2..7addc4e3e45 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1453,11 +1453,13 @@ dependencies = [ "rstest", "serde", "serde_json", + "serde_path_to_error", "sha2 0.10.9", "thiserror 2.0.19", "tokio", "tracing", "tracing-subscriber", + "url", ] [[package]] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 82de7f40069..00fa6faaaef 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -39,6 +39,7 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"] tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] } futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] } base64 = "0.22" +url = "2.5.8" [profile.release] opt-level = 3 diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs index dbe2d3a325b..098ce071efc 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs @@ -269,7 +269,7 @@ fn guardrail_error_to_core_error(error: GuardrailError) -> Error { fn core_error_kind(error: &Error) -> &'static str { match error { - Error::Auth(_) => "AuthError", + Error::Auth(_) | Error::MissingApiKey { .. } => "AuthError", Error::InvalidProvider(_) => "InvalidProvider", Error::InvalidRequest(_) => "InvalidRequest", Error::InvalidType { .. } => "InvalidType", diff --git a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs index ed41f1ff9e7..d1dd811f7ab 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs @@ -386,7 +386,7 @@ fn guardrail_error_to_core_error(error: GuardrailError) -> Error { fn core_error_kind(error: &Error) -> &'static str { match error { - Error::Auth(_) => "AuthError", + Error::Auth(_) | Error::MissingApiKey { .. } => "AuthError", Error::InvalidProvider(_) => "InvalidProvider", Error::InvalidRequest(_) => "InvalidRequest", Error::InvalidType { .. } => "InvalidType", diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs index bb9f3851a77..c22d05f5726 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -114,7 +114,8 @@ impl IntoResponse for MessagesRouteError { | Error::Connect(_) | Error::InvalidResponse(_) | Error::InvalidType { .. } - | Error::MissingField(_) => ( + | Error::MissingField(_) + | Error::MissingApiKey { .. } => ( StatusCode::BAD_GATEWAY, "messages provider request failed".to_string(), ), diff --git a/litellm-rust/crates/ai-gateway/src/trace_parity.rs b/litellm-rust/crates/ai-gateway/src/trace_parity.rs index 614852c541d..21123df3f1c 100644 --- a/litellm-rust/crates/ai-gateway/src/trace_parity.rs +++ b/litellm-rust/crates/ai-gateway/src/trace_parity.rs @@ -47,10 +47,10 @@ pub async fn messages_request( .header(CONTENT_TYPE, "application/json") .body(Body::from(body.to_string())) .map_err(|error| Error::InvalidRequest(error.to_string()))?; - let response = routes::app(state) - .oneshot(request) - .await - .map_err(|error| match error {})?; + let response = match routes::app(state).oneshot(request).await { + Ok(response) => response, + Err(error) => match error {}, + }; let status: StatusCode = response.status(); let bytes = to_bytes(response.into_body(), usize::MAX) .await diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index c0de7ff3977..b4ca88cf16f 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -4,6 +4,11 @@ version = "0.1.0" edition.workspace = true license.workspace = true repository.workspace = true +autotests = false + +[[test]] +name = "workspace_crate_allowlist" +path = "tests/workspace_crate_allowlist.rs" [dependencies] base64.workspace = true @@ -11,10 +16,13 @@ rand.workspace = true reqwest.workspace = true serde.workspace = true serde_json.workspace = true +serde_path_to_error = "0.1" +tokio.workspace = true thiserror.workspace = true tracing.workspace = true tracing-subscriber = { workspace = true, optional = true } sha2.workspace = true +url.workspace = true aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"], optional = true } aws-sdk-sts = { version = "1.108.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index fc81f4fa029..108d2a48e30 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -43,3 +43,6 @@ pub const EMPTY_TEXT_PLACEHOLDER: &str = "[System: Empty message content sanitised to satisfy protocol]"; pub const FUNCTION_TRACE_TARGET: &str = "litellm::function_trace"; +pub(crate) const OCR_HTTP_TIMEOUT_SECS: u64 = 600; +pub(crate) const OCR_CONNECT_TIMEOUT_SECS: u64 = 10; +pub(crate) const MISTRAL_OCR_API_BASE: &str = "https://api.mistral.ai/v1"; diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index db3fa2ec704..0382314057f 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -17,6 +17,10 @@ pub enum Error { InvalidRequest(String), #[error("{0}")] Auth(String), + #[error( + "Missing {provider} API Key - A call is being made to {provider} but no key is set either in the environment variables or via params" + )] + MissingApiKey { provider: &'static str }, #[error("upstream request failed with status {status}: {body}")] Http { status: u16, body: String }, #[error("upstream network error: {0}")] @@ -36,6 +40,59 @@ pub enum Error { Unsupported(&'static str), } +#[derive(Clone, Debug, ThisError, PartialEq, Eq)] +pub enum TransportError { + #[error("upstream request failed with status {status}: {body}")] + Http { status: u16, body: String }, + #[error("upstream network error: {0}")] + Network(String), + #[error("could not reach the provider: {0}")] + Connect(String), +} + +impl TransportError { + pub fn from_reqwest_before_dispatch(error: reqwest::Error) -> Self { + let before_dispatch = !error.is_timeout() && (error.is_connect() || error.is_builder()); + let message = error.without_url().to_string(); + if before_dispatch { + Self::Connect(message) + } else { + Self::Network(message) + } + } +} + +impl From for TransportError { + fn from(error: reqwest::Error) -> Self { + Self::Network(error.without_url().to_string()) + } +} + +impl From for Error { + fn from(error: crate::ocr::error::OcrRequestError) -> Self { + match error { + crate::ocr::error::OcrRequestError::MissingField(field) => Self::MissingField(field), + error => Self::InvalidRequest(error.to_string()), + } + } +} + +impl From for Error { + fn from(error: crate::ocr::error::OcrResponseError) -> Self { + Self::InvalidResponse(error.to_string()) + } +} + +impl From for Error { + fn from(error: TransportError) -> Self { + match error { + TransportError::Http { status, body } => Self::Http { status, body }, + TransportError::Network(message) => Self::Network(message), + TransportError::Connect(message) => Self::Connect(message), + } + } +} + pub fn json_type_name(value: &serde_json::Value) -> &'static str { match value { serde_json::Value::Null => "null", @@ -46,3 +103,53 @@ pub fn json_type_name(value: &serde_json::Value) -> &'static str { serde_json::Value::Object(_) => "object", } } + +#[cfg(test)] +mod transport_tests { + use super::*; + + #[tokio::test] + async fn transport_errors_remove_urls_and_keep_dispatch_context() { + let error = reqwest::Client::builder() + .no_proxy() + .build() + .expect("client") + .get("http://localhost:invalid/private?api_key=secret") + .send() + .await + .expect_err("invalid port"); + let error = TransportError::from_reqwest_before_dispatch(error); + assert!(matches!(error, TransportError::Connect(_))); + assert!(!error.to_string().contains("secret")); + assert!(!error.to_string().contains("private")); + } + + #[tokio::test] + async fn request_timeout_is_not_safe_to_retry_as_a_connect_failure() { + use std::time::Duration; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind"); + let address = listener.local_addr().expect("address"); + let request = reqwest::Client::builder() + .no_proxy() + .build() + .expect("client") + .get(format!("http://{address}")) + .timeout(Duration::from_millis(200)) + .send(); + let (response, accepted) = tokio::join!( + request, + tokio::time::timeout(Duration::from_secs(2), listener.accept()) + ); + let _connection = accepted + .expect("accept deadline") + .expect("accepted connection"); + let error = response.expect_err("server does not respond"); + assert!(error.is_timeout()); + assert!(matches!( + TransportError::from_reqwest_before_dispatch(error), + TransportError::Network(_) + )); + } +} diff --git a/litellm-rust/crates/core/src/http_utils.rs b/litellm-rust/crates/core/src/http_utils.rs index cb472dd5a57..9299bb77ac8 100644 --- a/litellm-rust/crates/core/src/http_utils.rs +++ b/litellm-rust/crates/core/src/http_utils.rs @@ -1,10 +1,43 @@ -//! Header and upstream-body helpers shared by every route module. - use serde_json::{Map, Value}; use crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS; use crate::error::{Error, json_type_name}; +#[allow( + dead_code, + reason = "used by the OCR architecture in the next stacked PR" +)] +pub(crate) enum HeaderPolicy<'a> { + All, + Only(&'a [&'a str]), + Except(&'a [&'a str]), +} + +#[allow( + dead_code, + reason = "used by the OCR architecture in the next stacked PR" +)] +pub(crate) fn with_headers( + builder: reqwest::RequestBuilder, + headers: &[(String, String)], + policy: HeaderPolicy<'_>, +) -> reqwest::RequestBuilder { + headers + .iter() + .filter(|(name, _)| match policy { + HeaderPolicy::All => true, + HeaderPolicy::Only(names) => names + .iter() + .any(|allowed| name.eq_ignore_ascii_case(allowed)), + HeaderPolicy::Except(names) => !names + .iter() + .any(|excluded| name.eq_ignore_ascii_case(excluded)), + }) + .fold(builder, |builder, (name, value)| { + builder.header(name, value) + }) +} + #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub async fn http_request( request: reqwest::RequestBuilder, @@ -12,8 +45,6 @@ pub async fn http_request( request.send().await } -/// Bound an upstream error body before it crosses a host boundary, so provider -/// bodies stay data-minimized. pub fn truncate_error_body(body: &str) -> String { if body.chars().count() <= UPSTREAM_ERROR_BODY_MAX_CHARS { return body.to_string(); @@ -61,11 +92,77 @@ pub fn has_bearer_auth(headers: &[(String, String)]) -> bool { }) } +#[allow( + dead_code, + reason = "used by the OCR architecture in the next stacked PR" +)] +pub(crate) fn deserialize_optional_param<'de, D, T>( + deserializer: D, +) -> Result>, D::Error> +where + D: serde::Deserializer<'de>, + T: serde::Deserialize<'de>, +{ + as serde::Deserialize>::deserialize(deserializer).map(Some) +} + #[cfg(test)] mod tests { use super::*; use serde_json::json; + #[rstest::rstest] + #[case(HeaderPolicy::All, true, true)] + #[case(HeaderPolicy::Only(&["authorization"]), true, false)] + #[case(HeaderPolicy::Except(&["authorization"]), false, true)] + fn forwarding_policy_preserves_matching_headers_and_duplicates( + #[case] policy: HeaderPolicy<'_>, + #[case] auth: bool, + #[case] trace: bool, + ) { + let request = with_headers( + reqwest::Client::new().get("https://example.com"), + &[ + ("AuThOrIzAtIoN".into(), "Bearer token".into()), + ("X-Trace".into(), "first".into()), + ("x-trace".into(), "second".into()), + ], + policy, + ) + .build() + .unwrap(); + assert_eq!(request.headers().contains_key("authorization"), auth); + let traces: Vec<_> = request.headers().get_all("x-trace").iter().collect(); + if trace { + assert_eq!(traces, ["first", "second"]); + } else { + assert!(traces.is_empty()); + } + } + + #[test] + fn multipart_policy_leaves_content_headers_to_reqwest() { + let request = with_headers( + reqwest::Client::new() + .post("https://example.com") + .multipart(reqwest::multipart::Form::new().text("file", "abc")), + &[ + ("Content-Type".into(), "application/json".into()), + ("CONTENT-LENGTH".into(), "0".into()), + ], + HeaderPolicy::Except(&["content-type", "content-length"]), + ) + .build() + .unwrap(); + assert!( + request.headers()["content-type"] + .to_str() + .unwrap() + .starts_with("multipart/form-data; boundary=") + ); + assert_ne!(request.headers()["content-length"], "0"); + } + #[test] fn truncate_leaves_short_bodies_untouched() { assert_eq!(truncate_error_body("short"), "short"); diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index b93e084f57e..3a0896a4d5a 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -14,5 +14,6 @@ pub mod realtime; pub mod responses; pub mod router; pub mod routing_utils; +mod url_utils; pub use error::Error; diff --git a/litellm-rust/crates/core/src/ocr/adapters/mistral.rs b/litellm-rust/crates/core/src/ocr/adapters/mistral.rs new file mode 100644 index 00000000000..ea569ffb34f --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/mistral.rs @@ -0,0 +1,147 @@ +use super::OcrAdapter; +use crate::Error; +use crate::constants::MISTRAL_OCR_API_BASE; +use crate::ocr::OcrClient; +use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse}; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{ + _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, +}; +use crate::ocr::registry::OcrProvider; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection}; +use crate::url_utils::ApiUrl; + +const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY"; + +#[derive(Clone, Debug)] +pub(crate) struct MistralAdapter; + +impl OcrAdapter for MistralAdapter { + type ProviderResponse = MistralOcrResponse; + const PROVIDER: OcrProvider = OcrProvider::Mistral; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let ParsedProviderParams { + known: params, + extra_params: _extra_params, + } = _prepare_ocr_request::(request)?; + let headers = validate_environment(&request.connection, &credential_env)?; + let url = get_complete_url(request.connection.api_base.as_deref())?; + let body = + mistral::transform_ocr_request(&request.model, request.document.clone(), ¶ms)?; + transform_request_body(client, request, &url, &headers, body, |_| Ok(())).await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result { + mistral::transform_ocr_response(&request.model, response) + } +} + +pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result { + let base = api_base + .map(str::trim) + .filter(|base| !base.is_empty()) + .unwrap_or(MISTRAL_OCR_API_BASE); + ApiUrl::parse(base) + .and_then(|url| url.complete_path(&["v1", "ocr"])) + .map(|url| url.into_string()) + .map_err(|_| { + OcrRequestError::RequestField { + path: "api_base".into(), + } + .into() + }) +} + +fn validate_environment( + connection: &OcrConnection, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Result, OcrError> { + if crate::http_utils::has_header(&connection.extra_headers, "authorization") { + return Ok(connection.extra_headers.clone()); + } + let api_key = connection + .api_key + .as_deref() + .map(str::trim) + .filter(|key| !key.is_empty()) + .map(str::to_string) + .or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty())) + .ok_or(Error::MissingApiKey { + provider: "Mistral", + })?; + Ok( + std::iter::once(("Authorization".into(), format!("Bearer {api_key}"))) + .chain(connection.extra_headers.clone()) + .collect(), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn complete_url_defaults_and_dedupes_v1() { + assert_eq!( + get_complete_url(None).unwrap(), + "https://api.mistral.ai/v1/ocr" + ); + assert_eq!( + get_complete_url(Some("https://example.com/v1?tenant=a")).unwrap(), + "https://example.com/v1/ocr?tenant=a" + ); + assert_eq!( + get_complete_url(Some("https://example.com/v1/ocr?tenant=a")).unwrap(), + "https://example.com/v1/ocr?tenant=a" + ); + } + + #[test] + fn environment_prefers_explicit_key_then_environment() { + let explicit = OcrConnection { + api_key: Some("explicit".into()), + ..OcrConnection::default() + }; + assert_eq!( + validate_environment(&explicit, &|_| Some("environment".into())).unwrap()[0], + ("Authorization".into(), "Bearer explicit".into()) + ); + + assert_eq!( + validate_environment(&OcrConnection::default(), &|_| Some("environment".into())) + .unwrap()[0], + ("Authorization".into(), "Bearer environment".into()) + ); + } + + #[test] + fn environment_preserves_forwarded_authorization() { + let connection = OcrConnection { + extra_headers: vec![("authorization".into(), "Bearer forwarded".into())], + ..OcrConnection::default() + }; + assert_eq!( + validate_environment(&connection, &|_| None).unwrap(), + connection.extra_headers + ); + } + + #[test] + fn environment_rejects_missing_key() { + assert!(matches!( + validate_environment(&OcrConnection::default(), &|_| None), + Err(OcrError::Public(Error::MissingApiKey { + provider: "Mistral" + })) + )); + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/mod.rs new file mode 100644 index 00000000000..7dfb08b4dc2 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/mod.rs @@ -0,0 +1,69 @@ +use std::future::Future; + +use serde::de::DeserializeOwned; + +use super::OcrClient; +use super::error::{OcrError, OcrResponseError}; +use super::registry::OcrProvider; +use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrResponseFormat}; +use super::wire::DecodedOcrResponse; + +mod mistral; + +pub(crate) use mistral::MistralAdapter; + +/// Converts a complete LiteLLM OCR call to provider HTTP and normalizes its response. +pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static { + /// Provider JSON schema; direct and Vertex Mistral share `MistralOcrResponse`. + type ProviderResponse: DeserializeOwned + Send; + + const PROVIDER: OcrProvider; + + /// Prepares the complete provider HTTP request. + /// `request` contains the model, document, connection, and unmapped caller options. + /// `client` supplies reusable provider and document HTTP clients. + /// Returns the complete HTTP request, whereas Python returns body data. + fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> impl Future> + Send; + + /// Python: `transform_ocr_response`. + /// `request` supplies caller context, including the fallback model. + /// `response` is the decoded provider payload; the output is the shared LiteLLM schema. + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result; + + /// Decodes provider HTTP; adapters may override this to poll asynchronous operations. + /// Python performs that polling inside `async_transform_ocr_response`. + /// `client` is reused for polling; `response` is the initial HTTP response. + /// `url` and `headers` describe the submitted call; `request` supplies limits and format. + fn read_response( + &self, + _client: &OcrClient, + response: reqwest::Response, + _url: &str, + _headers: &[(String, String)], + request: &LiteLLMOcrRequest, + ) -> impl Future, OcrError>> + Send + { + let retain_native = request + .response_format() + .map(|format| format == OcrResponseFormat::Native); + async move { super::client::read_json_response(response, retain_native?).await } + } +} + +macro_rules! for_each_ocr_adapter { + ($callback:ident) => { + $callback! { + Mistral, $crate::ocr::adapters::MistralAdapter, $crate::ocr::adapters::MistralAdapter, Mistral; + } + }; +} + +pub(crate) use for_each_ocr_adapter; diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs new file mode 100644 index 00000000000..36bfc036bae --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -0,0 +1,75 @@ +use std::sync::OnceLock; +use std::time::Duration; + +use serde::de::DeserializeOwned; + +use super::error::OcrError; +use super::handler::perform_ocr_request; +use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; +use super::wire::{DecodedOcrResponse, decode_response}; +use crate::Error; +use crate::constants::OCR_CONNECT_TIMEOUT_SECS; +use crate::error::TransportError; + +#[derive(Clone)] +pub struct OcrClient { + provider_http: reqwest::Client, +} + +impl OcrClient { + pub fn new(provider_http: reqwest::Client) -> Result { + Ok(Self { provider_http }) + } + + #[tracing::instrument( + name = "ocr", + target = "litellm::function_trace", + level = "trace", + skip_all + )] + pub async fn perform(&self, request: LiteLLMOcrRequest) -> Result { + perform_ocr_request(self, request).await + } + + pub(crate) fn provider_http(&self) -> &reqwest::Client { + &self.provider_http + } + + #[cfg(test)] + pub(crate) fn for_test(provider_http: reqwest::Client) -> Self { + Self { provider_http } + } +} + +pub async fn ocr(request: LiteLLMOcrRequest) -> Result { + static CLIENT: OnceLock> = OnceLock::new(); + let client = CLIENT + .get_or_init(|| { + reqwest::Client::builder() + .connect_timeout(Duration::from_secs(OCR_CONNECT_TIMEOUT_SECS)) + .build() + .map_err(TransportError::from) + .and_then(OcrClient::new) + }) + .clone()?; + client.perform(request).await +} + +pub async fn read_json_response( + response: reqwest::Response, + native: bool, +) -> Result, OcrError> { + let status = response.status(); + let bytes = response + .bytes() + .await + .map_err(crate::error::TransportError::from)?; + if !status.is_success() { + return Err(crate::error::TransportError::Http { + status: status.as_u16(), + body: crate::http_utils::truncate_error_body(&String::from_utf8_lossy(&bytes)), + } + .into()); + } + Ok(decode_response(&bytes, native)?) +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/mistral/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/mistral/mod.rs new file mode 100644 index 00000000000..eea4254779e --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/mistral/mod.rs @@ -0,0 +1,5 @@ +mod transformation; +mod types; + +pub(crate) use transformation::{transform_ocr_request, transform_ocr_response}; +pub(crate) use types::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse}; diff --git a/litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs b/litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs new file mode 100644 index 00000000000..cd0a1dc6b17 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs @@ -0,0 +1,102 @@ +use super::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse}; +use crate::ocr::error::{OcrRequestError, OcrResponseError}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub(crate) fn transform_ocr_request( + model: &str, + document: OcrDocument, + params: &MistralOcrParams, +) -> Result { + Ok(MistralOcrRequest { + model: model.to_string(), + document, + params: params.clone(), + }) +} + +pub(crate) fn transform_ocr_response( + model: &str, + response: MistralOcrResponse, +) -> Result { + Ok(LiteLLMOcrResponse { + pages: response.pages, + model: response.model.unwrap_or_else(|| model.to_string()), + document_annotation: response.document_annotation, + usage_info: response.usage_info, + object: "ocr".to_string(), + extra_fields: response.extra_fields, + provider_native_response: None, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use rstest::rstest; + use serde_json::{Value, json}; + + #[rstest] + #[case("pages", json!([0, 2]))] + #[case("include_image_base64", json!(true))] + #[case("image_limit", json!(2))] + #[case("image_min_size", json!(100))] + #[case("bbox_annotation_format", json!({"type":"json_schema"}))] + #[case("document_annotation_format", json!({"type":"json_schema"}))] + #[case("document_annotation_prompt", json!("extract"))] + #[case("extract_header", json!(true))] + #[case("extract_footer", json!(false))] + #[case("table_format", json!("html"))] + #[case("confidence_scores_granularity", json!("word"))] + #[case("include_blocks", json!(true))] + #[case("id", json!("req-123"))] + fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) { + let params: MistralOcrParams = + serde_json::from_value(json!({name: value.clone()})).unwrap(); + let document: OcrDocument = serde_json::from_value( + json!({"type":"document_url","document_url":"https://example.com/a.pdf"}), + ) + .unwrap(); + let result = + serde_json::to_value(transform_ocr_request("model", document, ¶ms).unwrap()) + .unwrap(); + assert_eq!(result["model"], "model"); + assert_eq!(result[name], value); + } + + #[test] + fn request_mapping_filters_unknown_fields() { + let params: MistralOcrParams = serde_json::from_value(json!({"unknown": true})).unwrap(); + let document: OcrDocument = serde_json::from_value( + json!({"type":"document_url","document_url":"https://example.com/a.pdf"}), + ) + .unwrap(); + let result = + serde_json::to_value(transform_ocr_request("model", document, ¶ms).unwrap()) + .unwrap(); + assert!(result.get("unknown").is_none()); + } + + #[test] + fn response_preserves_provider_fields() { + let response: MistralOcrResponse = serde_json::from_value(json!({ + "pages":[{"index":0,"markdown":"hello","header":"head","confidence_scores":{"mean":0.99}}], + "model":"returned-model", + "usage_info":{"pages_processed":1,"future_counter":5}, + "future_response_field":"kept" + })) + .unwrap(); + let result = transform_ocr_response("model", response) + .unwrap() + .into_json(); + assert_eq!(result["pages"][0]["header"], "head"); + assert_eq!(result["usage_info"]["future_counter"], 5); + assert_eq!(result["future_response_field"], "kept"); + assert_eq!(result["model"], "returned-model"); + } + + #[test] + fn response_rejects_null_pages() { + assert!(serde_json::from_value::(json!({"pages":null})).is_err()); + } +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/mistral/types.rs b/litellm-rust/crates/core/src/ocr/codecs/mistral/types.rs new file mode 100644 index 00000000000..0e601cd8319 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/mistral/types.rs @@ -0,0 +1,53 @@ +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; + +use crate::ocr::types::OcrDocument; + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub(crate) struct MistralOcrParams { + #[serde(skip_serializing_if = "Option::is_none")] + pub pages: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub include_image_base64: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub image_limit: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub image_min_size: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub bbox_annotation_format: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub document_annotation_format: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub document_annotation_prompt: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub extract_header: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub extract_footer: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub table_format: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub confidence_scores_granularity: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub include_blocks: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub id: Option, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct MistralOcrRequest { + pub model: String, + pub document: OcrDocument, + #[serde(flatten)] + pub params: MistralOcrParams, +} + +#[derive(Clone, Debug, Default, Deserialize)] +pub(crate) struct MistralOcrResponse { + #[serde(default)] + pub pages: Vec, + pub model: Option, + pub document_annotation: Option, + pub usage_info: Option, + #[serde(flatten)] + pub extra_fields: Map, +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/mod.rs new file mode 100644 index 00000000000..170ef5f68a7 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/mod.rs @@ -0,0 +1 @@ +pub(crate) mod mistral; diff --git a/litellm-rust/crates/core/src/ocr/error.rs b/litellm-rust/crates/core/src/ocr/error.rs new file mode 100644 index 00000000000..50793aad566 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/error.rs @@ -0,0 +1,42 @@ +use thiserror::Error; + +use crate::error::TransportError; + +#[derive(Debug, Clone, PartialEq, Eq, Error)] +pub enum OcrRequestError { + #[error("Invalid `req_format`. Expected 'native' or 'litellm'.")] + RequestFormat, + #[error("invalid OCR request field: {path}")] + RequestField { path: String }, + #[error("missing required field: {0}")] + MissingField(&'static str), +} + +#[derive(Debug, Clone, PartialEq, Eq, Error)] +pub enum OcrResponseError { + #[error("invalid OCR response field: {path}")] + ResponseField { path: String }, +} + +#[derive(Debug, Error)] +pub enum OcrError { + #[error("{0}")] + Request(#[from] OcrRequestError), + #[error("{0}")] + Response(#[from] OcrResponseError), + #[error("{0}")] + Transport(#[from] TransportError), + #[error("{0}")] + Public(#[from] crate::Error), +} + +impl From for crate::Error { + fn from(error: OcrError) -> Self { + match error { + OcrError::Request(error) => error.into(), + OcrError::Response(error) => error.into(), + OcrError::Transport(error) => error.into(), + OcrError::Public(error) => error, + } + } +} diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs new file mode 100644 index 00000000000..0b04319d966 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -0,0 +1,72 @@ +use super::OcrClient; +use super::adapters::OcrAdapter; +use super::hooks::OcrLifecycleHooks; +use super::registry::OcrAdapterKind; +use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; +use crate::Error; +use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext}; + +pub(crate) async fn perform_ocr_request( + client: &OcrClient, + request: LiteLLMOcrRequest, +) -> Result { + let context = CallLifecycleContext::new( + "ocr", + request.model.clone(), + request.adapter.provider().as_str(), + request + .litellm_call_id + .clone() + .unwrap_or_else(|| format!("ocr-{:032x}", rand::random::())), + ); + let hooks = OcrLifecycleHooks { + hooks: request.hooks.clone(), + provider_name: context.custom_llm_provider.clone(), + }; + CallLifecycle::default().run(context, request, &hooks, |request| async move { + macro_rules! execute_selected_adapter { + ($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => { + match request.adapter { + $( OcrAdapterKind::$variant => execute_ocr_provider_call(client, &$instance, request).await, )+ + } + }; + } + super::adapters::for_each_ocr_adapter!(execute_selected_adapter) + }).await +} + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +async fn execute_ocr_provider_call( + client: &OcrClient, + adapter: &A, + request: LiteLLMOcrRequest, +) -> Result { + let provider_request = adapter.prepare_request(&request, client).await?; + let url = provider_request.url().to_string(); + let headers = provider_request + .headers() + .iter() + .map(|(name, value)| { + value + .to_str() + .map(|value| (name.to_string(), value.to_string())) + .map_err(|_| super::error::OcrRequestError::RequestField { + path: "headers".into(), + }) + }) + .collect::, _>>()?; + let response = crate::http_utils::http_request(reqwest::RequestBuilder::from_parts( + client.provider_http().clone(), + provider_request, + )) + .await + .map_err(crate::error::TransportError::from)?; + let decoded = adapter + .read_response(client, response, &url, &headers, &request) + .await?; + let response = adapter.transform_ocr_response(&request, decoded.data)?; + Ok(LiteLLMOcrResponse { + provider_native_response: decoded.native, + ..response + }) +} diff --git a/litellm-rust/crates/core/src/ocr/hooks.rs b/litellm-rust/crates/core/src/ocr/hooks.rs new file mode 100644 index 00000000000..7dd3c6bf8b2 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/hooks.rs @@ -0,0 +1,146 @@ +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; + +use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument}; +use crate::Error; +use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; +use serde::Serialize; +use serde_json::Value; + +pub type OcrHookFuture<'a, T> = Pin> + Send + 'a>>; +pub type OcrLogFuture<'a> = Pin + Send + 'a>>; + +#[derive(Clone, Debug, Serialize)] +pub struct OcrPreCallRequest { + pub model: String, + pub custom_llm_provider: String, + pub document: OcrDocument, + pub optional_params: Value, +} + +#[derive(Clone, Debug, Serialize)] +pub struct OcrDuringCallRequest { + pub model: String, + pub custom_llm_provider: String, + pub url: String, + pub body: Value, +} + +pub trait OcrHooks: Send + Sync { + fn has_guardrails(&self) -> bool { + false + } + fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> { + Box::pin(async move { Ok(request) }) + } + fn during_call( + &self, + request: OcrDuringCallRequest, + ) -> OcrHookFuture<'_, OcrDuringCallRequest> { + Box::pin(async move { Ok(request) }) + } + fn success<'a>( + &'a self, + _context: &'a CallLifecycleContext, + _response: &'a LiteLLMOcrResponse, + _timing: &'a CallLifecycleTiming, + ) -> OcrLogFuture<'a> { + Box::pin(async {}) + } + fn failure<'a>( + &'a self, + _context: &'a CallLifecycleContext, + _error: &'a Error, + _timing: &'a CallLifecycleTiming, + ) -> OcrLogFuture<'a> { + Box::pin(async {}) + } +} + +pub struct NoopOcrHooks; +impl OcrHooks for NoopOcrHooks {} + +pub(crate) struct OcrLifecycleHooks { + pub hooks: Arc, + pub provider_name: String, +} + +impl CallLifecycleHooks + for OcrLifecycleHooks +{ + type PreCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>; + type DuringCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>; + type SuccessFuture<'a> = OcrLogFuture<'a>; + type FailureFuture<'a> = OcrLogFuture<'a>; + + fn async_pre_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + request: LiteLLMOcrRequest, + ) -> Self::PreCallFuture<'a> { + Box::pin(async move { + if !self.hooks.has_guardrails() { + return Ok(request); + } + let changed = self + .hooks + .pre_call(OcrPreCallRequest { + model: request.model.clone(), + custom_llm_provider: self.provider_name.clone(), + document: request.document, + optional_params: Value::Object(request.optional_params), + }) + .await?; + let Value::Object(optional_params) = changed.optional_params else { + return Err(super::error::OcrRequestError::RequestField { + path: "guardrail.optional_params".into(), + } + .into()); + }; + Ok(LiteLLMOcrRequest { + document: changed.document, + optional_params, + ..request + }) + }) + } + + fn async_during_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + request: LiteLLMOcrRequest, + ) -> Self::DuringCallFuture<'a> { + Box::pin(async move { Ok(request) }) + } + + #[tracing::instrument( + name = "success_callback", + target = "litellm::function_trace", + level = "trace", + skip_all + )] + fn async_log_success_event<'a>( + &'a self, + context: &'a CallLifecycleContext, + response: &'a LiteLLMOcrResponse, + timing: &'a CallLifecycleTiming, + ) -> Self::SuccessFuture<'a> { + self.hooks.success(context, response, timing) + } + + #[tracing::instrument( + name = "failure_callback", + target = "litellm::function_trace", + level = "trace", + skip_all + )] + fn async_log_failure_event<'a>( + &'a self, + context: &'a CallLifecycleContext, + error: &'a Error, + timing: &'a CallLifecycleTiming, + ) -> Self::FailureFuture<'a> { + self.hooks.failure(context, error, timing) + } +} diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index ec2fbb969a6..78c86aeaa0e 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -1,2 +1,21 @@ +mod adapters; +pub mod client; +mod codecs; +pub mod error; +mod handler; +pub mod hooks; +mod prepare; +mod registry; pub mod transformation; pub mod types; +pub mod wire; + +pub use client::{OcrClient, ocr}; +pub use types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument}; + +#[cfg(test)] +#[path = "../../tests/ocr/support.rs"] +pub(crate) mod test_support; +#[cfg(test)] +#[path = "../../tests/ocr.rs"] +pub(crate) mod tests; diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs new file mode 100644 index 00000000000..cff40b7de50 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -0,0 +1,142 @@ +use serde::{Deserialize, Serialize, de::DeserializeOwned}; +use serde_json::{Map, Value}; + +use super::OcrClient; +use super::error::{OcrError, OcrRequestError}; +use super::hooks::OcrDuringCallRequest; +use super::types::LiteLLMOcrRequest; + +#[derive(Debug, Deserialize)] +pub(crate) struct ParsedProviderParams { + #[serde(flatten)] + pub known: T, + #[serde(default, flatten)] + pub extra_params: Map, +} + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub(crate) fn _prepare_ocr_request( + request: &LiteLLMOcrRequest, +) -> Result, OcrRequestError> { + super::wire::decode_request_value( + Value::Object(request.optional_params.clone()), + "optional_params", + ) +} + +pub(crate) async fn transform_request_body( + client: &OcrClient, + request: &LiteLLMOcrRequest, + url: &str, + headers: &[(String, String)], + body: B, + validate: impl FnOnce(&B) -> Result<(), OcrRequestError>, +) -> Result +where + B: Serialize + DeserializeOwned, +{ + let body = if request.hooks.has_guardrails() { + let changed = request + .hooks + .during_call(OcrDuringCallRequest { + model: request.model.clone(), + custom_llm_provider: request.adapter.provider().as_str().into(), + url: url.into(), + body: serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField { + path: "body".into(), + })?, + }) + .await?; + let body = OcrWireBody::::decode(changed.body)?; + validate(&body.body)?; + body + } else { + OcrWireBody { + body, + extra: Map::new(), + } + }; + build_http_request(client, request, url, headers, &body) +} + +pub(crate) fn build_http_request( + client: &OcrClient, + request: &LiteLLMOcrRequest, + url: &str, + headers: &[(String, String)], + body: &B, +) -> Result { + let builder = client + .provider_http() + .post(url) + .json(body) + .timeout(request.connection.timeout); + crate::http_utils::with_headers(builder, headers, crate::http_utils::HeaderPolicy::All) + .build() + .map_err(crate::error::TransportError::from) + .map_err(OcrError::from) +} + +#[derive(Serialize)] +struct OcrWireBody { + #[serde(flatten)] + body: B, + #[serde(flatten)] + extra: Map, +} + +impl OcrWireBody { + fn decode(value: Value) -> Result { + let body: B = super::wire::decode_request_value(value.clone(), "guardrail.body")?; + let Value::Object(fields) = value else { + return Err(OcrRequestError::RequestField { + path: "guardrail.body".into(), + }); + }; + let known = serde_json::to_value(&body).map_err(|_| OcrRequestError::RequestField { + path: "guardrail.body".into(), + })?; + let extra = fields + .into_iter() + .filter(|(key, _)| known.get(key).is_none()) + .collect(); + Ok(Self { body, extra }) + } +} + +pub(crate) fn credential_env(name: &str) -> Option { + std::env::var(name).ok() +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::*; + + #[derive(Debug, Deserialize, PartialEq)] + struct KnownParams { + pages: Option>, + } + + #[test] + fn parsed_provider_params_separates_known_and_extra_params() { + let parsed: ParsedProviderParams = super::super::wire::decode_request_value( + json!({ + "pages": [0, 2], + "future_ocr_option": true, + "extra_body": {"provider_option": "value"} + }), + "optional_params", + ) + .unwrap(); + + assert_eq!(parsed.known.pages, Some(vec![0, 2])); + assert_eq!(parsed.extra_params["future_ocr_option"], true); + assert_eq!( + parsed.extra_params["extra_body"], + json!({"provider_option": "value"}) + ); + assert_eq!(parsed.extra_params.len(), 2); + } +} diff --git a/litellm-rust/crates/core/src/ocr/registry.rs b/litellm-rust/crates/core/src/ocr/registry.rs new file mode 100644 index 00000000000..9bae8153ee8 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/registry.rs @@ -0,0 +1,53 @@ +use super::adapters::OcrAdapter; +use crate::Error; +use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; + +macro_rules! define_adapter_types { + ($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => { + #[derive(Clone, Copy, Debug, PartialEq, Eq)] + pub(crate) enum OcrAdapterKind { + $( $variant, )+ + } + + impl OcrAdapterKind { + pub(crate) const fn provider(self) -> OcrProvider { + match self { + $( Self::$variant => <$adapter>::PROVIDER, )+ + } + } + } + }; +} + +super::adapters::for_each_ocr_adapter!(define_adapter_types); + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum OcrProvider { + Mistral, +} + +impl OcrProvider { + pub(crate) const fn as_str(self) -> &'static str { + match self { + Self::Mistral => "mistral", + } + } +} + +pub(crate) fn resolve_wire_adapter( + model: &str, + custom_llm_provider: Option<&str>, +) -> Result<(String, OcrAdapterKind), Error> { + let provider = + get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider { + model, + custom_llm_provider: OcrProvider::Mistral.as_str(), + }); + let typed_provider = match provider.custom_llm_provider { + "mistral" => OcrProvider::Mistral, + value => return Err(Error::InvalidProvider(value.to_string())), + }; + match typed_provider { + OcrProvider::Mistral => Ok((provider.model.to_string(), OcrAdapterKind::Mistral)), + } +} diff --git a/litellm-rust/crates/core/src/ocr/transformation.rs b/litellm-rust/crates/core/src/ocr/transformation.rs index 62299faf9ed..ac4f10bf15b 100644 --- a/litellm-rust/crates/core/src/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/ocr/transformation.rs @@ -1,7 +1,7 @@ use crate::Error; use serde_json::{Map, Value}; -use super::types::{OcrRequestData, OcrResponseData}; +use super::types::{LiteLLMOcrResponse, OcrRequestData}; #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum OcrAuthStrategy { @@ -49,14 +49,14 @@ pub trait OcrProviderConfig: Sync { &self, model: &str, response_json: Value, - ) -> Result; + ) -> Result; fn transform_ocr_response_with_params( &self, model: &str, response_json: Value, _optional_params: &Map, - ) -> Result { + ) -> Result { self.transform_ocr_response(model, response_json) } diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 71cdb232a87..02b5330188c 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -1,6 +1,14 @@ +use std::sync::Arc; +use std::time::Duration; + use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use super::hooks::{NoopOcrHooks, OcrHooks}; +use super::registry::{OcrAdapterKind, resolve_wire_adapter}; +use crate::Error; +use crate::constants::OCR_HTTP_TIMEOUT_SECS; + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct OcrRequestData { pub data: Value, @@ -8,31 +16,145 @@ pub struct OcrRequestData { } #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct OcrResponseData { +#[serde(tag = "type")] +pub enum OcrDocument { + #[serde(rename = "document_url")] + DocumentUrl { + document_url: String, + #[serde(flatten)] + extra_fields: Map, + }, + #[serde(rename = "image_url")] + ImageUrl { + image_url: String, + #[serde(flatten)] + extra_fields: Map, + }, +} + +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum OcrResponseFormat { + #[default] + Litellm, + Native, +} + +#[derive(Clone)] +pub struct OcrConnection { + pub api_key: Option, + pub api_base: Option, + pub extra_headers: Vec<(String, String)>, + pub timeout: Duration, +} + +impl Default for OcrConnection { + fn default() -> Self { + Self { + api_key: None, + api_base: None, + extra_headers: Vec::new(), + timeout: Duration::from_secs(OCR_HTTP_TIMEOUT_SECS), + } + } +} + +pub struct LiteLLMOcrRequest { + pub model: String, + pub document: OcrDocument, + pub connection: OcrConnection, + pub hooks: Arc, + pub litellm_call_id: Option, + pub optional_params: Map, + pub(crate) adapter: OcrAdapterKind, +} + +impl LiteLLMOcrRequest { + pub fn new( + model: String, + document: OcrDocument, + custom_llm_provider: Option<&str>, + optional_params: Map, + ) -> Result { + let (model, adapter_kind) = resolve_wire_adapter(&model, custom_llm_provider)?; + + Ok(Self { + model, + document, + connection: OcrConnection::default(), + hooks: Arc::new(NoopOcrHooks), + litellm_call_id: None, + optional_params, + adapter: adapter_kind, + }) + } + + pub(crate) fn response_format( + &self, + ) -> Result { + self.optional_params + .get("req_format") + .map(|value| { + serde_json::from_value(value.clone()) + .map_err(|_| super::error::OcrRequestError::RequestFormat) + }) + .transpose() + .map(|format| format.unwrap_or_default()) + } + + pub fn with_host_hooks( + self, + hooks: Arc, + litellm_call_id: Option, + ) -> Self { + Self { + hooks, + litellm_call_id, + ..self + } + } +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct LiteLLMOcrResponse { pub pages: Vec, pub model: String, pub document_annotation: Option, pub usage_info: Option, pub object: String, + #[serde(flatten)] pub extra_fields: Map, + #[serde(skip_serializing_if = "Option::is_none")] pub provider_native_response: Option, } -impl OcrResponseData { +impl LiteLLMOcrResponse { pub fn into_json(self) -> Value { - let mut response = serde_json::json!({ - "pages": self.pages, - "model": self.model, - "document_annotation": self.document_annotation, - "usage_info": self.usage_info, - "object": self.object, - }); - if let Value::Object(object) = &mut response { - object.extend(self.extra_fields); - if let Some(native_response) = self.provider_native_response { - object.insert("provider_native_response".to_string(), native_response); - } - } - response + serde_json::to_value(self).expect("OCR response fields are JSON-compatible") + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn response_serialization_flattens_extra_fields_and_omits_absent_native_response() { + let response = LiteLLMOcrResponse { + pages: vec![], + model: "model".into(), + document_annotation: None, + usage_info: None, + object: "ocr".into(), + extra_fields: json!({"provider_field":"kept"}) + .as_object() + .unwrap() + .clone(), + provider_native_response: None, + }; + let serialized = response.into_json(); + assert_eq!(serialized["provider_field"], "kept"); + assert!(serialized.get("provider_native_response").is_none()); } } diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs new file mode 100644 index 00000000000..f37c06a01ac --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -0,0 +1,154 @@ +use crate::ocr::error::OcrRequestError; +use crate::ocr::error::OcrResponseError; +use std::time::Duration; + +use super::hooks::{OcrDuringCallRequest, OcrPreCallRequest}; +use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument}; +use crate::Error; +use serde::{ + Deserialize, + de::{DeserializeOwned, IntoDeserializer}, +}; +use serde_json::{Map, Value}; + +#[derive(Debug)] +pub struct DecodedOcrResponse { + pub data: T, + pub native: Option, +} + +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +pub struct OcrWireRequest { + pub model: String, + pub document: Value, + pub api_key: Option, + pub api_base: Option, + pub custom_llm_provider: Option, + pub extra_headers: Option>, + #[serde(default)] + pub optional_params: Map, + pub timeout_seconds: Option, +} + +pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> bool { + super::registry::resolve_wire_adapter(model, custom_llm_provider).is_ok() +} + +pub fn decode_request(wire: OcrWireRequest) -> Result { + let document = decode_request_value(wire.document, "document")?; + let headers = wire + .extra_headers + .unwrap_or_default() + .into_iter() + .map(|(name, value)| { + let value = value + .as_str() + .ok_or_else(|| OcrRequestError::RequestField { + path: format!("extra_headers.{name}"), + })?; + Ok((name, value.to_string())) + }) + .collect::, OcrRequestError>>()?; + let timeout = wire + .timeout_seconds + .map(|seconds| { + Duration::try_from_secs_f64(seconds).map_err(|_| OcrRequestError::RequestField { + path: "timeout_seconds".into(), + }) + }) + .transpose()?; + let defaults = OcrConnection::default(); + let request = LiteLLMOcrRequest::new( + wire.model, + document, + wire.custom_llm_provider.as_deref(), + wire.optional_params, + )?; + let connection = OcrConnection { + api_key: nonblank(wire.api_key), + api_base: nonblank(wire.api_base), + extra_headers: headers, + timeout: timeout.unwrap_or(defaults.timeout), + }; + Ok(LiteLLMOcrRequest { + connection, + ..request + }) +} + +fn nonblank(value: Option) -> Option { + value + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()) +} +pub fn decode_request_value( + value: Value, + prefix: &str, +) -> Result { + serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| { + OcrRequestError::RequestField { + path: format!("{prefix}.{}", error.path()), + } + }) +} + +pub fn decode_response( + bytes: &[u8], + native: bool, +) -> Result, OcrResponseError> { + let mut deserializer = serde_json::Deserializer::from_slice(bytes); + let data = serde_path_to_error::deserialize(&mut deserializer).map_err(|error| { + OcrResponseError::ResponseField { + path: error.path().to_string(), + } + })?; + deserializer + .end() + .map_err(|_| OcrResponseError::ResponseField { + path: "response".into(), + })?; + let native = if native { + Some( + serde_json::from_slice(bytes).map_err(|_| OcrResponseError::ResponseField { + path: "response".into(), + })?, + ) + } else { + None + }; + Ok(DecodedOcrResponse { data, native }) +} + +pub fn decode_pre_call_result( + original: OcrPreCallRequest, + value: Value, +) -> Result { + #[derive(Deserialize)] + struct Changed { + document: OcrDocument, + #[serde(default)] + optional_params: Map, + } + let changed: Changed = decode_request_value(value, "guardrail")?; + Ok(OcrPreCallRequest { + document: changed.document, + optional_params: Value::Object(changed.optional_params), + ..original + }) +} + +pub fn decode_during_call_result( + original: OcrDuringCallRequest, + value: Value, +) -> Result { + #[derive(Deserialize)] + struct Changed { + body: Value, + } + let changed: Changed = decode_request_value(value, "guardrail")?; + Ok(OcrDuringCallRequest { + body: changed.body, + ..original + }) +} diff --git a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs index d15c032f0bc..2af25a5639a 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs @@ -2,7 +2,7 @@ use std::collections::BTreeSet; use crate::error::{Error, json_type_name}; use crate::ocr::transformation::{OcrAuthStrategy, OcrProviderConfig, OcrResponseHandling}; -use crate::ocr::types::{OcrRequestData, OcrResponseData}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData}; use serde_json::{Map, Value, json}; use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; @@ -440,7 +440,7 @@ fn transform_document_intelligence_response( model: &str, response_json: Value, preserve_native_response: bool, -) -> Result { +) -> Result { let response = response_json .as_object() .ok_or_else(|| Error::InvalidType { @@ -488,7 +488,7 @@ fn transform_document_intelligence_response( }) .collect(); - Ok(OcrResponseData { + Ok(LiteLLMOcrResponse { usage_info: Some(json!({ "pages_processed": pages.len(), "doc_size_bytes": null, @@ -521,7 +521,7 @@ impl OcrProviderConfig for AzureAiOcrConfig { &self, model: &str, response_json: Value, - ) -> Result { + ) -> Result { MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json) } @@ -599,7 +599,7 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig { &self, model: &str, response_json: Value, - ) -> Result { + ) -> Result { transform_document_intelligence_response(model, response_json, false) } @@ -608,7 +608,7 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig { model: &str, response_json: Value, optional_params: &Map, - ) -> Result { + ) -> Result { transform_document_intelligence_response( model, response_json, @@ -718,7 +718,7 @@ mod tests { }) } - fn assert_native_fields_preserved(response: &OcrResponseData, operation: &Value) { + fn assert_native_fields_preserved(response: &LiteLLMOcrResponse, operation: &Value) { let analyze_result = &operation["analyzeResult"]; assert_eq!(response.extra_fields["content"], analyze_result["content"]); diff --git a/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs index 11e8fe7db18..044fc587c22 100644 --- a/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs @@ -1,6 +1,6 @@ use crate::error::{Error, json_type_name}; use crate::ocr::transformation::OcrProviderConfig; -use crate::ocr::types::{OcrRequestData, OcrResponseData}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData}; use serde_json::{Map, Value}; const SUPPORTED_OCR_PARAMS: &[&str] = &[ @@ -107,7 +107,7 @@ impl OcrProviderConfig for MistralOcrConfig { &self, model: &str, response_json: Value, - ) -> Result { + ) -> Result { let response_object = response_json .as_object() .ok_or_else(|| Error::InvalidType { @@ -128,7 +128,7 @@ impl OcrProviderConfig for MistralOcrConfig { let document_annotation = response_object.get("document_annotation").cloned(); let usage_info = response_object.get("usage_info").cloned(); - Ok(OcrResponseData { + Ok(LiteLLMOcrResponse { pages, model, document_annotation, @@ -178,7 +178,10 @@ pub fn transform_ocr_request( } #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub fn transform_ocr_response(model: &str, response_json: Value) -> Result { +pub fn transform_ocr_response( + model: &str, + response_json: Value, +) -> Result { MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json) } diff --git a/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs index b8507541e18..f025887846f 100644 --- a/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs @@ -6,7 +6,7 @@ use serde_json::{Map, Value, json}; use crate::error::{Error, json_type_name}; use crate::ocr::transformation::OcrProviderConfig; -use crate::ocr::types::{OcrRequestData, OcrResponseData}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData}; pub const REDUCTO_API_BASE: &str = "https://platform.reducto.ai"; pub const REDUCTO_API_KEY_ENV: &str = "REDUCTO_API_KEY"; @@ -283,7 +283,7 @@ fn build_pages(result: &Map) -> Vec { pub fn transform_reducto_response( model: &str, response_json: Value, -) -> Result { +) -> Result { let response = response_json .as_object() .ok_or_else(|| Error::InvalidType { @@ -311,7 +311,7 @@ pub fn transform_reducto_response( "credits": usage.get("credits").cloned().unwrap_or(Value::Null), })); - Ok(OcrResponseData { + Ok(LiteLLMOcrResponse { pages: build_pages(result), model: model.to_string(), document_annotation: None, @@ -341,7 +341,7 @@ impl OcrProviderConfig for ReductoParseV3Config { &self, model: &str, response_json: Value, - ) -> Result { + ) -> Result { transform_reducto_response(model, response_json) } @@ -383,7 +383,7 @@ impl OcrProviderConfig for ReductoParseLegacyConfig { &self, model: &str, response_json: Value, - ) -> Result { + ) -> Result { transform_reducto_response(model, response_json) } diff --git a/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs index c324de8cb45..b0e5a278d0b 100644 --- a/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs @@ -1,6 +1,6 @@ use crate::error::{Error, json_type_name}; use crate::ocr::transformation::OcrProviderConfig; -use crate::ocr::types::{OcrRequestData, OcrResponseData}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData}; use serde_json::{Map, Value, json}; use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; @@ -226,7 +226,7 @@ impl OcrProviderConfig for VertexAiOcrConfig { &self, model: &str, response_json: Value, - ) -> Result { + ) -> Result { MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json) } @@ -301,7 +301,7 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig { &self, model: &str, response_json: Value, - ) -> Result { + ) -> Result { let response = response_json .as_object() .ok_or_else(|| Error::InvalidType { @@ -339,7 +339,7 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig { .get("usage_info") .cloned() .or_else(|| response.get("usage").cloned()); - Ok(OcrResponseData { + Ok(LiteLLMOcrResponse { pages, model: object .get("model") diff --git a/litellm-rust/crates/core/src/url_utils.rs b/litellm-rust/crates/core/src/url_utils.rs new file mode 100644 index 00000000000..982dca0dbe3 --- /dev/null +++ b/litellm-rust/crates/core/src/url_utils.rs @@ -0,0 +1,95 @@ +use std::marker::PhantomData; + +use thiserror::Error; +use url::Url; + +#[derive(Debug, Error)] +pub(crate) enum ApiUrlError { + #[error("invalid URL: {0}")] + Parse(#[from] url::ParseError), + #[error("URL cannot be used as a base")] + CannotBeBase, +} + +pub(crate) struct Base; +pub(crate) struct Complete; + +pub(crate) struct ApiUrl { + url: Url, + state: PhantomData, +} + +impl ApiUrl { + pub(crate) fn parse(value: &str) -> Result { + Ok(Self { + url: Url::parse(value.trim())?, + state: PhantomData, + }) + } + + pub(crate) fn complete_path( + mut self, + target: &[&str], + ) -> Result, ApiUrlError> { + let existing: Vec = self + .url + .path_segments() + .ok_or(ApiUrlError::CannotBeBase)? + .filter(|segment| !segment.is_empty()) + .map(str::to_string) + .collect(); + let overlap = (0..=existing.len().min(target.len())) + .rev() + .find(|&length| { + existing[existing.len() - length..] + .iter() + .map(String::as_str) + .eq(target[..length].iter().copied()) + }) + .unwrap_or(0); + self.url + .path_segments_mut() + .map_err(|()| ApiUrlError::CannotBeBase)? + .pop_if_empty() + .extend(target[overlap..].iter().copied()); + Ok(ApiUrl { + url: self.url, + state: PhantomData, + }) + } +} + +impl ApiUrl { + pub(crate) fn into_string(self) -> String { + self.url.into() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn completion_appends_only_the_missing_path_suffix() { + for (base, expected) in [ + ("https://example.test", "https://example.test/v1/ocr"), + ("https://example.test/v1", "https://example.test/v1/ocr"), + ("https://example.test/v1/ocr", "https://example.test/v1/ocr"), + ] { + let actual = ApiUrl::parse(base) + .and_then(|url| url.complete_path(&["v1", "ocr"])) + .map(|url| url.into_string()) + .expect("url builds"); + assert_eq!(actual, expected); + } + } + + #[test] + fn completion_places_paths_before_queries() { + let actual = ApiUrl::parse("https://example.test/v1?tenant=a") + .and_then(|url| url.complete_path(&["v1", "ocr"])) + .map(|url| url.into_string()) + .expect("url builds"); + assert_eq!(actual, "https://example.test/v1/ocr?tenant=a"); + } +} diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs new file mode 100644 index 00000000000..d1828dfb816 --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -0,0 +1,227 @@ +use std::sync::{Arc, Mutex}; + +use serde_json::{Value, json}; + +use super::OcrClient; +use super::hooks::{OcrHookFuture, OcrHooks, OcrLogFuture, OcrPreCallRequest}; +use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; +use super::wire::{OcrWireRequest, decode_request}; +use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming}; + +#[test] +fn request_boundary_selects_mistral_and_rejects_unknown_providers() { + let request = OcrWireRequest { + model: "mistral/model".into(), + document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}), + api_key: Some("key".into()), + api_base: None, + custom_llm_provider: None, + extra_headers: None, + optional_params: json!({"extract_header":true,"unknown":42}) + .as_object() + .unwrap() + .clone(), + timeout_seconds: None, + }; + assert!(decode_request(request).is_ok()); + assert!( + decode_request(OcrWireRequest { + model: "model".into(), + document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}), + api_key: Some("key".into()), + api_base: None, + custom_llm_provider: Some("unknown".into()), + extra_headers: None, + optional_params: serde_json::Map::new(), + timeout_seconds: None, + }) + .is_err() + ); +} + +#[tokio::test] +async fn facade_executes_direct_mistral_once() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ + "pages":[{"index":0,"markdown":"hello","custom":"preserved"}], + "usage_info":{"pages_processed":1} + }))]) + .await; + let result = perform_ocr(wire_request( + "mistral/model", + &base, + json!({"extract_header":true,"unknown":"ignored"}), + )) + .await + .unwrap(); + server.await.unwrap(); + assert_eq!(result.pages[0]["markdown"], "hello"); + assert_eq!(result.pages[0]["custom"], "preserved"); + let requests = seen.lock().unwrap(); + assert_eq!(requests.len(), 1); + assert!(requests[0].starts_with("POST /v1/ocr ")); + assert!( + requests[0] + .to_ascii_lowercase() + .contains("authorization: bearer test-key\r\n") + ); + let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap(); + assert_eq!( + body, + json!({ + "model":"model", + "document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}, + "extract_header":true + }) + ); +} + +#[tokio::test] +async fn facade_retains_native_response_when_requested() { + let provider_response = json!({ + "pages":[{"index":0,"markdown":"hello"}], + "usage_info":{"pages_processed":1}, + "provider_only":"preserved" + }); + let (base, _, server) = mock_server(vec![MockResponse::json(provider_response.clone())]).await; + let response = perform_ocr(wire_request( + "mistral/model", + &base, + json!({"req_format":"native"}), + )) + .await + .unwrap(); + + server.await.unwrap(); + assert_eq!(response.provider_native_response, Some(provider_response)); +} + +#[tokio::test] +async fn facade_uses_the_injected_http_client() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; + let mut default_headers = reqwest::header::HeaderMap::new(); + default_headers.insert( + "x-transport-owner", + reqwest::header::HeaderValue::from_static("host"), + ); + let provider_http = reqwest::Client::builder() + .default_headers(default_headers) + .build() + .unwrap(); + OcrClient::new(provider_http) + .unwrap() + .perform(wire_request("mistral/model", &base, json!({}))) + .await + .unwrap(); + server.await.unwrap(); + assert!(seen.lock().unwrap()[0].contains("x-transport-owner: host")); +} + +struct RecordingHooks { + events: Arc>>, + block: bool, +} + +impl OcrHooks for RecordingHooks { + fn has_guardrails(&self) -> bool { + true + } + + fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> { + Box::pin(async move { + self.events.lock().unwrap().push("pre"); + if self.block { + return Err(crate::Error::InvalidRequest("blocked".into())); + } + Ok(request) + }) + } + + fn during_call( + &self, + request: super::hooks::OcrDuringCallRequest, + ) -> OcrHookFuture<'_, super::hooks::OcrDuringCallRequest> { + Box::pin(async move { + self.events.lock().unwrap().push("during"); + Ok(request) + }) + } + + fn success<'a>( + &'a self, + _context: &'a CallLifecycleContext, + _response: &'a super::LiteLLMOcrResponse, + _timing: &'a CallLifecycleTiming, + ) -> OcrLogFuture<'a> { + Box::pin(async move { + self.events.lock().unwrap().push("success"); + }) + } + + fn failure<'a>( + &'a self, + _context: &'a CallLifecycleContext, + _error: &'a crate::Error, + _timing: &'a CallLifecycleTiming, + ) -> OcrLogFuture<'a> { + Box::pin(async move { + self.events.lock().unwrap().push("failure"); + }) + } +} + +#[tokio::test] +async fn lifecycle_orders_hooks_and_emits_one_success() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; + let events = Arc::new(Mutex::new(Vec::new())); + let request = wire_request("mistral/model", &base, json!({})); + let request = super::LiteLLMOcrRequest { + hooks: Arc::new(RecordingHooks { + events: events.clone(), + block: false, + }), + ..request + }; + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + assert_eq!(*events.lock().unwrap(), ["pre", "during", "success"]); + assert_eq!(seen.lock().unwrap().len(), 1); +} + +#[tokio::test] +async fn lifecycle_blocking_prevents_execution_and_emits_one_failure() { + let events = Arc::new(Mutex::new(Vec::new())); + let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({})); + let request = super::LiteLLMOcrRequest { + hooks: Arc::new(RecordingHooks { + events: events.clone(), + block: true, + }), + ..request + }; + let error = perform_ocr(request).await.unwrap_err(); + assert!(matches!(error, crate::Error::InvalidRequest(_))); + assert_eq!(*events.lock().unwrap(), ["pre", "failure"]); +} + +#[tokio::test] +async fn upstream_failure_emits_one_terminal_failure() { + let (base, seen, server) = mock_server(vec![MockResponse { + status: 500, + headers: vec![], + body: json!({"error":"failed"}), + }]) + .await; + let events = Arc::new(Mutex::new(Vec::new())); + let request = wire_request("mistral/model", &base, json!({})); + let request = super::LiteLLMOcrRequest { + hooks: Arc::new(RecordingHooks { + events: events.clone(), + block: false, + }), + ..request + }; + assert!(perform_ocr(request).await.is_err()); + server.await.unwrap(); + assert_eq!(*events.lock().unwrap(), ["pre", "during", "failure"]); + assert_eq!(seen.lock().unwrap().len(), 1); +} diff --git a/litellm-rust/crates/core/tests/ocr/support.rs b/litellm-rust/crates/core/tests/ocr/support.rs new file mode 100644 index 00000000000..d9720019419 --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/support.rs @@ -0,0 +1,106 @@ +use std::sync::{Arc, Mutex}; + +use serde_json::{Value, json}; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::TcpListener; + +use crate::ocr::wire::{OcrWireRequest, decode_request}; +use crate::ocr::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient}; + +pub(crate) fn ocr_client() -> OcrClient { + OcrClient::for_test(reqwest::Client::new()) +} + +pub(crate) async fn perform_ocr( + request: LiteLLMOcrRequest, +) -> Result { + ocr_client().perform(request).await +} + +pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest { + decode_request(OcrWireRequest { + model: model.into(), + document: json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}), + api_key: Some("test-key".into()), + api_base: Some(base.into()), + custom_llm_provider: None, + extra_headers: None, + optional_params: options.as_object().unwrap().clone(), + timeout_seconds: Some(2.0), + }) + .unwrap() +} + +pub(crate) struct MockResponse { + pub status: u16, + pub headers: Vec<(&'static str, String)>, + pub body: Value, +} + +impl MockResponse { + pub fn json(body: Value) -> Self { + Self { + status: 200, + headers: vec![], + body, + } + } +} + +pub(crate) async fn mock_server( + responses: Vec, +) -> (String, Arc>>, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let base = format!("http://{}", listener.local_addr().unwrap()); + let requests = Arc::new(Mutex::new(Vec::new())); + let seen = requests.clone(); + let server_base = base.clone(); + let task = tokio::spawn(async move { + for response in responses { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut bytes = Vec::new(); + let mut buffer = [0u8; 4096]; + let header_end = loop { + let n = socket.read(&mut buffer).await.unwrap(); + assert!(n > 0); + bytes.extend_from_slice(&buffer[..n]); + if let Some(index) = bytes.windows(4).position(|s| s == b"\r\n\r\n") { + break index + 4; + } + }; + let length = String::from_utf8_lossy(&bytes[..header_end]) + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().unwrap()) + }) + .unwrap_or(0); + while bytes.len() < header_end + length { + let n = socket.read(&mut buffer).await.unwrap(); + assert!(n > 0); + bytes.extend_from_slice(&buffer[..n]); + } + seen.lock() + .unwrap() + .push(String::from_utf8_lossy(&bytes).into_owned()); + let body = serde_json::to_vec(&response.body).unwrap(); + let headers = response + .headers + .into_iter() + .map(|(name, value)| { + format!("{name}: {}\r\n", value.replace("{base}", &server_base)) + }) + .collect::(); + let head = format!( + "HTTP/1.1 {} OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n{}\r\n", + response.status, + body.len(), + headers + ); + socket.write_all(head.as_bytes()).await.unwrap(); + socket.write_all(&body).await.unwrap(); + } + }); + (base, requests, task) +} diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 76c298abf89..407475dedce 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -41,6 +41,7 @@ pub(crate) fn chat_completions_error_to_pyerr(err: Error) -> PyErr { | Error::InvalidRequest(_) | Error::InvalidType { .. } | Error::MissingField(_) + | Error::MissingApiKey { .. } | Error::Routing(_) // Nothing reached the provider, so serving it on Python cannot double // bill and is the only way the caller gets an answer at all. From 93bd3d61b56a167b794c6b8cc2dcc6dcefaf4a47 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Thu, 10 Sep 2026 13:20:49 -0700 Subject: [PATCH 49/60] fix(mcp): preserve verified identity when warming OAuth cache --- .../mcp_server/discoverable_endpoints.py | 1 + .../mcp_server/test_discoverable_endpoints.py | 43 +++++++++++++++++++ 2 files changed, 44 insertions(+) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index c8e02be308b..920fe17e43e 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -634,6 +634,7 @@ async def _store_per_user_token_server_side( server_id=server.server_id, access_token=access_token, ttl=ttl, + identity_binding_proof=identity_binding_proof, ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 42187cb1771..ff649075468 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -11103,3 +11103,46 @@ async def test_identity_bound_authorization_carries_nonce_and_caller_through_cal assert sealed.upstream_code == "upstream-code" assert sealed.oauth_nonce == query["nonce"][0] assert returned["state"] == ["client-state"] + + +@pytest.mark.asyncio +async def test_enforced_login_warms_verified_token_readable_without_database_lookup(monkeypatch): + from types import SimpleNamespace + from litellm.caching.caching import DualCache + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import db, mcp_server_manager + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _store_per_user_token_server_side + from litellm.proxy._experimental.mcp_server.oauth2_token_cache import mcp_per_user_token_cache + from litellm.proxy._experimental.mcp_server.oauth_identity_binding import current_binding_proof + from litellm.proxy._experimental.mcp_server.outbound_credentials.dual_cache_token_backend import DualCacheTokenCacheBackend + from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import CachedOAuthTokenStore + from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import LazyPerUserOAuthTokenStore + from litellm.proxy._experimental.mcp_server.outbound_credentials.v2_token_store import V2PerUserTokenStore + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="srv", name="srv", transport="http", auth_type="oauth2", + oauth_identity_binding={"mode": "enforce", "issuer": "https://idp.example", "audiences": ["client"], + "caller_field": "user_id", "principal_claim": "sub"}, + ) + proof = await current_binding_proof(server.oauth_identity_binding, "alice", "srv") + cache = DualCache() + monkeypatch.setenv("LITELLM_SALT_KEY", "test-cache-warm-encryption-salt") + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", SimpleNamespace(invalidate_user_oauth_token_cache=AsyncMock())) + monkeypatch.setattr(db, "store_user_oauth_credential", AsyncMock()) + await _store_per_user_token_server_side( + server=server, user_id="alice", token_response={"access_token": "alice-token", "refresh_token": "private", "expires_in": 3600}, + identity_binding_proof=proof, + ) + read = AsyncMock(return_value=None) + cached = CachedOAuthTokenStore( + V2PerUserTokenStore(read), default_ttl_seconds=300, backend=DualCacheTokenCacheBackend(cache, mcp_per_user_token_cache._codec()), + ) + store = LazyPerUserOAuthTokenStore(lambda _: server, store_builder=lambda _: (cached, True), redis_available=lambda: True) + token = await store.fetch("alice", "srv") + assert token is not None and token.access_token == "alice-token" + assert token.identity_binding_proof == proof + assert token.refresh_token is None + read.assert_not_awaited() From a9cec50960e76f1e311bbbb679d6b3dd3394c93e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 13:36:23 -0700 Subject: [PATCH 50/60] feat(infra): scale gateway on per-pod RPS and TPS in Helm and Terraform (#40479) * feat(infra): scale gateway on per-pod RPM and TPM in Helm and Terraform Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(helm): require the metrics server before rendering the gateway ServiceMonitor The http port serves /metrics/ behind virtual-key auth, so a ServiceMonitor pointed at it only collects 401s and the RPM/TPM HPA metrics never appear Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(infra): express gateway HPA, KEDA and ECS workload targets per second Rename the per-pod request and token targets in both Helm charts and the AWS module from per minute to per second, and shorten the recommended Prometheus rate window to [1m] with no * 60 so the adapter and KEDA signals are what the HPA compares against. ECS keeps CloudWatch's 60-second aggregation: the ALB target is 60x the per-second variable and the token metric math divides the period Sum by 60 before dividing by the running task count. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- helm/litellm-helm/templates/hpa.yaml | 18 ++ helm/litellm-helm/templates/keda.yaml | 21 ++ helm/litellm-helm/tests/hpa_tests.yaml | 78 +++++++ helm/litellm-helm/tests/keda_tests.yaml | 106 +++++++++ helm/litellm-helm/values.yaml | 36 +++ helm/litellm/templates/gateway/hpa.yaml | 18 ++ .../templates/gateway/servicemonitor.yaml | 28 +++ .../tests/hpa_workload_metrics_tests.yaml | 213 ++++++++++++++++++ helm/litellm/values.yaml | 29 +++ terraform/litellm/aws/README.md | 55 +++++ terraform/litellm/aws/autoscaling.tf | 99 ++++++++ .../aws/tests/workload_autoscaling.tftest.hcl | 194 ++++++++++++++++ terraform/litellm/aws/variables.tf | 41 ++++ terraform/litellm/gcp/README.md | 18 ++ 14 files changed, 954 insertions(+) create mode 100644 helm/litellm-helm/tests/keda_tests.yaml create mode 100644 helm/litellm/templates/gateway/servicemonitor.yaml create mode 100644 helm/litellm/tests/hpa_workload_metrics_tests.yaml create mode 100644 terraform/litellm/aws/tests/workload_autoscaling.tftest.hcl diff --git a/helm/litellm-helm/templates/hpa.yaml b/helm/litellm-helm/templates/hpa.yaml index fec4d1f5c5e..010c4095ab3 100644 --- a/helm/litellm-helm/templates/hpa.yaml +++ b/helm/litellm-helm/templates/hpa.yaml @@ -33,4 +33,22 @@ spec: type: Utilization averageUtilization: {{ .Values.autoscaling.targetMemoryUtilizationPercentage }} {{- end }} + {{- with .Values.autoscaling.targetRequestsPerSecond }} + - type: Pods + pods: + metric: + name: litellm_requests_per_second + target: + type: AverageValue + averageValue: {{ toJson . | trimAll "\"" | quote }} + {{- end }} + {{- with .Values.autoscaling.targetTokensPerSecond }} + - type: Pods + pods: + metric: + name: litellm_tokens_per_second + target: + type: AverageValue + averageValue: {{ toJson . | trimAll "\"" | quote }} + {{- end }} {{- end }} diff --git a/helm/litellm-helm/templates/keda.yaml b/helm/litellm-helm/templates/keda.yaml index fe5190fffc6..bf585d0d4be 100644 --- a/helm/litellm-helm/templates/keda.yaml +++ b/helm/litellm-helm/templates/keda.yaml @@ -23,6 +23,27 @@ spec: triggers: {{- with .Values.keda.triggers }} {{- toYaml . | nindent 2 }} +{{- end }} +{{- $prom := .Values.keda.prometheus }} +{{- if or $prom.requestsPerSecond $prom.tokensPerSecond }} +{{- if not $prom.serverAddress }} +{{- fail "keda.prometheus.serverAddress is required when keda.prometheus.requestsPerSecond or tokensPerSecond is set" }} +{{- end }} +{{- $selector := printf "namespace=%q,job=%q" .Release.Namespace (printf "%s%s" (include "litellm.fullname" .) (ternary "-metrics" "" .Values.metricsServer.enabled)) }} +{{- with $prom.requestsPerSecond }} + - type: prometheus + metadata: + serverAddress: {{ $prom.serverAddress | quote }} + threshold: {{ toJson . | trimAll "\"" | quote }} + query: {{ printf "sum(rate(litellm_proxy_total_requests_metric_total{%s}[1m]))" $selector | quote }} +{{- end }} +{{- with $prom.tokensPerSecond }} + - type: prometheus + metadata: + serverAddress: {{ $prom.serverAddress | quote }} + threshold: {{ toJson . | trimAll "\"" | quote }} + query: {{ printf "sum(rate(litellm_total_tokens_metric_total{%s}[1m]))" $selector | quote }} +{{- end }} {{- end }} advanced: restoreToOriginalReplicaCount: {{ .Values.keda.restoreToOriginalReplicaCount }} diff --git a/helm/litellm-helm/tests/hpa_tests.yaml b/helm/litellm-helm/tests/hpa_tests.yaml index cd062dd5971..e446f58c8fe 100644 --- a/helm/litellm-helm/tests/hpa_tests.yaml +++ b/helm/litellm-helm/tests/hpa_tests.yaml @@ -61,6 +61,84 @@ tests: - equal: { path: "spec.metrics[1].resource.name", value: memory } - equal: { path: "spec.metrics[1].resource.target.averageUtilization", value: 80 } + - it: "renders no workload metrics by default" + set: + autoscaling.enabled: true + autoscaling.targetMemoryUtilizationPercentage: 80 + asserts: + - lengthEqual: { path: spec.metrics, count: 2 } + - notContains: { path: spec.metrics, content: { type: Pods }, any: true } + + - it: "adds a requests-per-second Pods metric after the cpu metric" + set: + autoscaling.enabled: true + autoscaling.targetRequestsPerSecond: 90 + asserts: + - lengthEqual: { path: spec.metrics, count: 2 } + - equal: { path: "spec.metrics[0].resource.name", value: cpu } + - equal: + path: "spec.metrics[1]" + value: + type: Pods + pods: + metric: { name: litellm_requests_per_second } + target: { type: AverageValue, averageValue: "90" } + + - it: "adds a tokens-per-second Pods metric on its own" + set: + autoscaling.enabled: true + autoscaling.targetTokensPerSecond: 6M + asserts: + - lengthEqual: { path: spec.metrics, count: 2 } + - equal: + path: "spec.metrics[1]" + value: + type: Pods + pods: + metric: { name: litellm_tokens_per_second } + target: { type: AverageValue, averageValue: "6M" } + - notContains: + path: spec.metrics + content: { type: Pods, pods: { metric: { name: litellm_requests_per_second } } } + any: true + + - it: "renders requests, tokens, cpu and memory metrics together" + set: + autoscaling.enabled: true + autoscaling.targetMemoryUtilizationPercentage: 80 + autoscaling.targetRequestsPerSecond: 90 + autoscaling.targetTokensPerSecond: 6000000 + asserts: + - lengthEqual: { path: spec.metrics, count: 4 } + - equal: { path: "spec.metrics[0].resource.name", value: cpu } + - equal: { path: "spec.metrics[1].resource.name", value: memory } + - equal: { path: "spec.metrics[2].pods.metric.name", value: litellm_requests_per_second } + - equal: { path: "spec.metrics[2].pods.target.averageValue", value: "90" } + - equal: { path: "spec.metrics[3].pods.metric.name", value: litellm_tokens_per_second } + - equal: { path: "spec.metrics[3].pods.target.averageValue", value: "6000000" } + + - it: "scales on workload metrics alone when the cpu target is cleared" + set: + autoscaling.enabled: true + autoscaling.targetCPUUtilizationPercentage: null + autoscaling.targetRequestsPerSecond: 90 + autoscaling.targetTokensPerSecond: 6000000 + asserts: + - lengthEqual: { path: spec.metrics, count: 2 } + - notContains: { path: spec.metrics, content: { type: Resource }, any: true } + - equal: { path: "spec.metrics[0].pods.metric.name", value: litellm_requests_per_second } + - equal: { path: "spec.metrics[1].pods.metric.name", value: litellm_tokens_per_second } + - notMatchRegexRaw: { pattern: per_minute } + + - it: "ignores the per-minute keys, which the chart never shipped" + set: + autoscaling.enabled: true + autoscaling.targetRequestsPerMinute: 5400 + autoscaling.targetTokensPerMinute: 360000000 + asserts: + - lengthEqual: { path: spec.metrics, count: 1 } + - notContains: { path: spec.metrics, content: { type: Pods }, any: true } + - it: "renders no hpa when autoscaling is disabled" asserts: - hasDocuments: { count: 0 } diff --git a/helm/litellm-helm/tests/keda_tests.yaml b/helm/litellm-helm/tests/keda_tests.yaml new file mode 100644 index 00000000000..c9598646223 --- /dev/null +++ b/helm/litellm-helm/tests/keda_tests.yaml @@ -0,0 +1,106 @@ +suite: "keda" +templates: + - keda.yaml +release: + name: rel + namespace: llm +tests: + - it: "renders no scaled object by default" + asserts: + - hasDocuments: { count: 0 } + + - it: "passes user triggers through and adds no prometheus triggers by default" + set: + keda.enabled: true + keda.triggers: + - type: cpu + metricType: Utilization + metadata: { value: "60" } + asserts: + - isKind: { of: ScaledObject } + - equal: + path: spec.triggers + value: + - type: cpu + metricType: Utilization + metadata: { value: "60" } + + - it: "scales on release-wide requests per second divided by the per-replica target" + set: + keda.enabled: true + keda.prometheus.serverAddress: http://prometheus-operated.monitoring.svc:9090 + keda.prometheus.requestsPerSecond: 90 + asserts: + - lengthEqual: { path: spec.triggers, count: 1 } + - equal: + path: "spec.triggers[0]" + value: + type: prometheus + metadata: + serverAddress: http://prometheus-operated.monitoring.svc:9090 + threshold: "90" + query: sum(rate(litellm_proxy_total_requests_metric_total{namespace="llm",job="rel-litellm"}[1m])) + + - it: "scales on tokens per second on its own" + set: + keda.enabled: true + keda.prometheus.serverAddress: http://prom:9090 + keda.prometheus.tokensPerSecond: 6000000 + asserts: + - lengthEqual: { path: spec.triggers, count: 1 } + - equal: { path: "spec.triggers[0].type", value: prometheus } + - equal: { path: "spec.triggers[0].metadata.threshold", value: "6000000" } + - equal: + path: "spec.triggers[0].metadata.query" + value: sum(rate(litellm_total_tokens_metric_total{namespace="llm",job="rel-litellm"}[1m])) + + - it: "appends requests and tokens triggers after user triggers and selects the metrics service job" + set: + keda.enabled: true + metricsServer.enabled: true + keda.triggers: + - type: cpu + metricType: Utilization + metadata: { value: "60" } + keda.prometheus.serverAddress: http://prom:9090 + keda.prometheus.requestsPerSecond: 90 + keda.prometheus.tokensPerSecond: 6000000 + asserts: + - lengthEqual: { path: spec.triggers, count: 3 } + - equal: { path: "spec.triggers[0].type", value: cpu } + - equal: { path: "spec.triggers[1].metadata.threshold", value: "90" } + - equal: + path: "spec.triggers[1].metadata.query" + value: sum(rate(litellm_proxy_total_requests_metric_total{namespace="llm",job="rel-litellm-metrics"}[1m])) + - equal: { path: "spec.triggers[2].metadata.threshold", value: "6000000" } + - equal: + path: "spec.triggers[2].metadata.query" + value: sum(rate(litellm_total_tokens_metric_total{namespace="llm",job="rel-litellm-metrics"}[1m])) + - notMatchRegexRaw: { pattern: "\\* *60|per_minute|PerMinute" } + + - it: "ignores the per-minute keys, which the chart never shipped" + set: + keda.enabled: true + keda.prometheus.serverAddress: http://prom:9090 + keda.prometheus.requestsPerMinute: 5400 + keda.prometheus.tokensPerMinute: 360000000 + asserts: + - isKind: { of: ScaledObject } + - isNullOrEmpty: { path: spec.triggers } + + - it: "refuses a workload target without a prometheus server address" + set: + keda.enabled: true + keda.prometheus.requestsPerSecond: 90 + asserts: + - failedTemplate: + errorMessage: keda.prometheus.serverAddress is required when keda.prometheus.requestsPerSecond or tokensPerSecond is set + + - it: "yields to the hpa when both autoscalers are enabled" + set: + autoscaling.enabled: true + keda.enabled: true + keda.prometheus.serverAddress: http://prom:9090 + keda.prometheus.requestsPerSecond: 90 + asserts: + - hasDocuments: { count: 0 } diff --git a/helm/litellm-helm/values.yaml b/helm/litellm-helm/values.yaml index 8dc7b967e11..907f85de8fc 100644 --- a/helm/litellm-helm/values.yaml +++ b/helm/litellm-helm/values.yaml @@ -222,6 +222,25 @@ autoscaling: # Memory is a floor to provision under 'resources', not a signal to scale on. # targetMemoryUtilizationPercentage: 80 # behavior: {} + # Opt-in per-pod workload targets, rendered as autoscaling/v2 `Pods` metrics + # named `litellm_requests_per_second` and `litellm_tokens_per_second` with an + # AverageValue target, alongside whichever resource targets are set (the HPA + # follows the metric asking for the most replicas). A Prometheus Adapter must + # serve those two names on custom.metrics.k8s.io from the proxy's counters, + # grouped by the scrape target's `pod` label (enable serviceMonitor below so + # every pod is scraped on its own): + # litellm_requests_per_second: + # sum(rate(litellm_proxy_total_requests_metric_total{<<.LabelMatchers>>}[1m])) by (<<.GroupBy>>) + # litellm_tokens_per_second: + # sum(rate(litellm_total_tokens_metric_total{<<.LabelMatchers>>}[1m])) by (<<.GroupBy>>) + # rate() over [1m] is already per second, so no `* 60`. How fast the HPA + # reacts is set by that window, the scrape interval and the HPA sync period + # (15s by default), not by the unit: keep serviceMonitor.interval at 15s or + # faster so a 1m window holds at least 4 samples. averageValue takes SI + # suffixes, so "6M" is six million tokens per second per pod. Tokens are + # counted when a response completes, so TPS trails long streams. + targetRequestsPerSecond: "" + targetTokensPerSecond: "" # Autoscaling with keda is mutually exclusive with hpa keda: @@ -243,6 +262,23 @@ keda: # metricName: http_requests_total # threshold: '100' # query: sum(rate(http_requests_total{deployment="my-deployment"}[2m])) + # First-class Prometheus triggers on the proxy's own request and token + # counters, appended to `triggers`. Each target is the per-second load one + # replica should carry: KEDA divides the release-wide + # `sum(rate([1m]))` by it to pick the replica count. Thresholds + # are plain numbers (KEDA parses them as floats, no SI suffixes). The + # queries select samples by the release namespace and the `job` label the + # chart's ServiceMonitor produces (the metrics Service name), so enable + # serviceMonitor below together with metricsServer: the http port serves + # /metrics/ behind virtual-key auth and answers an unauthenticated scrape + # with 401. Reaction time comes from the [1m] window, the scrape interval + # and pollingInterval above, so keep both at 15s or faster. Tokens are + # counted at completion, so TPS trails long streams. serverAddress is + # required once either target is set. + prometheus: + serverAddress: "" + requestsPerSecond: "" + tokensPerSecond: "" behavior: {} # scaleDown: # stabilizationWindowSeconds: 300 diff --git a/helm/litellm/templates/gateway/hpa.yaml b/helm/litellm/templates/gateway/hpa.yaml index e97cef95ffb..4183079cd31 100644 --- a/helm/litellm/templates/gateway/hpa.yaml +++ b/helm/litellm/templates/gateway/hpa.yaml @@ -30,6 +30,24 @@ spec: type: Utilization averageUtilization: {{ .Values.gateway.hpa.targetMemoryUtilizationPercentage }} {{- end }} + {{- with .Values.gateway.hpa.targetRequestsPerSecond }} + - type: Pods + pods: + metric: + name: litellm_requests_per_second + target: + type: AverageValue + averageValue: {{ toJson . | trimAll "\"" | quote }} + {{- end }} + {{- with .Values.gateway.hpa.targetTokensPerSecond }} + - type: Pods + pods: + metric: + name: litellm_tokens_per_second + target: + type: AverageValue + averageValue: {{ toJson . | trimAll "\"" | quote }} + {{- end }} {{- with .Values.gateway.hpa.behavior }} behavior: {{- toYaml . | nindent 4 }} diff --git a/helm/litellm/templates/gateway/servicemonitor.yaml b/helm/litellm/templates/gateway/servicemonitor.yaml new file mode 100644 index 00000000000..e1bafa6e388 --- /dev/null +++ b/helm/litellm/templates/gateway/servicemonitor.yaml @@ -0,0 +1,28 @@ +{{- if and .Values.gateway.enabled .Values.gateway.serviceMonitor.enabled }} +{{- if not .Values.gateway.metricsServer.enabled }} +{{- fail "gateway.serviceMonitor.enabled requires gateway.metricsServer.enabled: the http port serves /metrics/ behind virtual-key auth, so an unauthenticated scrape gets 401" }} +{{- end }} +apiVersion: monitoring.coreos.com/v1 +kind: ServiceMonitor +metadata: + name: {{ include "litellm.gateway.fullname" . }} + labels: + {{- include "litellm.commonLabels" . | nindent 4 }} + app.kubernetes.io/component: gateway + {{- with .Values.gateway.serviceMonitor.labels }} + {{- toYaml . | nindent 4 }} + {{- end }} +spec: + selector: + matchLabels: + {{- include "litellm.gateway.selectorLabels" . | nindent 6 }} + namespaceSelector: + matchNames: + - {{ .Release.Namespace | quote }} + endpoints: + - port: metrics + path: /metrics/ + interval: {{ .Values.gateway.serviceMonitor.interval }} + scrapeTimeout: {{ .Values.gateway.serviceMonitor.scrapeTimeout }} + scheme: http +{{- end }} diff --git a/helm/litellm/tests/hpa_workload_metrics_tests.yaml b/helm/litellm/tests/hpa_workload_metrics_tests.yaml new file mode 100644 index 00000000000..a29c53c5ed9 --- /dev/null +++ b/helm/litellm/tests/hpa_workload_metrics_tests.yaml @@ -0,0 +1,213 @@ +suite: test gateway HPA per-pod requests-per-second and tokens-per-second targets +templates: + - gateway/hpa.yaml + - gateway/servicemonitor.yaml +values: + - ./values/required.yaml +tests: + - it: scales on CPU and memory only by default + template: gateway/hpa.yaml + asserts: + - equal: + path: spec.metrics + value: + - type: Resource + resource: + name: cpu + target: + type: Utilization + averageUtilization: 70 + - type: Resource + resource: + name: memory + target: + type: Utilization + averageUtilization: 80 + + - it: adds a requests-per-second Pods metric next to the resource metrics + template: gateway/hpa.yaml + set: + gateway.hpa.targetRequestsPerSecond: 90 + asserts: + - lengthEqual: + path: spec.metrics + count: 3 + - contains: + path: spec.metrics + content: + type: Resource + resource: + name: cpu + target: + type: Utilization + averageUtilization: 70 + - equal: + path: spec.metrics[2] + value: + type: Pods + pods: + metric: + name: litellm_requests_per_second + target: + type: AverageValue + averageValue: "90" + - notContains: + path: spec.metrics + content: + type: Pods + pods: + metric: + name: litellm_tokens_per_second + any: true + + - it: adds a tokens-per-second Pods metric on its own + template: gateway/hpa.yaml + set: + gateway.hpa.targetTokensPerSecond: 6M + asserts: + - lengthEqual: + path: spec.metrics + count: 3 + - equal: + path: spec.metrics[2] + value: + type: Pods + pods: + metric: + name: litellm_tokens_per_second + target: + type: AverageValue + averageValue: "6M" + - notContains: + path: spec.metrics + content: + type: Pods + pods: + metric: + name: litellm_requests_per_second + any: true + + - it: renders requests and tokens targets together and keeps CPU and memory + template: gateway/hpa.yaml + set: + gateway.hpa.targetRequestsPerSecond: 90 + gateway.hpa.targetTokensPerSecond: 6000000 + asserts: + - lengthEqual: + path: spec.metrics + count: 4 + - equal: + path: spec.metrics[0].resource.name + value: cpu + - equal: + path: spec.metrics[1].resource.name + value: memory + - equal: + path: spec.metrics[2].pods.metric.name + value: litellm_requests_per_second + - equal: + path: spec.metrics[3].pods.metric.name + value: litellm_tokens_per_second + - equal: + path: spec.metrics[3].pods.target.averageValue + value: "6000000" + + - it: scales on workload metrics alone when the resource targets are cleared + template: gateway/hpa.yaml + set: + gateway.hpa.targetCPUUtilizationPercentage: null + gateway.hpa.targetMemoryUtilizationPercentage: null + gateway.hpa.targetRequestsPerSecond: 90 + gateway.hpa.targetTokensPerSecond: 6000000 + asserts: + - lengthEqual: + path: spec.metrics + count: 2 + - notContains: + path: spec.metrics + content: + type: Resource + any: true + - equal: + path: spec.metrics[0].pods.metric.name + value: litellm_requests_per_second + - equal: + path: spec.metrics[1].pods.metric.name + value: litellm_tokens_per_second + - notMatchRegexRaw: + pattern: per_minute + + - it: ignores the per-minute keys, which the chart never shipped + template: gateway/hpa.yaml + set: + gateway.hpa.targetRequestsPerMinute: 5400 + gateway.hpa.targetTokensPerMinute: 360000000 + asserts: + - lengthEqual: + path: spec.metrics + count: 2 + - notContains: + path: spec.metrics + content: + type: Pods + any: true + + - it: renders no ServiceMonitor by default + template: gateway/servicemonitor.yaml + asserts: + - hasDocuments: + count: 0 + + - it: refuses a ServiceMonitor without the metrics server, whose http port needs a bearer token + template: gateway/servicemonitor.yaml + set: + gateway.serviceMonitor.enabled: true + asserts: + - failedTemplate: + errorPattern: gateway.serviceMonitor.enabled requires gateway.metricsServer.enabled + + - it: scrapes each gateway pod through the metrics port + template: gateway/servicemonitor.yaml + release: + name: rel + namespace: llm + set: + gateway.serviceMonitor.enabled: true + gateway.metricsServer.enabled: true + gateway.serviceMonitor.labels: + release: kube-prometheus-stack + asserts: + - isKind: + of: ServiceMonitor + - equal: + path: metadata.labels.release + value: kube-prometheus-stack + - equal: + path: spec.selector.matchLabels + value: + app.kubernetes.io/name: litellm + app.kubernetes.io/instance: rel + app.kubernetes.io/component: gateway + - equal: + path: spec.namespaceSelector.matchNames + value: + - llm + - equal: + path: spec.endpoints + value: + - port: metrics + path: /metrics/ + interval: 15s + scrapeTimeout: 10s + scheme: http + + - it: honours a custom scrape interval + template: gateway/servicemonitor.yaml + set: + gateway.serviceMonitor.enabled: true + gateway.serviceMonitor.interval: 30s + gateway.metricsServer.enabled: true + asserts: + - equal: + path: spec.endpoints[0].interval + value: 30s diff --git a/helm/litellm/values.yaml b/helm/litellm/values.yaml index b5d535c992d..be1e6d43987 100644 --- a/helm/litellm/values.yaml +++ b/helm/litellm/values.yaml @@ -284,6 +284,16 @@ gateway: memory: 128Mi limits: memory: 512Mi + # Prometheus Operator ServiceMonitor for the gateway pods. Scrapes the + # `-metrics` Service, so it requires metricsServer above (the http + # port serves /metrics/ behind virtual-key auth). Every pod is its own scrape + # target, so the samples carry the `pod` label the per-pod autoscaling + # queries below group by. + serviceMonitor: + enabled: false + labels: {} + interval: 15s + scrapeTimeout: 10s image: repository: ghcr.io/berriai/litellm-gateway tag: "" # defaults to .Chart.AppVersion @@ -340,6 +350,25 @@ gateway: # policies: # - { type: Percent, value: 100, periodSeconds: 30 } behavior: {} + # Opt-in per-pod workload targets, rendered as autoscaling/v2 `Pods` metrics + # named `litellm_requests_per_second` and `litellm_tokens_per_second` with an + # AverageValue target. They coexist with the CPU/memory targets above: the + # HPA scales on whichever metric asks for the most replicas. Kubernetes has + # no idea what a token is, so a Prometheus Adapter must serve those two + # names on custom.metrics.k8s.io from the proxy's counters, grouped by the + # scrape target's `pod` label (enable serviceMonitor above): + # litellm_requests_per_second: + # sum(rate(litellm_proxy_total_requests_metric_total{<<.LabelMatchers>>}[1m])) by (<<.GroupBy>>) + # litellm_tokens_per_second: + # sum(rate(litellm_total_tokens_metric_total{<<.LabelMatchers>>}[1m])) by (<<.GroupBy>>) + # rate() over [1m] is already per second, so no `* 60`. How fast the HPA + # reacts is set by that window, the scrape interval and the HPA sync period + # (15s by default), not by the unit: keep serviceMonitor.interval at 15s or + # faster so a 1m window holds at least 4 samples. averageValue takes SI + # suffixes, so "6M" is six million tokens per second per pod. Tokens are + # counted when a response completes, so TPS trails long streams. + targetRequestsPerSecond: "" + targetTokensPerSecond: "" # PodDisruptionBudget for the gateway pods. Set exactly one of # `minAvailable` / `maxUnavailable` (minAvailable wins if both are set; # enabling without either falls back to `maxUnavailable: 1`). Disabled by diff --git a/terraform/litellm/aws/README.md b/terraform/litellm/aws/README.md index 923a4ca2da8..6d3e0269a15 100644 --- a/terraform/litellm/aws/README.md +++ b/terraform/litellm/aws/README.md @@ -258,6 +258,61 @@ gateway_metrics_port = 4001 gateway_metrics_scrape_cidrs = ["10.0.0.0/16"] ``` +### Scaling the gateway on requests and tokens + +By default the gateway service target-tracks CPU (`gateway_cpu_target`) and +memory (`gateway_memory_target`). Two more targets add workload signals next +to them. Application Auto Scaling evaluates every attached policy and follows +the one asking for the most tasks, so the resource policies keep working as a +floor while requests or tokens drive scale-out + +Both targets are per task per second, the way load is usually quoted (1k +rps, 75M tok/s). CloudWatch is the limit on how fast they react: target +tracking evaluates every metric, predefined or custom, aggregated over +60-second periods and has no period setting, so ECS reacts on a roughly +one-minute cadence whatever unit the variable is written in. The Kubernetes +charts get a faster signal because the Prometheus `rate()` window and scrape +interval are theirs to shorten + +`gateway_target_requests_per_second` adds an `ALBRequestCountPerTarget` +policy on the gateway target group. The ALB publishes that metric as requests +per minute per registered task, so the policy's target value is 60 times the +variable: 90 rps becomes a target of 5,400 per minute. No agent or sidecar is +needed + +`gateway_target_tokens_per_second` adds a metric-math policy over a +CloudWatch metric of the gateway's `litellm_total_tokens_metric_total` +counter and the service's `RunningTaskCount` from Container Insights. Nothing +native to ECS carries token throughput, so you publish that metric yourself +with the CloudWatch agent's Prometheus scraper pointed at the metrics sidecar +above. The agent emits the delta of a counter between scrapes, so `Sum` over +the 60-second period is the tokens served in that minute; the expression +divides by 60 (`tokens_per_second`) and then by the task count +(`tokens_per_second_per_task`). Tokens are counted when a response completes, +so long streams show up late in this signal. `gateway_tokens_metric` tells the +policy where the agent publishes: the namespace, the metric name (defaults to +the counter name) and the dimensions from your `metric_declaration` + +```hcl +gateway_metrics_port = 4001 +gateway_target_requests_per_second = 90 +gateway_target_tokens_per_second = 6000000 +gateway_tokens_metric = { + namespace = "LiteLLM/Prometheus" + dimensions = { ClusterName = "acme-litellm-prod", TaskDefinitionFamily = "acme-litellm-prod-gateway" } +} +``` + +Worked example for the request policy: 1,000 rps across 10 tasks is 100 rps +per task (the ALB reports it as 6,000 per minute per target) against a target +of 90 (5,400), so target tracking sizes the service to +`ceil(10 * 100 / 90) = 12` tasks. The token policy does the same arithmetic: +ten tasks handle 4,200,000,000 tokens in a minute, `tokens / 60` is +70,000,000 tokens per second and `tokens_per_second / running_tasks` is +7,000,000 against a target of 6,000,000, so the service grows to +`ceil(10 * 7000000 / 6000000) = 12`. Container Insights must be enabled on the +cluster for `RunningTaskCount` to exist + ## Tenant deployment Every resource the stack creates is named `${tenant}-litellm-${env}` (or diff --git a/terraform/litellm/aws/autoscaling.tf b/terraform/litellm/aws/autoscaling.tf index 71b6c24fac7..5197311d1b3 100644 --- a/terraform/litellm/aws/autoscaling.tf +++ b/terraform/litellm/aws/autoscaling.tf @@ -52,6 +52,105 @@ resource "aws_appautoscaling_policy" "gateway_memory" { } } +resource "aws_appautoscaling_policy" "gateway_requests" { + count = var.gateway_autoscaling_enabled && var.gateway_target_requests_per_second > 0 ? 1 : 0 + name = "${local.name}-gateway-requests" + policy_type = "TargetTrackingScaling" + service_namespace = aws_appautoscaling_target.gateway[0].service_namespace + resource_id = aws_appautoscaling_target.gateway[0].resource_id + scalable_dimension = aws_appautoscaling_target.gateway[0].scalable_dimension + + target_tracking_scaling_policy_configuration { + predefined_metric_specification { + predefined_metric_type = "ALBRequestCountPerTarget" + resource_label = "${aws_lb.this.arn_suffix}/${aws_lb_target_group.gateway.arn_suffix}" + } + # ALBRequestCountPerTarget is a per-minute count + target_value = var.gateway_target_requests_per_second * 60 + } +} + +resource "aws_appautoscaling_policy" "gateway_tokens" { + count = var.gateway_autoscaling_enabled && var.gateway_target_tokens_per_second > 0 ? 1 : 0 + name = "${local.name}-gateway-tokens" + policy_type = "TargetTrackingScaling" + service_namespace = aws_appautoscaling_target.gateway[0].service_namespace + resource_id = aws_appautoscaling_target.gateway[0].resource_id + scalable_dimension = aws_appautoscaling_target.gateway[0].scalable_dimension + + lifecycle { + precondition { + condition = var.gateway_tokens_metric != null + error_message = "gateway_tokens_metric is required when gateway_target_tokens_per_second > 0." + } + } + + target_tracking_scaling_policy_configuration { + target_value = var.gateway_target_tokens_per_second + + # target tracking has no period setting and always aggregates over 60s + customized_metric_specification { + metrics { + id = "tokens" + return_data = false + + metric_stat { + stat = "Sum" + + metric { + namespace = var.gateway_tokens_metric.namespace + metric_name = var.gateway_tokens_metric.name + + dynamic "dimensions" { + for_each = var.gateway_tokens_metric.dimensions + content { + name = dimensions.key + value = dimensions.value + } + } + } + } + } + + metrics { + id = "running_tasks" + return_data = false + + metric_stat { + stat = "Average" + + metric { + namespace = "ECS/ContainerInsights" + metric_name = "RunningTaskCount" + + dimensions { + name = "ClusterName" + value = aws_ecs_cluster.this.name + } + dimensions { + name = "ServiceName" + value = aws_ecs_service.gateway.name + } + } + } + } + + metrics { + id = "tokens_per_second" + expression = "tokens / 60" + return_data = false + } + + metrics { + id = "tokens_per_second_per_task" + expression = "tokens_per_second / running_tasks" + label = "Tokens per second per gateway task" + return_data = true + } + } + } +} + # ---------- Backend ---------- resource "aws_appautoscaling_target" "backend" { count = var.backend_autoscaling_enabled ? 1 : 0 diff --git a/terraform/litellm/aws/tests/workload_autoscaling.tftest.hcl b/terraform/litellm/aws/tests/workload_autoscaling.tftest.hcl new file mode 100644 index 00000000000..94e281faf94 --- /dev/null +++ b/terraform/litellm/aws/tests/workload_autoscaling.tftest.hcl @@ -0,0 +1,194 @@ +# Plan-only coverage for the gateway request and token autoscaling policies. +# Offline via mock_provider, same as byo_infrastructure.tftest.hcl. + +mock_provider "aws" { + mock_data "aws_iam_policy_document" { + defaults = { + json = "{\"Version\":\"2012-10-17\",\"Statement\":[]}" + } + } +} +mock_provider "random" {} + +variables { + region = "us-east-1" + tenant = "acme" + env = "test" + allow_plaintext_alb = true + azs = ["us-east-1a", "us-east-1b"] +} + +run "defaults_scale_on_cpu_and_memory_only" { + command = plan + + assert { + condition = alltrue([ + length(aws_appautoscaling_policy.gateway_cpu) == 1, + length(aws_appautoscaling_policy.gateway_memory) == 1, + length(aws_appautoscaling_policy.gateway_requests) == 0, + length(aws_appautoscaling_policy.gateway_tokens) == 0, + ]) + error_message = "Request and token policies must be absent by default while the CPU and memory policies stay." + } +} + +run "requests_per_second_adds_an_alb_request_count_policy" { + command = plan + + variables { + gateway_target_requests_per_second = 90 + } + + assert { + condition = length(aws_appautoscaling_policy.gateway_requests) == 1 && length(aws_appautoscaling_policy.gateway_tokens) == 0 + error_message = "A request target alone must add exactly the request policy." + } + + assert { + condition = alltrue([ + aws_appautoscaling_policy.gateway_requests[0].name == "acme-litellm-test-gateway-requests", + aws_appautoscaling_policy.gateway_requests[0].policy_type == "TargetTrackingScaling", + aws_appautoscaling_policy.gateway_requests[0].service_namespace == "ecs", + aws_appautoscaling_policy.gateway_requests[0].resource_id == "service/acme-litellm-test/acme-litellm-test-gateway", + aws_appautoscaling_policy.gateway_requests[0].scalable_dimension == "ecs:service:DesiredCount", + ]) + error_message = "The request policy must be a target-tracking policy on the gateway service's desired count." + } + + assert { + condition = alltrue([ + one(aws_appautoscaling_policy.gateway_requests[0].target_tracking_scaling_policy_configuration).target_value == 5400, + one(one(aws_appautoscaling_policy.gateway_requests[0].target_tracking_scaling_policy_configuration).predefined_metric_specification).predefined_metric_type == "ALBRequestCountPerTarget", + length(one(aws_appautoscaling_policy.gateway_requests[0].target_tracking_scaling_policy_configuration).customized_metric_specification) == 0, + ]) + error_message = "The request policy must track ALBRequestCountPerTarget at 60 times the configured requests per second per task." + } +} + +run "tokens_per_second_adds_a_metric_math_policy" { + command = plan + + variables { + gateway_target_tokens_per_second = 6000000 + gateway_tokens_metric = { + namespace = "LiteLLM/Prometheus" + dimensions = { ClusterName = "acme-litellm-test", TaskDefinitionFamily = "acme-litellm-test-gateway" } + } + } + + assert { + condition = length(aws_appautoscaling_policy.gateway_tokens) == 1 && length(aws_appautoscaling_policy.gateway_requests) == 0 + error_message = "A token target alone must add exactly the token policy." + } + + assert { + condition = alltrue([ + aws_appautoscaling_policy.gateway_tokens[0].name == "acme-litellm-test-gateway-tokens", + aws_appautoscaling_policy.gateway_tokens[0].resource_id == "service/acme-litellm-test/acme-litellm-test-gateway", + one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).target_value == 6000000, + length(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).predefined_metric_specification) == 0, + ]) + error_message = "The token policy must track a customized metric at the configured tokens per second per task." + } + + assert { + condition = alltrue([ + length(one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics) == 4, + { for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens"].id == "tokens", + { for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens"].return_data == false, + one({ for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens"].metric_stat).stat == "Sum", + one(one({ for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens"].metric_stat).metric).namespace == "LiteLLM/Prometheus", + one(one({ for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens"].metric_stat).metric).metric_name == "litellm_total_tokens_metric_total", + { for d in one(one({ for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens"].metric_stat).metric).dimensions : d.name => d.value } == { ClusterName = "acme-litellm-test", TaskDefinitionFamily = "acme-litellm-test-gateway" }, + ]) + error_message = "The first metric must sum the published token counter deltas under the configured namespace and dimensions." + } + + assert { + condition = alltrue([ + { for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["running_tasks"].id == "running_tasks", + { for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["running_tasks"].return_data == false, + one({ for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["running_tasks"].metric_stat).stat == "Average", + one(one({ for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["running_tasks"].metric_stat).metric).namespace == "ECS/ContainerInsights", + one(one({ for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["running_tasks"].metric_stat).metric).metric_name == "RunningTaskCount", + { for d in one(one({ for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["running_tasks"].metric_stat).metric).dimensions : d.name => d.value } == { ClusterName = "acme-litellm-test", ServiceName = "acme-litellm-test-gateway" }, + ]) + error_message = "The second metric must read the gateway service's Container Insights RunningTaskCount." + } + + assert { + condition = alltrue([ + { for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens_per_second"].expression == "tokens / 60", + { for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens_per_second"].return_data == false, + ]) + error_message = "The 60s period Sum must be divided by 60 to yield tokens per second." + } + + assert { + condition = alltrue([ + { for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens_per_second_per_task"].expression == "tokens_per_second / running_tasks", + { for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens_per_second_per_task"].return_data == true, + length([for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m if m.return_data]) == 1, + ]) + error_message = "Only the per-task tokens per second may return data to the scaling policy." + } +} + +run "tokens_per_second_needs_the_metric_location" { + command = plan + + variables { + gateway_target_tokens_per_second = 6000000 + } + + expect_failures = [ + aws_appautoscaling_policy.gateway_tokens, + ] +} + +run "requests_and_tokens_scale_next_to_cpu_and_memory" { + command = plan + + variables { + gateway_target_requests_per_second = 90 + gateway_target_tokens_per_second = 6000000 + gateway_tokens_metric = { namespace = "LiteLLM/Prometheus" } + } + + assert { + condition = alltrue([ + length(aws_appautoscaling_policy.gateway_cpu) == 1, + length(aws_appautoscaling_policy.gateway_memory) == 1, + length(aws_appautoscaling_policy.gateway_requests) == 1, + length(aws_appautoscaling_policy.gateway_tokens) == 1, + one(aws_appautoscaling_policy.gateway_cpu[0].target_tracking_scaling_policy_configuration).target_value == 70, + one(aws_appautoscaling_policy.gateway_memory[0].target_tracking_scaling_policy_configuration).target_value == 80, + ]) + error_message = "Workload policies must coexist with the CPU and memory policies at their default targets." + } + + assert { + condition = length(one(one({ for m in one(one(aws_appautoscaling_policy.gateway_tokens[0].target_tracking_scaling_policy_configuration).customized_metric_specification).metrics : m.id => m }["tokens"].metric_stat).metric).dimensions) == 0 + error_message = "Omitting dimensions must query the token metric without any." + } +} + +run "workload_targets_are_ignored_when_autoscaling_is_off" { + command = plan + + variables { + gateway_autoscaling_enabled = false + gateway_target_requests_per_second = 90 + gateway_target_tokens_per_second = 6000000 + gateway_tokens_metric = { namespace = "LiteLLM/Prometheus" } + } + + assert { + condition = alltrue([ + length(aws_appautoscaling_target.gateway) == 0, + length(aws_appautoscaling_policy.gateway_requests) == 0, + length(aws_appautoscaling_policy.gateway_tokens) == 0, + ]) + error_message = "Disabling gateway autoscaling must drop the workload policies with the target." + } +} diff --git a/terraform/litellm/aws/variables.tf b/terraform/litellm/aws/variables.tf index 667f6db63c9..10a9b392673 100644 --- a/terraform/litellm/aws/variables.tf +++ b/terraform/litellm/aws/variables.tf @@ -272,6 +272,47 @@ variable "gateway_memory_target" { default = 80 } +variable "gateway_target_requests_per_second" { + description = <<-EOT + Requests per second one gateway task should serve. Adds an + ALBRequestCountPerTarget target-tracking policy next to the CPU/memory + ones (Application Auto Scaling follows whichever asks for more tasks). + CloudWatch publishes that metric as a 1-minute count, so the policy + targets 60x this value and ECS reacts on a ~1 minute cadence. 0 skips + the policy. + EOT + type = number + default = 0 +} + +variable "gateway_target_tokens_per_second" { + description = <<-EOT + Tokens per second one gateway task should serve. Adds a target-tracking + policy on gateway_tokens_metric summed over each 60s period, divided by + 60 and by the service's Container Insights RunningTaskCount. Tokens are + counted when a response completes, so the signal trails long streams. + 0 skips the policy. + EOT + type = number + default = 0 +} + +variable "gateway_tokens_metric" { + description = <<-EOT + CloudWatch metric carrying the gateway's litellm_total_tokens_metric_total + counter, as published by the CloudWatch agent's Prometheus scraper (it + emits the delta between scrapes, so Sum over a period is the tokens + served in it). Required when gateway_target_tokens_per_second > 0. + dimensions must match the metric_declaration the agent publishes with. + EOT + type = object({ + namespace = string + name = optional(string, "litellm_total_tokens_metric_total") + dimensions = optional(map(string), {}) + }) + default = null +} + variable "backend_autoscaling_enabled" { description = "Toggle Application Auto Scaling target-tracking on the backend service." type = bool diff --git a/terraform/litellm/gcp/README.md b/terraform/litellm/gcp/README.md index c93e5f6b303..00ec7e8afd0 100644 --- a/terraform/litellm/gcp/README.md +++ b/terraform/litellm/gcp/README.md @@ -238,6 +238,24 @@ this with `litellm_license`. To tune the export cadence, set Behavior matches the AWS stack 1:1; the variable names are identical +### Autoscaling + +Cloud Run scales the gateway on request concurrency (plus its built-in CPU +target), not on a metric you attach. Each instance takes up to +`gateway_max_instance_request_concurrency` requests at once (default 80) +and Cloud Run adds instances between `gateway_min_instances` and +`gateway_max_instances` when the in-flight count fills up. That is the +request-rate signal for this stack: lower the concurrency for LLM streams +that hold a worker for tens of seconds, since a stream counts as one request +for as long as it is open + +There is no tokens-per-second path here. Cloud Run's autoscaler has no +custom-metric input, so the `litellm_total_tokens_metric_total` counter the +proxy exposes cannot drive it. If you need token-based scaling on GCP, run +the gateway on GKE with the Helm chart's `targetTokensPerSecond` (see +"Dependencies only" below) rather than wiring the counter into Cloud +Monitoring, which the autoscaler would ignore + ## Tenant deployment Every resource the stack creates is named `${tenant}-litellm-${env}` (or From d266b76a8bcbb6cbec408eeeb8758a7af1a92644 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Thu, 10 Sep 2026 13:48:21 -0700 Subject: [PATCH 51/60] fix(router): honor Codex reminders and map classifier failures --- litellm/responses/main.py | 48 ++++++++++++++++--- .../complexity_router/complexity_router.py | 4 +- .../test_responses_api_bridge_flag.py | 24 ++++++++++ .../router_strategy/test_complexity_router.py | 34 +++++++++++++ 4 files changed, 102 insertions(+), 8 deletions(-) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index e88cd618a8b..a68cd02e61b 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -4,10 +4,11 @@ from collections.abc import Coroutine, Generator, Iterable, Mapping from contextlib import contextmanager from dataclasses import dataclass from functools import partial -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast +from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Optional, TypeAlias, cast import httpx from pydantic import BaseModel +from typing_extensions import assert_never import litellm from litellm._logging import verbose_logger @@ -407,6 +408,37 @@ def _bridges_to_chat_completions( return responses_api_provider_config is None or use_chat_completions_api is True +_ResponsesCompatibilityFailure: TypeAlias = Literal["encrypted_task_unsupported"] + + +def _encrypted_task_support_failure( + responses_api_provider_config: BaseResponsesAPIConfig | None, use_chat_completions_api: bool +) -> _ResponsesCompatibilityFailure | None: + if ( + responses_api_provider_config is None + or _bridges_to_chat_completions(responses_api_provider_config, use_chat_completions_api) + or not responses_api_provider_config.supports_encrypted_agent_messages() + ): + return "encrypted_task_unsupported" + return None + + +def _raise_responses_compatibility_failure( + failure: _ResponsesCompatibilityFailure, model: str, custom_llm_provider: str | None +) -> NoReturn: + match failure: + case "encrypted_task_unsupported": + raise litellm.exception_type( + model=model, + custom_llm_provider=custom_llm_provider, + original_exception=ValueError( + "Encrypted task classification requires a compatible native Responses deployment" + ), + ) + case _: + assert_never(failure) + + def _deployment_passes_through_responses(model_info: object) -> bool: """Whether ``model_info.supported_endpoints`` opts the deployment into native ``{api_base}/responses``.""" if not isinstance(model_info, dict): @@ -1187,12 +1219,16 @@ def responses( model, custom_llm_provider, deployment_model_info ) - if require_encrypted_task_support and ( - _bridges_to_chat_completions(responses_api_provider_config, use_chat_completions_api) - or responses_api_provider_config is None - or not responses_api_provider_config.supports_encrypted_agent_messages() + if ( + require_encrypted_task_support + and ( + compatibility_failure := _encrypted_task_support_failure( + responses_api_provider_config, use_chat_completions_api + ) + ) + is not None ): - raise ValueError("Encrypted task classification requires a compatible native Responses deployment") + _raise_responses_compatibility_failure(compatibility_failure, model, custom_llm_provider) local_vars.update(kwargs) # Map reasoning_effort (from litellm_params/proxy config) to reasoning when not set diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 6df74b26f60..214872c17c9 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -1690,7 +1690,7 @@ class ComplexityRouter(CustomLogger): if self.config.classifier_type == "custom": return await self._classify_with_plugin(prompt, system_prompt, request_kwargs, raw_messages) if self.config.classifier_type in ("heuristic_first", "hybrid") and _encrypted_classifier_task( - request_kwargs, self._reminder_markers + request_kwargs, self._reminder_markers_for_request(request_kwargs or EMPTY_MAPPING) ): return await self._llm_classifier_outcome(prompt, system_prompt, request_kwargs, messages) if self.config.classifier_type == "heuristic_first" and self.config.classifier_llm_config is not None: @@ -2028,7 +2028,7 @@ class ComplexityRouter(CustomLogger): > 1 ) - encrypted_task: Final = _encrypted_classifier_task(request_kwargs, self._reminder_markers) + encrypted_task: Final = _encrypted_classifier_task(request_kwargs, marker_pairs) user_payload: Final = self._build_classifier_user_payload( prompt="The delegated task in the following agent_message." if encrypted_task is not None else prompt, system_prompt=system_prompt, diff --git a/tests/test_litellm/responses/test_responses_api_bridge_flag.py b/tests/test_litellm/responses/test_responses_api_bridge_flag.py index 57aa2a6baa2..ed44a9f4545 100644 --- a/tests/test_litellm/responses/test_responses_api_bridge_flag.py +++ b/tests/test_litellm/responses/test_responses_api_bridge_flag.py @@ -7,10 +7,14 @@ calls so routed requests do not hit a custom api_base /v1/responses endpoint. """ from importlib import import_module +from typing import Final from unittest.mock import MagicMock, patch +import httpx +import pytest import litellm +from litellm.llms.custom_httpx.http_handler import HTTPHandler from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse from litellm.types.utils import Choices, Message, ModelResponse, Usage @@ -18,6 +22,26 @@ from litellm.types.utils import Choices, Message, ModelResponse, Usage class TestUseResponsesApiBridgeFlag: """Test that bridge opt-in forces the chat completions path.""" + @pytest.mark.parametrize("model", ["openai/chat_completions/gpt-6-astra", "xai/test-classifier"]) + def test_encrypted_classifier_rejection_preserves_public_error(self, model: str) -> None: + respond: Final = MagicMock(side_effect=AssertionError("Incompatible classifier sent an upstream request")) + with httpx.Client(transport=httpx.MockTransport(respond)) as client: + with pytest.raises( + litellm.APIConnectionError, + match="Encrypted task classification requires a compatible native Responses deployment", + ) as error: + litellm.responses( + model=model, + input="Delegated task", + api_key="test-key", + api_base="https://classifier.test/v1", + client=HTTPHandler(client=client), + _require_encrypted_task_support=True, + num_retries=0, + ) + assert error.value.status_code == 500 + respond.assert_not_called() + @patch.object( import_module("litellm.responses.main").litellm_completion_transformation_handler, "response_api_handler" ) diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 2ea8a8b53fe..cd8f05d1e41 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -2577,6 +2577,40 @@ async def native_classifier_http() -> AsyncIterator[tuple[AsyncHTTPHandler, Magi class TestEncryptedTaskClassifier: + @pytest.mark.asyncio + @pytest.mark.parametrize("classifier_type", ["llm", "heuristic_first", "hybrid"]) + @pytest.mark.parametrize("codex", [True, False]) + @pytest.mark.parametrize( + "reminder", + [ + "cwd=/repo", + "Keep answers concise", + ], + ) + async def test_encrypted_task_detection_uses_request_reminder_markers( + self, classifier_type: str, codex: bool, reminder: str + ): + router, dependency = _native_classifier_router(classifier_type=classifier_type) + task: Final = _encrypted_agent_task() + request: Final = { + "input": [task, {"role": "user", "content": reminder}], + "metadata": {"user_agent": "codex-tui" if codex else "curl/8.7.1"}, + } + original: Final = deepcopy(request) + + result: Final = await router.async_pre_routing_hook(model="encrypted-router", request_kwargs=request) + + assert request == original + assert result.model == ("deep-model" if codex else "cheap-model") + if codex: + assert result.routing_decision["cause"] == "llm_classifier" + assert result.routing_decision["tier"] == "REASONING" + dependency.aresponses.assert_awaited_once() + assert dependency.aresponses.call_args.kwargs["input"][-1] == task + dependency.acompletion.assert_not_called() + else: + dependency.aresponses.assert_not_called() + @pytest.mark.asyncio @pytest.mark.parametrize("classifier_type", ["llm", "heuristic_first", "hybrid"]) @pytest.mark.parametrize("tier,model", [("SIMPLE", "cheap-model"), ("REASONING", "deep-model")]) From d3a0b0d45b814494e64edf82cd13b0683cd4aba8 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 10 Sep 2026 13:51:45 -0700 Subject: [PATCH 52/60] fix(voyage): let caller params win for contextual auto-chunking and drop duplicate cost map entries A flat list[str] sent to voyage-context-4 is treated as independent inputs and forwarded flat with enable_auto_chunking=True, chunk_size=32000, and input_type=document unless the caller already set input_type=query. Caller-supplied params now override the defaults instead of being clobbered. The voyage-4 family and voyage-context-4 cost map entries already exist on litellm_internal_staging, and voyage-4-nano is not served by the Voyage API, so those additions and their pricing test are dropped. --- .../embedding/transformation_contextual.py | 69 ++------- ...odel_prices_and_context_window_backup.json | 40 ------ model_prices_and_context_window.json | 40 ------ tests/llm_translation/test_voyage_ai.py | 131 ------------------ .../test_voyage_contextual_embedding.py | 51 ++++++- 5 files changed, 59 insertions(+), 272 deletions(-) diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index a9a4a01f8e3..870de8756bb 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -3,6 +3,7 @@ This module is used to transform the request and response for the Voyage context This would be used for all the contextualized embeddings models in Voyage. """ +from collections.abc import Mapping from typing import Final import httpx @@ -100,7 +101,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): "Authorization": f"Bearer {api_key}", } - AUTO_CHUNK_SIZE = 32000 + AUTO_CHUNK_SIZE: Final = 32000 def transform_embedding_request( self, @@ -109,71 +110,27 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): optional_params: dict, headers: dict, ) -> dict: - inputs, extra_params = self._transform_contextual_inputs(input, optional_params) return { - "inputs": inputs, + "inputs": [input] if isinstance(input, str) else input, "model": model, + **self._auto_chunk_params(input, optional_params), **optional_params, - **extra_params, } @classmethod - def _transform_contextual_inputs( + def _auto_chunk_params( cls, - input: AllEmbeddingInputValues | list[list[str]], # mutable-ok: union with AllEmbeddingInputValues - optional_params: dict, # mutable-ok: matches public API - ) -> tuple[list[str] | list[list[str]], dict]: # mutable-ok: returned to caller who owns it - """ - Normalize ``input`` for Voyage's contextualized embeddings API and - return ``(inputs, extra_params)`` where ``extra_params`` carries any - request fields (e.g. auto-chunking) needed for the chosen shape. - - The API contract (verified against the live endpoint) is: - - - A flat ``list[str]`` is only accepted with ``input_type="query"`` or - with ``enable_auto_chunking=True`` (which itself requires - ``input_type="document"``). - - A ``list[list[str]]`` (each inner list = one document's chunks) is - always accepted. - - So we prefer to send a flat ``list[str]`` and let the API auto-chunk, - instead of pre-wrapping into ``list[list[str]]``: - - - ``str`` -> ``[str]`` + ``enable_auto_chunking`` (input_type=document) - - flat ``list[str]`` + ``input_type="query"`` -> kept flat, as-is - - flat ``list[str]`` otherwise -> kept flat + ``enable_auto_chunking`` - (input_type=document) - - ``list[list[str]]`` -> passed through unchanged - - Reference: https://docs.voyageai.com/reference/contextualized-embeddings-api - """ - if isinstance(input, str): - if optional_params.get("input_type") == "query": - return [input], {} # mutable-ok: fresh list returned to caller - return [input], cls._auto_chunk_params(optional_params) # mutable-ok: fresh list returned to caller - - if all(isinstance(i, str) for i in input): - if optional_params.get("input_type") == "query": - return input, {} # pyright: ignore[reportReturnType] # narrowed to list[str] by all(isinstance) check - return input, cls._auto_chunk_params(optional_params) # pyright: ignore[reportReturnType] # narrowed to list[str] - - return input, {} # pyright: ignore[reportReturnType] # list[list[str]] branch - - @classmethod - def _auto_chunk_params(cls, optional_params: dict) -> dict: # mutable-ok: matches public API - """ - Params required to send a flat ``list[str]`` to the contextualized API. - - ``enable_auto_chunking=True`` requires ``input_type="document"``, so set - it unless the caller already provided an ``input_type``. - """ - params: dict[str, object] = { # mutable-ok: building return value + input: AllEmbeddingInputValues | list[list[str]], + optional_params: Mapping[str, object], + ) -> Mapping[str, object]: + is_flat: Final = isinstance(input, str) or all(isinstance(item, str) for item in input) + if not is_flat or optional_params.get("input_type") == "query": + return {} + return { "enable_auto_chunking": True, "chunk_size": cls.AUTO_CHUNK_SIZE, + "input_type": "document", } - if not optional_params.get("input_type"): - params["input_type"] = "document" - return params def transform_embedding_response( self, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index fb1cef7ff58..12c3c5e2be9 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -39676,38 +39676,6 @@ "mode": "embedding", "output_cost_per_token": 0.0 }, - "voyage/voyage-4": { - "input_cost_per_token": 6e-08, - "litellm_provider": "voyage", - "max_input_tokens": 32000, - "max_tokens": 32000, - "mode": "embedding", - "output_cost_per_token": 0.0 - }, - "voyage/voyage-4-large": { - "input_cost_per_token": 1.2e-07, - "litellm_provider": "voyage", - "max_input_tokens": 32000, - "max_tokens": 32000, - "mode": "embedding", - "output_cost_per_token": 0.0 - }, - "voyage/voyage-4-lite": { - "input_cost_per_token": 2e-08, - "litellm_provider": "voyage", - "max_input_tokens": 32000, - "max_tokens": 32000, - "mode": "embedding", - "output_cost_per_token": 0.0 - }, - "voyage/voyage-4-nano": { - "input_cost_per_token": 0.0, - "litellm_provider": "voyage", - "max_input_tokens": 32000, - "max_tokens": 32000, - "mode": "embedding", - "output_cost_per_token": 0.0 - }, "voyage/voyage-code-2": { "input_cost_per_token": 1.2e-07, "litellm_provider": "voyage", @@ -39732,14 +39700,6 @@ "mode": "embedding", "output_cost_per_token": 0.0 }, - "voyage/voyage-context-4": { - "input_cost_per_token": 1.2e-07, - "litellm_provider": "voyage", - "max_input_tokens": 120000, - "max_tokens": 120000, - "mode": "embedding", - "output_cost_per_token": 0.0 - }, "voyage/voyage-finance-2": { "input_cost_per_token": 1.2e-07, "litellm_provider": "voyage", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index fb1cef7ff58..12c3c5e2be9 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -39676,38 +39676,6 @@ "mode": "embedding", "output_cost_per_token": 0.0 }, - "voyage/voyage-4": { - "input_cost_per_token": 6e-08, - "litellm_provider": "voyage", - "max_input_tokens": 32000, - "max_tokens": 32000, - "mode": "embedding", - "output_cost_per_token": 0.0 - }, - "voyage/voyage-4-large": { - "input_cost_per_token": 1.2e-07, - "litellm_provider": "voyage", - "max_input_tokens": 32000, - "max_tokens": 32000, - "mode": "embedding", - "output_cost_per_token": 0.0 - }, - "voyage/voyage-4-lite": { - "input_cost_per_token": 2e-08, - "litellm_provider": "voyage", - "max_input_tokens": 32000, - "max_tokens": 32000, - "mode": "embedding", - "output_cost_per_token": 0.0 - }, - "voyage/voyage-4-nano": { - "input_cost_per_token": 0.0, - "litellm_provider": "voyage", - "max_input_tokens": 32000, - "max_tokens": 32000, - "mode": "embedding", - "output_cost_per_token": 0.0 - }, "voyage/voyage-code-2": { "input_cost_per_token": 1.2e-07, "litellm_provider": "voyage", @@ -39732,14 +39700,6 @@ "mode": "embedding", "output_cost_per_token": 0.0 }, - "voyage/voyage-context-4": { - "input_cost_per_token": 1.2e-07, - "litellm_provider": "voyage", - "max_input_tokens": 120000, - "max_tokens": 120000, - "mode": "embedding", - "output_cost_per_token": 0.0 - }, "voyage/voyage-finance-2": { "input_cost_per_token": 1.2e-07, "litellm_provider": "voyage", diff --git a/tests/llm_translation/test_voyage_ai.py b/tests/llm_translation/test_voyage_ai.py index e6c73240a0e..438bc4f3507 100644 --- a/tests/llm_translation/test_voyage_ai.py +++ b/tests/llm_translation/test_voyage_ai.py @@ -72,26 +72,6 @@ class TestVoyageAI(BaseLLMEmbeddingTest): assert response.usage.total_tokens > 0 -@pytest.mark.parametrize( - "model, expected_input_cost", - [ - ("voyage/voyage-4", 6e-08), - ("voyage/voyage-4-large", 1.2e-07), - ("voyage/voyage-4-lite", 2e-08), - ("voyage/voyage-4-nano", 0.0), - ("voyage/voyage-context-4", 1.2e-07), - ], -) -def test_voyage_4_family_pricing_registered(model, expected_input_cost): - """The voyage-4 family and voyage-context-4 must be in the cost map.""" - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - info = litellm.get_model_info(model=model) - assert info["litellm_provider"] == "voyage" - assert info["mode"] == "embedding" - assert info["input_cost_per_token"] == expected_input_cost - - def test_voyage_ai_embedding_extra_params(): """Test Voyage AI embedding with extra parameters""" try: @@ -220,117 +200,6 @@ class TestVoyageContextualEmbeddings: assert transformed["model"] == "voyage-context-3" assert transformed["encoding_format"] == "float" - def test_contextual_flat_list_str_is_auto_chunked(self): - """Flat list[str] without input_type stays flat and is auto-chunked. - - The API accepts a flat list[str] only with input_type="query" or with - enable_auto_chunking=True (which requires input_type="document"), so we - keep the list flat and let the API auto-chunk it. - """ - from litellm.llms.voyage.embedding.transformation_contextual import ( - VoyageContextualEmbeddingConfig, - ) - - config = VoyageContextualEmbeddingConfig() - - transformed = config.transform_embedding_request( - "voyage-context-4", ["Hello", "world"], {}, {} - ) - - assert transformed["inputs"] == ["Hello", "world"] - assert transformed["model"] == "voyage-context-4" - assert transformed["enable_auto_chunking"] is True - assert transformed["chunk_size"] == 32000 - assert transformed["input_type"] == "document" - - def test_contextual_flat_list_str_query_stays_flat(self): - """Flat list[str] with input_type='query' is sent as-is (list[str]).""" - from litellm.llms.voyage.embedding.transformation_contextual import ( - VoyageContextualEmbeddingConfig, - ) - - config = VoyageContextualEmbeddingConfig() - - transformed = config.transform_embedding_request( - "voyage-context-4", - ["Hello", "world"], - {"input_type": "query"}, - {}, - ) - - assert transformed["inputs"] == ["Hello", "world"] - assert transformed["input_type"] == "query" - # A query list is accepted as-is, no auto-chunking needed. - assert "enable_auto_chunking" not in transformed - - def test_contextual_flat_list_str_document_input_type_preserved(self): - """Explicit input_type='document' is preserved while auto-chunking.""" - from litellm.llms.voyage.embedding.transformation_contextual import ( - VoyageContextualEmbeddingConfig, - ) - - config = VoyageContextualEmbeddingConfig() - - transformed = config.transform_embedding_request( - "voyage-context-4", - ["Hello", "world"], - {"input_type": "document"}, - {}, - ) - - assert transformed["inputs"] == ["Hello", "world"] - assert transformed["input_type"] == "document" - assert transformed["enable_auto_chunking"] is True - assert transformed["chunk_size"] == 32000 - - def test_contextual_single_string_is_auto_chunked(self): - """A single string becomes a one-element flat list and is auto-chunked.""" - from litellm.llms.voyage.embedding.transformation_contextual import ( - VoyageContextualEmbeddingConfig, - ) - - config = VoyageContextualEmbeddingConfig() - - transformed = config.transform_embedding_request( - "voyage-context-4", "Hello", {}, {} - ) - - assert transformed["inputs"] == ["Hello"] - assert transformed["enable_auto_chunking"] is True - assert transformed["chunk_size"] == 32000 - assert transformed["input_type"] == "document" - - def test_contextual_single_string_query_no_auto_chunk(self): - """A single string with input_type='query' is not auto-chunked.""" - from litellm.llms.voyage.embedding.transformation_contextual import ( - VoyageContextualEmbeddingConfig, - ) - - config = VoyageContextualEmbeddingConfig() - - transformed = config.transform_embedding_request( - "voyage-context-4", "Hello", {"input_type": "query"}, {} - ) - - assert transformed["inputs"] == ["Hello"] - assert transformed["input_type"] == "query" - assert "enable_auto_chunking" not in transformed - - def test_contextual_nested_input_passthrough(self): - """Already-nested list[list[str]] input is passed through unchanged.""" - from litellm.llms.voyage.embedding.transformation_contextual import ( - VoyageContextualEmbeddingConfig, - ) - - config = VoyageContextualEmbeddingConfig() - - nested = [["Hello", "world"], ["Test"]] - transformed = config.transform_embedding_request( - "voyage-context-4", nested, {}, {} - ) - - assert transformed["inputs"] == nested - def test_contextual_embedding_response_transformation(self): """Test response transformation for contextual embeddings""" from litellm.llms.voyage.embedding.transformation_contextual import ( diff --git a/tests/test_litellm/llms/voyage/test_voyage_contextual_embedding.py b/tests/test_litellm/llms/voyage/test_voyage_contextual_embedding.py index 865608e286f..d3e912ce6af 100644 --- a/tests/test_litellm/llms/voyage/test_voyage_contextual_embedding.py +++ b/tests/test_litellm/llms/voyage/test_voyage_contextual_embedding.py @@ -74,15 +74,11 @@ class TestVoyageContextualEmbeddings: assert headers == {"Authorization": "Bearer test-key"} def test_validate_environment_secret_fallback(self, monkeypatch): - import litellm.llms.voyage.embedding.transformation_contextual as module from litellm.llms.voyage.embedding.transformation_contextual import ( VoyageContextualEmbeddingConfig, ) - def fake_get_secret(name): - return "secret-key" if name == "VOYAGE_API_KEY" else None - - monkeypatch.setattr(module, "get_secret_str", fake_get_secret) + monkeypatch.setenv("VOYAGE_API_KEY", "secret-key") config = VoyageContextualEmbeddingConfig() headers = config.validate_environment( {}, "voyage-context-4", [], {}, {}, api_key=None @@ -142,6 +138,51 @@ class TestVoyageContextualEmbeddings: assert transformed["input_type"] == "document" assert transformed["enable_auto_chunking"] is True + def test_flat_list_str_caller_chunk_params_win(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + transformed = config.transform_embedding_request( + "voyage-context-4", + ["Hello", "world"], + {"input_type": "document", "chunk_size": 512, "chunk_overlap": 32}, + {}, + ) + assert transformed["enable_auto_chunking"] is True + assert transformed["chunk_size"] == 512 + assert transformed["chunk_overlap"] == 32 + assert transformed["input_type"] == "document" + + def test_flat_list_str_caller_can_disable_auto_chunking(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + transformed = config.transform_embedding_request( + "voyage-context-4", ["Hello"], {"enable_auto_chunking": False}, {} + ) + assert transformed["enable_auto_chunking"] is False + assert transformed["input_type"] == "document" + + def test_nested_list_keeps_caller_params(self): + from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, + ) + + config = VoyageContextualEmbeddingConfig() + transformed = config.transform_embedding_request( + "voyage-context-4", [["Hello", "world"]], {"input_type": "document", "output_dimension": 512}, {} + ) + assert transformed == { + "inputs": [["Hello", "world"]], + "model": "voyage-context-4", + "input_type": "document", + "output_dimension": 512, + } + def test_single_string_auto_chunked(self): from litellm.llms.voyage.embedding.transformation_contextual import ( VoyageContextualEmbeddingConfig, From 9cd1c4c29ce45dea3c285d7bc7a1c46183757c80 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 13:51:59 -0700 Subject: [PATCH 53/60] feat(terraform): prometheus metrics sidecar for the GCP Cloud Run gateway (#40614) * feat(terraform): Prometheus metrics sidecar for the GCP Cloud Run gateway gateway_metrics_port adds a metrics container running litellm.proxy.prometheus_metrics_server next to the gateway, sharing the PROMETHEUS_MULTIPROC_DIR over an in-memory volume, plus a Managed Service for Prometheus collector sidecar that scrapes it over localhost and writes to Cloud Monitoring. The gateway stays on port 4000 and the load balancer routing is unchanged. Resolves LIT-7502 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(terraform): reject fractional and collector health ports for gateway_metrics_port Greptile review on #40614: 4000.5 fails integer port parsing and 13133 collides with the gmp sidecar liveness listener. Also stop claiming the load balancer's own /metrics goes away Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- terraform/litellm/gcp/README.md | 31 +++ terraform/litellm/gcp/cloudrun.tf | 120 ++++++++++- .../litellm/gcp/examples/default/main.tf | 2 + .../examples/default/terraform.tfvars.example | 6 + .../litellm/gcp/examples/default/variables.tf | 6 + terraform/litellm/gcp/iam.tf | 24 +++ terraform/litellm/gcp/secrets.tf | 17 ++ .../gcp/tests/metrics_sidecar.tftest.hcl | 200 ++++++++++++++++++ terraform/litellm/gcp/variables.tf | 38 ++++ 9 files changed, 442 insertions(+), 2 deletions(-) create mode 100644 terraform/litellm/gcp/tests/metrics_sidecar.tftest.hcl diff --git a/terraform/litellm/gcp/README.md b/terraform/litellm/gcp/README.md index 00ec7e8afd0..66095ec0108 100644 --- a/terraform/litellm/gcp/README.md +++ b/terraform/litellm/gcp/README.md @@ -238,6 +238,37 @@ this with `litellm_license`. To tune the export cadence, set Behavior matches the AWS stack 1:1; the variable names are identical +### Prometheus metrics sidecar + +`gateway_metrics_port` adds a `metrics` sidecar +(`python -m litellm.proxy.prometheus_metrics_server`) to the gateway Cloud Run +service that aggregates the workers' samples over an in-memory volume shared +with the gateway container, so the collector's scrape never runs on an +inference worker. Cloud Run only routes traffic to the gateway container, so +the load balancer keeps hitting port 4000 (including the gateway's own +authenticated `/metrics`, which stays as it was) and the sidecar port is +reachable on localhost inside the instance only. To get the series out, the +stack also adds Google's +[Managed Service for Prometheus sidecar](https://cloud.google.com/stackdriver/docs/managed-prometheus/cloudrun-sidecar) +(`gateway_metrics_collector_image`) with a `RunMonitoring` config stored in +Secret Manager that scrapes `localhost:/metrics` every 30s and writes to +Cloud Monitoring as `prometheus.googleapis.com/...` metrics. Enabling it grants +the runtime service account `roles/monitoring.metricWriter` and +`roles/logging.logWriter` on the project. Needs `gateway_image` v1.101.0 or +newer. See [Prometheus metrics](https://docs.litellm.ai/docs/proxy/prometheus) +for the metrics themselves + +```hcl +gateway_metrics_port = 4001 +``` + +The collector scrapes from inside the instance, so scrapes on an instance with +no in-flight requests can fail when CPU is throttled between requests. Keep +`gateway_min_instances` at 1 or more and, if you see gaps, enable +instance-based billing on the gateway service. Unlike the AWS stack there is +no `gateway_metrics_scrape_cidrs`: nothing outside the instance can reach the +sidecar port, so there is no network rule to open + ### Autoscaling Cloud Run scales the gateway on request concurrency (plus its built-in CPU diff --git a/terraform/litellm/gcp/cloudrun.tf b/terraform/litellm/gcp/cloudrun.tf index 84ae8b9247f..405893b132b 100644 --- a/terraform/litellm/gcp/cloudrun.tf +++ b/terraform/litellm/gcp/cloudrun.tf @@ -150,6 +150,21 @@ locals { [local.gateway_launch_cmd], )) + metrics_enabled = var.create_runtime && var.gateway_metrics_port != null + metrics_multiproc_dir = "/tmp/litellm_prometheus_multiproc" + metrics_volume = "prometheus-multiproc" + metrics_env_kv = local.metrics_enabled ? [{ name = "PROMETHEUS_MULTIPROC_DIR", value = local.metrics_multiproc_dir }] : [] + metrics_config_volume = "gmp-config" + + metrics_run_monitoring_yaml = local.metrics_enabled ? yamlencode({ + apiVersion = "monitoring.googleapis.com/v1beta" + kind = "RunMonitoring" + metadata = { name = "${local.name}-gateway" } + spec = { + endpoints = [{ port = var.gateway_metrics_port, path = "/metrics", interval = "30s" }] + } + }) : "" + backend_args = join(" && ", concat( local.redis_ca_fragment, local.database_url_fragment, @@ -197,6 +212,7 @@ resource "google_cloud_run_v2_service" "gateway" { } containers { + name = "gateway" image = local.gateway_image command = ["sh", "-c"] args = [local.gateway_args] @@ -213,7 +229,7 @@ resource "google_cloud_run_v2_service" "gateway" { } dynamic "env" { - for_each = concat(local.shared_env_kv, local.gateway_otel_env_kv, local.billing_metrics_env_kv, local.gateway_extra_env_kv, local.proxy_config_env) + for_each = concat(local.shared_env_kv, local.gateway_otel_env_kv, local.billing_metrics_env_kv, local.gateway_extra_env_kv, local.proxy_config_env, local.metrics_env_kv) content { name = env.value.name value = env.value.value @@ -241,6 +257,14 @@ resource "google_cloud_run_v2_service" "gateway" { } } + dynamic "volume_mounts" { + for_each = local.metrics_enabled ? [1] : [] + content { + name = local.metrics_volume + mount_path = local.metrics_multiproc_dir + } + } + startup_probe { http_get { path = "/health/readiness" @@ -262,6 +286,71 @@ resource "google_cloud_run_v2_service" "gateway" { } } + dynamic "containers" { + for_each = local.metrics_enabled ? [1] : [] + content { + name = "metrics" + image = local.gateway_image + command = ["python", "-m", "litellm.proxy.prometheus_metrics_server"] + args = ["--port", tostring(var.gateway_metrics_port)] + + dynamic "env" { + for_each = local.metrics_env_kv + content { + name = env.value.name + value = env.value.value + } + } + + volume_mounts { + name = local.metrics_volume + mount_path = local.metrics_multiproc_dir + } + + startup_probe { + http_get { + path = "/health" + port = var.gateway_metrics_port + } + period_seconds = 5 + timeout_seconds = 3 + failure_threshold = 12 + } + + liveness_probe { + http_get { + path = "/health" + port = var.gateway_metrics_port + } + period_seconds = 30 + timeout_seconds = 5 + } + } + } + + dynamic "containers" { + for_each = local.metrics_enabled ? [1] : [] + content { + name = "collector" + image = var.gateway_metrics_collector_image + depends_on = ["metrics"] + + volume_mounts { + name = local.metrics_config_volume + mount_path = "/etc/rungmp" + } + + liveness_probe { + http_get { + path = "/liveness" + port = 13133 + } + period_seconds = 30 + timeout_seconds = 30 + } + } + } + dynamic "volumes" { for_each = local.proxy_config_enabled ? [1] : [] content { @@ -272,6 +361,31 @@ resource "google_cloud_run_v2_service" "gateway" { } } } + + dynamic "volumes" { + for_each = local.metrics_enabled ? [1] : [] + content { + name = local.metrics_volume + empty_dir { + medium = "MEMORY" + size_limit = "256Mi" + } + } + } + + dynamic "volumes" { + for_each = local.metrics_enabled ? [1] : [] + content { + name = local.metrics_config_volume + secret { + secret = google_secret_manager_secret.metrics_run_monitoring[0].secret_id + items { + version = "latest" + path = "config.yaml" + } + } + } + } } depends_on = [ @@ -283,6 +397,8 @@ resource "google_cloud_run_v2_service" "gateway" { google_secret_manager_secret_iam_member.billing_metrics_client_cert, google_secret_manager_secret_iam_member.billing_metrics_client_key, google_secret_manager_secret_iam_member.billing_metrics_ca_cert, + google_secret_manager_secret_iam_member.metrics_run_monitoring, + google_project_iam_member.runtime_metric_writer, google_storage_bucket_iam_member.proxy_config_runtime, google_sql_user.app, # Don't go live until the schema is migrated; otherwise the proxy boots, @@ -332,7 +448,7 @@ resource "google_cloud_run_v2_service" "backend" { } dynamic "env" { - for_each = concat(local.shared_env_kv, local.backend_default_env_kv, local.backend_otel_env_kv, local.billing_metrics_env_kv, local.backend_extra_env_kv, local.proxy_config_env) + for_each = concat(local.shared_env_kv, local.backend_default_env_kv, local.backend_otel_env_kv, local.billing_metrics_env_kv, local.backend_extra_env_kv, local.proxy_config_env, local.metrics_env_kv) content { name = env.value.name value = env.value.value diff --git a/terraform/litellm/gcp/examples/default/main.tf b/terraform/litellm/gcp/examples/default/main.tf index f44b2a9a001..782d4ee4b65 100644 --- a/terraform/litellm/gcp/examples/default/main.tf +++ b/terraform/litellm/gcp/examples/default/main.tf @@ -53,4 +53,6 @@ module "litellm" { backend_extra_env = var.backend_extra_env gateway_extra_secrets = var.gateway_extra_secrets backend_extra_secrets = var.backend_extra_secrets + + gateway_metrics_port = var.gateway_metrics_port } diff --git a/terraform/litellm/gcp/examples/default/terraform.tfvars.example b/terraform/litellm/gcp/examples/default/terraform.tfvars.example index c35206503bb..ec6206734e2 100644 --- a/terraform/litellm/gcp/examples/default/terraform.tfvars.example +++ b/terraform/litellm/gcp/examples/default/terraform.tfvars.example @@ -107,3 +107,9 @@ env = "stage" # main.tf (otel_endpoint, otel_exporter, otel_environment_name, # otel_capture_message_content, otel_headers_secret). Full docs in # ../../variables.tf. + +# ---------- Prometheus metrics sidecar ---------- +# Serve /metrics from a sidecar in the gateway service instead of the inference +# workers. Scraped inside the instance by the Managed Service for Prometheus +# sidecar and written to Cloud Monitoring; see ../../README.md. +# gateway_metrics_port = 4001 diff --git a/terraform/litellm/gcp/examples/default/variables.tf b/terraform/litellm/gcp/examples/default/variables.tf index 88b57ce27eb..08b78346df3 100644 --- a/terraform/litellm/gcp/examples/default/variables.tf +++ b/terraform/litellm/gcp/examples/default/variables.tf @@ -142,3 +142,9 @@ variable "backend_extra_secrets" { type = map(string) default = {} } + +variable "gateway_metrics_port" { + description = "Port for the Prometheus metrics sidecar in the gateway service. Null keeps /metrics on the gateway port only." + type = number + default = null +} diff --git a/terraform/litellm/gcp/iam.tf b/terraform/litellm/gcp/iam.tf index 509e6d48ffd..5c2bc38bc12 100644 --- a/terraform/litellm/gcp/iam.tf +++ b/terraform/litellm/gcp/iam.tf @@ -116,3 +116,27 @@ resource "google_secret_manager_secret_iam_member" "billing_metrics_ca_cert" { role = "roles/secretmanager.secretAccessor" member = "serviceAccount:${google_service_account.runtime.email}" } + +resource "google_secret_manager_secret_iam_member" "metrics_run_monitoring" { + count = local.metrics_enabled ? 1 : 0 + + secret_id = google_secret_manager_secret.metrics_run_monitoring[0].id + role = "roles/secretmanager.secretAccessor" + member = "serviceAccount:${google_service_account.runtime.email}" +} + +resource "google_project_iam_member" "runtime_metric_writer" { + count = local.metrics_enabled ? 1 : 0 + + project = var.project_id + role = "roles/monitoring.metricWriter" + member = "serviceAccount:${google_service_account.runtime.email}" +} + +resource "google_project_iam_member" "runtime_log_writer" { + count = local.metrics_enabled ? 1 : 0 + + project = var.project_id + role = "roles/logging.logWriter" + member = "serviceAccount:${google_service_account.runtime.email}" +} diff --git a/terraform/litellm/gcp/secrets.tf b/terraform/litellm/gcp/secrets.tf index 6ec77139996..b3b656721ab 100644 --- a/terraform/litellm/gcp/secrets.tf +++ b/terraform/litellm/gcp/secrets.tf @@ -118,3 +118,20 @@ resource "google_secret_manager_secret_version" "billing_metrics_ca_cert" { secret = google_secret_manager_secret.billing_metrics_ca_cert[0].id secret_data = var.billing_metrics_ca_cert_pem } + +resource "google_secret_manager_secret" "metrics_run_monitoring" { + count = local.metrics_enabled ? 1 : 0 + + secret_id = "${local.name}-gateway-run-monitoring" + labels = local.labels + replication { + auto {} + } +} + +resource "google_secret_manager_secret_version" "metrics_run_monitoring" { + count = local.metrics_enabled ? 1 : 0 + + secret = google_secret_manager_secret.metrics_run_monitoring[0].id + secret_data = local.metrics_run_monitoring_yaml +} diff --git a/terraform/litellm/gcp/tests/metrics_sidecar.tftest.hcl b/terraform/litellm/gcp/tests/metrics_sidecar.tftest.hcl new file mode 100644 index 00000000000..7de62065d23 --- /dev/null +++ b/terraform/litellm/gcp/tests/metrics_sidecar.tftest.hcl @@ -0,0 +1,200 @@ +mock_provider "google" { + mock_resource "google_redis_instance" { + defaults = { + host = "10.0.0.4" + port = 6379 + server_ca_certs = [{ + cert = "-----BEGIN CERTIFICATE-----\nmock\n-----END CERTIFICATE-----" + }] + } + } +} + +mock_provider "google-beta" {} +mock_provider "random" {} + +variables { + project_id = "test-project" + tenant = "tenant" + env = "test" + allow_plaintext_lb = true + image_registry = "us-central1-docker.pkg.dev/test-project/litellm" +} + +run "metrics_sidecar_off_by_default" { + command = plan + + assert { + condition = length(google_cloud_run_v2_service.gateway[0].template[0].containers) == 1 + error_message = "The gateway service must run only the gateway container when gateway_metrics_port is null." + } + + assert { + condition = length([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[0].env : e if e.name == "PROMETHEUS_MULTIPROC_DIR"]) == 0 + error_message = "PROMETHEUS_MULTIPROC_DIR must not be set when the metrics sidecar is off." + } + + assert { + condition = length(google_cloud_run_v2_service.gateway[0].template[0].volumes) == 0 + error_message = "No shared multiproc or collector config volume must exist when the metrics sidecar is off." + } + + assert { + condition = alltrue([ + length(google_secret_manager_secret.metrics_run_monitoring) == 0, + length(google_secret_manager_secret_version.metrics_run_monitoring) == 0, + length(google_secret_manager_secret_iam_member.metrics_run_monitoring) == 0, + length(google_project_iam_member.runtime_metric_writer) == 0, + length(google_project_iam_member.runtime_log_writer) == 0, + ]) + error_message = "No RunMonitoring secret or monitoring IAM must be created when the metrics sidecar is off." + } +} + +run "metrics_sidecar_enabled" { + command = plan + + variables { + gateway_metrics_port = 4001 + } + + assert { + condition = join(",", [for c in google_cloud_run_v2_service.gateway[0].template[0].containers : c.name]) == "gateway,metrics,collector" + error_message = "gateway_metrics_port must add the metrics and collector sidecars after the gateway container." + } + + assert { + condition = alltrue([ + google_cloud_run_v2_service.gateway[0].template[0].containers[1].image == local.gateway_image, + join(" ", google_cloud_run_v2_service.gateway[0].template[0].containers[1].command) == "python -m litellm.proxy.prometheus_metrics_server", + join(" ", google_cloud_run_v2_service.gateway[0].template[0].containers[1].args) == "--port 4001", + ]) + error_message = "The metrics sidecar must run the gateway image's prometheus_metrics_server on the configured port." + } + + assert { + condition = alltrue([ + length([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[0].env : e if e.name == "PROMETHEUS_MULTIPROC_DIR" && e.value == "/tmp/litellm_prometheus_multiproc"]) == 1, + length([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[1].env : e if e.name == "PROMETHEUS_MULTIPROC_DIR" && e.value == "/tmp/litellm_prometheus_multiproc"]) == 1, + ]) + error_message = "Gateway and metrics containers must share PROMETHEUS_MULTIPROC_DIR." + } + + assert { + condition = alltrue([ + length([for m in google_cloud_run_v2_service.gateway[0].template[0].containers[0].volume_mounts : m if m.name == "prometheus-multiproc" && m.mount_path == "/tmp/litellm_prometheus_multiproc"]) == 1, + length([for m in google_cloud_run_v2_service.gateway[0].template[0].containers[1].volume_mounts : m if m.name == "prometheus-multiproc" && m.mount_path == "/tmp/litellm_prometheus_multiproc"]) == 1, + length([for v in google_cloud_run_v2_service.gateway[0].template[0].volumes : v if v.name == "prometheus-multiproc" && length(v.empty_dir) == 1 && v.empty_dir[0].medium == "MEMORY"]) == 1, + ]) + error_message = "Gateway and metrics containers must mount the same in-memory empty_dir at the multiproc dir." + } + + assert { + condition = alltrue([ + google_cloud_run_v2_service.gateway[0].template[0].containers[1].startup_probe[0].http_get[0].path == "/health", + google_cloud_run_v2_service.gateway[0].template[0].containers[1].startup_probe[0].http_get[0].port == 4001, + google_cloud_run_v2_service.gateway[0].template[0].containers[1].liveness_probe[0].http_get[0].path == "/health", + google_cloud_run_v2_service.gateway[0].template[0].containers[1].liveness_probe[0].http_get[0].port == 4001, + ]) + error_message = "The metrics sidecar must be probed on /health at the configured port." + } + + assert { + condition = length(google_cloud_run_v2_service.gateway[0].template[0].containers[1].ports) == 0 && length(google_cloud_run_v2_service.gateway[0].template[0].containers[2].ports) == 0 + error_message = "Only the gateway container may declare a port; Cloud Run routes ingress to exactly one container." + } + + assert { + condition = alltrue([ + google_cloud_run_v2_service.gateway[0].template[0].containers[0].ports[0].container_port == 4000, + google_cloud_run_v2_service.gateway[0].template[0].containers[0].startup_probe[0].http_get[0].port == 4000, + google_cloud_run_v2_service.gateway[0].template[0].containers[0].liveness_probe[0].http_get[0].port == 4000, + google_compute_region_network_endpoint_group.gateway[0].cloud_run[0].service == "${local.name}-gateway", + ]) + error_message = "The gateway must stay on port 4000 and remain the load balancer's Cloud Run target." + } + + assert { + condition = alltrue([ + google_cloud_run_v2_service.gateway[0].template[0].containers[2].image == var.gateway_metrics_collector_image, + join(",", google_cloud_run_v2_service.gateway[0].template[0].containers[2].depends_on) == "metrics", + length([for m in google_cloud_run_v2_service.gateway[0].template[0].containers[2].volume_mounts : m if m.name == "gmp-config" && m.mount_path == "/etc/rungmp"]) == 1, + google_cloud_run_v2_service.gateway[0].template[0].containers[2].liveness_probe[0].http_get[0].port == 13133, + ]) + error_message = "The collector sidecar must start after the metrics server and read its RunMonitoring config from /etc/rungmp." + } + + assert { + condition = alltrue([ + length([for v in google_cloud_run_v2_service.gateway[0].template[0].volumes : v if v.name == "gmp-config" && length(v.secret) == 1 && v.secret[0].items[0].path == "config.yaml"]) == 1, + google_secret_manager_secret.metrics_run_monitoring[0].secret_id == "${local.name}-gateway-run-monitoring", + google_secret_manager_secret_iam_member.metrics_run_monitoring[0].role == "roles/secretmanager.secretAccessor", + ]) + error_message = "The RunMonitoring config must be mounted from a Secret Manager secret readable by the runtime SA." + } + + assert { + condition = alltrue([ + yamldecode(google_secret_manager_secret_version.metrics_run_monitoring[0].secret_data).kind == "RunMonitoring", + yamldecode(google_secret_manager_secret_version.metrics_run_monitoring[0].secret_data).spec.endpoints[0].port == 4001, + yamldecode(google_secret_manager_secret_version.metrics_run_monitoring[0].secret_data).spec.endpoints[0].path == "/metrics", + ]) + error_message = "The RunMonitoring config must scrape /metrics on the configured metrics port." + } + + assert { + condition = alltrue([ + google_project_iam_member.runtime_metric_writer[0].role == "roles/monitoring.metricWriter", + google_project_iam_member.runtime_log_writer[0].role == "roles/logging.logWriter", + google_project_iam_member.runtime_metric_writer[0].project == "test-project", + ]) + error_message = "The runtime SA must be able to write metrics and logs for the collector sidecar." + } +} + +run "metrics_sidecar_ignored_in_deps_only" { + command = plan + + variables { + create_runtime = false + gateway_metrics_port = 4001 + } + + assert { + condition = alltrue([ + length(google_secret_manager_secret.metrics_run_monitoring) == 0, + length(google_project_iam_member.runtime_metric_writer) == 0, + ]) + error_message = "Dependencies-only mode must not create metrics sidecar resources." + } +} + +run "metrics_port_rejects_gateway_port" { + command = plan + + variables { + gateway_metrics_port = 4000 + } + + expect_failures = [var.gateway_metrics_port] +} + +run "metrics_port_rejects_collector_health_port" { + command = plan + + variables { + gateway_metrics_port = 13133 + } + + expect_failures = [var.gateway_metrics_port] +} + +run "metrics_port_rejects_fractional_port" { + command = plan + + variables { + gateway_metrics_port = 4000.5 + } + + expect_failures = [var.gateway_metrics_port] +} diff --git a/terraform/litellm/gcp/variables.tf b/terraform/litellm/gcp/variables.tf index 9c68ed3db76..a753b4584a6 100644 --- a/terraform/litellm/gcp/variables.tf +++ b/terraform/litellm/gcp/variables.tf @@ -517,6 +517,44 @@ variable "otel_capture_message_content" { } } +# ---------- Prometheus metrics sidecar ---------- + +variable "gateway_metrics_port" { + description = <<-EOT + Serve Prometheus /metrics from a `metrics` sidecar container in the + gateway Cloud Run service on this port (a whole number 1-65535, not 4000 + or 13133), so the collector's scrape never runs on an inference worker. + The sidecar runs the gateway image with + `python -m litellm.proxy.prometheus_metrics_server` and aggregates the + workers' PROMETHEUS_MULTIPROC_DIR samples over an in-memory volume shared + with the gateway container. Cloud Run only routes ingress to the gateway + container, so the sidecar port is reachable on localhost inside the + instance only; a Managed Service for Prometheus collector sidecar + (gateway_metrics_collector_image) scrapes it and writes the series to + Cloud Monitoring. The load balancer keeps serving the authenticated + /metrics on the gateway port as before. Null (the default) leaves /metrics + on the gateway port only. Needs gateway_image v1.101.0 or newer. + EOT + type = number + default = null + + validation { + condition = var.gateway_metrics_port == null || (var.gateway_metrics_port >= 1 && var.gateway_metrics_port <= 65535 && floor(var.gateway_metrics_port) == var.gateway_metrics_port && !contains([4000, 13133], var.gateway_metrics_port)) + error_message = "gateway_metrics_port must be a whole number between 1 and 65535 and must not be 4000 (the gateway port) or 13133 (the collector health port)." + } +} + +variable "gateway_metrics_collector_image" { + description = <<-EOT + Managed Service for Prometheus sidecar image that scrapes + localhost:/metrics and writes to Cloud Monitoring. + Override only to pin a different release or pull through your own + Artifact Registry. Ignored when gateway_metrics_port is null. + EOT + type = string + default = "us-docker.pkg.dev/cloud-ops-agents-artifacts/cloud-run-gmp-sidecar/cloud-run-gmp-sidecar:1.9.2" +} + # ---------- Enterprise billing metrics ---------- # # License-gated request metering. Opt-in and gated entirely on From 0e35c8fee9d81593ff733be695562a053f0ebf30 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 13:53:42 -0700 Subject: [PATCH 54/60] fix(proxy): recreate the Prisma client when the writer session turns read-only (#40610) The writer health probe only ran SELECT 1, which a read-only Postgres session answers fine, so a pooled connection left pointing at a demoted primary kept failing every write with SQLSTATE 25006 until the pod was restarted. Probe transaction_read_only instead, treat a 25006 on the request path as a signal to recreate the client, and back off exponentially while the database as a whole stays read-only so a replica or an in-progress failover does not get its engine killed every cycle. Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/exception_handler.py | 12 ++ litellm/proxy/utils.py | 86 +++++++++++-- .../proxy/db/test_exception_handler.py | 37 ++++++ .../proxy/db/test_prisma_self_heal.py | 4 +- .../proxy/test_prisma_engine_watchdog.py | 9 +- .../proxy/utils/helpers/test_error_helpers.py | 45 +++++++ .../test_prisma_client_reconnect.py | 117 +++++++++++++++++- 7 files changed, 295 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index f469587ab8e..2cee5128c66 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -221,6 +221,18 @@ class PrismaDBExceptionHandler: or "write conflict or a deadlock" in error_message ) + @staticmethod + def is_read_only_transaction_error(e: Exception) -> bool: + """True iff ``e`` is Postgres SQLSTATE 25006 surfaced through prisma: the + pooled session answers reads but rejects writes, so the connection is + poisoned until the client is recreated.""" + import prisma + + if not isinstance(e, _exception_types(prisma.errors.PrismaError)): + return False + error_message: Final = str(e).lower() + return '"25006"' in error_message or "read-only transaction" in error_message + @staticmethod def is_prisma_engine_internal_error(e: Exception) -> bool: """True iff ``e`` is a non-``PrismaError`` exception raised from inside diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 00ccad33b6d..13cba9a25ce 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -68,6 +68,7 @@ except ImportError: raise ImportError("backoff is not installed. Please install it via 'pip install backoff'") from fastapi import HTTPException, status +from pydantic import TypeAdapter import litellm import litellm.litellm_core_utils @@ -3902,6 +3903,11 @@ async def prefetch_config_params(prisma_client: "PrismaClient | None", param_nam ) +_WRITER_WRITABILITY_PROBE_SQL: Final = "SELECT current_setting('transaction_read_only') AS transaction_read_only" +_WRITER_WRITABILITY_PROBE_ROWS: Final = TypeAdapter(list[dict[str, object]]) +_READ_ONLY_RECREATE_BACKOFF_CAP_SECONDS: Final = 600 + + class _ForcedRecreateDeclined(Exception): """A forced recreate was declined by the engine-generation guard. @@ -4079,6 +4085,8 @@ class PrismaClient: self._db_health_watchdog_task: asyncio.Task | None = None self._db_last_reconnect_attempt_ts: float = 0.0 self._db_reconnect_cooldown_seconds: int = max(1, int(os.getenv("PRISMA_RECONNECT_COOLDOWN_SECONDS", "15"))) + self._db_read_only_recreate_ts: float = 0.0 + self._db_read_only_recreate_streak: int = 0 self._db_health_watchdog_interval_seconds: int = max( 5, int(os.getenv("PRISMA_HEALTH_WATCHDOG_INTERVAL_SECONDS", "30")) ) @@ -5796,15 +5804,20 @@ class PrismaClient: writer: Final = self.writer_db if force_recreate is False: try: - await writer.query_raw("SELECT 1") - verbose_proxy_logger.info( - "Writer healthy on probe; skipping recreate (engine " - "likely already replaced by a token refresh)." - ) - if isinstance(self.db, RoutingPrismaWrapper): - self.db.mark_writer_recovered() - await self._start_engine_watcher() - return + if await self._writer_is_read_only(writer): + verbose_proxy_logger.warning( + "Writer answers the probe but its session is read-only " + "(writes fail with SQLSTATE 25006); recreating Prisma client." + ) + else: + verbose_proxy_logger.info( + "Writer healthy on probe; skipping recreate (engine " + "likely already replaced by a token refresh)." + ) + if isinstance(self.db, RoutingPrismaWrapper): + self.db.mark_writer_recovered() + await self._start_engine_watcher() + return except Exception as probe_err: verbose_proxy_logger.warning( "Writer probe failed (%s); recreating Prisma client.", @@ -6124,6 +6137,18 @@ class PrismaClient: reason="db_health_watchdog_writer_unavailable", timeout_seconds=self._db_watchdog_reconnect_timeout_seconds, ) + continue + if await asyncio.wait_for( + self._writer_is_read_only(self.writer_db), + timeout=self._db_health_watchdog_probe_timeout_seconds, + ): + await self.recreate_read_only_writer( + reason="db_health_watchdog_writer_read_only", + timeout_seconds=self._db_watchdog_reconnect_timeout_seconds, + ) + continue + self._db_read_only_recreate_streak = 0 + self._db_read_only_recreate_ts = 0.0 except asyncio.CancelledError: break except Exception as e: @@ -6135,6 +6160,39 @@ class PrismaClient: else: verbose_proxy_logger.debug("Prisma DB health watchdog observed non-DB error: %s", e) + async def recreate_read_only_writer(self, reason: str, timeout_seconds: float | None = None) -> bool: + """Force-recreate the client behind a writer session that rejects writes + (SQLSTATE 25006). Each recreate doubles the wait before the next one + until the watchdog sees a writable session again, so a database that is + read-only as a whole (replica, failover in progress) does not get its + engine killed on every watchdog cycle or failed write.""" + backoff_seconds: Final = min( + self._db_reconnect_cooldown_seconds * 2 ** min(self._db_read_only_recreate_streak, 10), + _READ_ONLY_RECREATE_BACKOFF_CAP_SECONDS, + ) + if time.time() - self._db_read_only_recreate_ts < backoff_seconds: + verbose_proxy_logger.debug( + "Writer session still read-only after %s recreate(s); backing off %ss. reason=%s", + self._db_read_only_recreate_streak, + backoff_seconds, + reason, + ) + return False + verbose_proxy_logger.warning( + "Writer session is read-only (writes fail with SQLSTATE 25006); recreating Prisma client. reason=%s", + reason, + ) + self._db_read_only_recreate_ts = time.time() + self._db_read_only_recreate_streak += 1 + return await self.attempt_db_reconnect(reason=reason, timeout_seconds=timeout_seconds, force_recreate=True) + + async def _writer_is_read_only(self, writer: PrismaWrapper) -> bool: + """True iff the pooled writer session answers reads but rejects writes (SQLSTATE 25006).""" + rows: Final = _WRITER_WRITABILITY_PROBE_ROWS.validate_python( + await writer.query_raw(_WRITER_WRITABILITY_PROBE_SQL) + ) + return any(row.get("transaction_read_only") == "on" for row in rows) + def _probe_target_wrapper(self) -> PrismaWrapper: """The Prisma wrapper a `SELECT 1` health probe actually reaches. @@ -7547,6 +7605,12 @@ def _get_openapi_url() -> str | None: return "/openapi.json" +def _recreate_writer_on_read_only_transaction(prisma_client: "PrismaClient | None") -> None: + if prisma_client is None: + return + asyncio.create_task(prisma_client.recreate_read_only_writer(reason="postgres_read_only_transaction")) + + def handle_exception_on_proxy(e: Exception) -> ProxyException: """ Returns an Exception as ProxyException, this ensures all exceptions are OpenAI API compatible @@ -7554,6 +7618,10 @@ def handle_exception_on_proxy(e: Exception) -> ProxyException: from fastapi import status verbose_proxy_logger.exception("Exception: %s", e) + if PrismaDBExceptionHandler.is_read_only_transaction_error(e): + from litellm.proxy.proxy_server import prisma_client + + _recreate_writer_on_read_only_transaction(prisma_client) if isinstance(e, HTTPException): return ProxyException( diff --git a/tests/test_litellm/proxy/db/test_exception_handler.py b/tests/test_litellm/proxy/db/test_exception_handler.py index c685d778c0e..3f009137a1c 100644 --- a/tests/test_litellm/proxy/db/test_exception_handler.py +++ b/tests/test_litellm/proxy/db/test_exception_handler.py @@ -665,10 +665,47 @@ def test_is_deadlock_error_excludes_non_deadlocks(error): assert PrismaDBExceptionHandler.is_deadlock_error(error) is False +READ_ONLY_CONNECTOR_ERROR: Final = ( + "Error occurred during query execution:\nConnectorError(ConnectorError { user_facing_error: None, " + 'kind: QueryError(PostgresError { code: "25006", message: "cannot execute UPDATE in a read-only transaction", ' + 'severity: "ERROR", detail: None, column: None, hint: None }), transient: false })' +) + + +@pytest.mark.parametrize( + "error", + [ + DataError(data={"user_facing_error": {"message": READ_ONLY_CONNECTOR_ERROR}}), + RawQueryError(data={"user_facing_error": {"message": "cannot execute INSERT in a read-only transaction"}}), + PrismaError( + 'PostgresError { code: "25006", message: "kann DELETE in einer Read-Only-Transaktion nicht ausführen" }' + ), + ], +) +def test_is_read_only_transaction_error_matches_sqlstate_25006(error): + assert PrismaDBExceptionHandler.is_read_only_transaction_error(error) is True + + +@pytest.mark.parametrize( + "error", + [ + UniqueViolationError(data={"user_facing_error": {"error_code": "P2002", "meta": {"table": "t"}}}), + PrismaError("can't reach database server"), + RawQueryError(data={"user_facing_error": {"message": "deadlock detected", "meta": {"table": "t"}}}), + httpx.ConnectError("connection refused"), + RuntimeError("cannot execute UPDATE in a read-only transaction"), + ValueError('"25006"'), + ], +) +def test_is_read_only_transaction_error_excludes_other_failures(error): + assert PrismaDBExceptionHandler.is_read_only_transaction_error(error) is False + + MOCKED_PRISMA_PREDICATES: Final = ( PrismaDBExceptionHandler.is_database_infrastructure_error, PrismaDBExceptionHandler.is_database_transport_error, PrismaDBExceptionHandler.is_deadlock_error, + PrismaDBExceptionHandler.is_read_only_transaction_error, PrismaDBExceptionHandler.is_prisma_engine_internal_error, PrismaDBExceptionHandler.is_database_service_unavailable_error, ) diff --git a/tests/test_litellm/proxy/db/test_prisma_self_heal.py b/tests/test_litellm/proxy/db/test_prisma_self_heal.py index 10a48941693..e6914213939 100644 --- a/tests/test_litellm/proxy/db/test_prisma_self_heal.py +++ b/tests/test_litellm/proxy/db/test_prisma_self_heal.py @@ -626,7 +626,7 @@ async def test_direct_reconnect_probe_success_clears_writer_unavailable( database_url="mock://test", proxy_logging_obj=mock_proxy_logging ) writer = MagicMock() - writer.query_raw = AsyncMock(return_value=[{"result": 1}]) + writer.query_raw = AsyncMock(return_value=[{"transaction_read_only": "off"}]) reader = MagicMock() routing = RoutingPrismaWrapper(writer=writer, reader=reader) routing._writer_unavailable = True @@ -636,5 +636,5 @@ async def test_direct_reconnect_probe_success_clears_writer_unavailable( with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}): await client._run_reconnect_cycle(timeout_seconds=5.0) - writer.query_raw.assert_awaited_once_with("SELECT 1") + writer.query_raw.assert_awaited_once_with("SELECT current_setting('transaction_read_only') AS transaction_read_only") assert routing.writer_unavailable is False diff --git a/tests/test_litellm/proxy/test_prisma_engine_watchdog.py b/tests/test_litellm/proxy/test_prisma_engine_watchdog.py index d73f74c5cd2..ae750b0770a 100644 --- a/tests/test_litellm/proxy/test_prisma_engine_watchdog.py +++ b/tests/test_litellm/proxy/test_prisma_engine_watchdog.py @@ -18,12 +18,15 @@ import asyncio import os import threading import time +from typing import Final from unittest.mock import ANY, AsyncMock, MagicMock, patch import pytest from litellm.proxy.utils import PrismaClient, ProxyLogging +WRITER_PROBE_SQL: Final = "SELECT current_setting('transaction_read_only') AS transaction_read_only" + @pytest.fixture(autouse=True) def mock_prisma_binary(): @@ -260,7 +263,7 @@ async def test_run_reconnect_cycle_uses_direct_path_when_engine_alive( """Direct reconnect (engine alive) probes the writer first and skips the recreate when the probe is healthy. - The engine-alive path now runs a SELECT 1 probe before recreating. A + The engine-alive path now runs a writability probe before recreating. A healthy probe means the connection is fine — e.g. an IAM token refresh already replaced the engine (issue #29176) — so recreating would kill a working engine. Recreate happens only when the probe fails (covered in @@ -278,7 +281,7 @@ async def test_run_reconnect_cycle_uses_direct_path_when_engine_alive( await engine_client._run_reconnect_cycle(timeout_seconds=5.0) engine_client.db.recreate_prisma_client.assert_not_awaited() - engine_client.db.query_raw.assert_awaited_once_with("SELECT 1") + engine_client.db.query_raw.assert_awaited_once_with(WRITER_PROBE_SQL) engine_client.db.disconnect.assert_not_awaited() engine_client._start_engine_watcher.assert_awaited_once() @@ -296,7 +299,7 @@ async def test_run_reconnect_cycle_uses_direct_path_when_pid_unknown( await engine_client._run_reconnect_cycle(timeout_seconds=5.0) engine_client.db.recreate_prisma_client.assert_not_awaited() - engine_client.db.query_raw.assert_awaited_once_with("SELECT 1") + engine_client.db.query_raw.assert_awaited_once_with(WRITER_PROBE_SQL) engine_client.db.disconnect.assert_not_awaited() engine_client._start_engine_watcher.assert_awaited_once() diff --git a/tests/test_litellm/proxy/utils/helpers/test_error_helpers.py b/tests/test_litellm/proxy/utils/helpers/test_error_helpers.py index dc30798df55..117c5aa3081 100644 --- a/tests/test_litellm/proxy/utils/helpers/test_error_helpers.py +++ b/tests/test_litellm/proxy/utils/helpers/test_error_helpers.py @@ -1,7 +1,10 @@ +import asyncio import json +from unittest.mock import AsyncMock, MagicMock import pytest from fastapi import HTTPException +from prisma.errors import DataError from litellm.proxy._types import ProxyErrorTypes, ProxyException from litellm.proxy.utils import get_error_message_str, handle_exception_on_proxy @@ -171,3 +174,45 @@ def test_handle_exception_on_proxy_error_path_none_input_wraps_as_500(): "code": "500", "type": ProxyErrorTypes.internal_server_error.value, } + + +@pytest.mark.asyncio +async def test_handle_exception_on_proxy_read_only_transaction_forces_writer_recreate( + monkeypatch: pytest.MonkeyPatch, +) -> None: + prisma_client = MagicMock() + prisma_client.recreate_read_only_writer = AsyncMock(return_value=True) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + exc = DataError( + data={ + "user_facing_error": { + "message": 'PostgresError { code: "25006", message: "cannot execute UPDATE in a read-only transaction" }' + } + } + ) + + result = handle_exception_on_proxy(exc) + await asyncio.sleep(0) + + snapshot = { + "code": result.code, + "recreate_kwargs": prisma_client.recreate_read_only_writer.await_args.kwargs, + } + assert snapshot == { + "code": "500", + "recreate_kwargs": {"reason": "postgres_read_only_transaction"}, + } + + +@pytest.mark.asyncio +async def test_handle_exception_on_proxy_leaves_writer_alone_for_other_db_errors( + monkeypatch: pytest.MonkeyPatch, +) -> None: + prisma_client = MagicMock() + prisma_client.recreate_read_only_writer = AsyncMock(return_value=True) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + + handle_exception_on_proxy(DataError(data={"user_facing_error": {"message": "deadlock detected"}})) + await asyncio.sleep(0) + + assert prisma_client.recreate_read_only_writer.await_count == 0 diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py index 41b0eb3cf95..cec6dd99ce8 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py @@ -29,7 +29,7 @@ from __future__ import annotations import asyncio import time from typing import Any, Final -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, call import pytest @@ -111,6 +111,34 @@ async def test_run_reconnect_cycle_direct_path_recreates_when_probe_fails( } +@pytest.mark.asyncio +async def test_run_reconnect_cycle_direct_path_recreates_when_writer_is_read_only( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("DATABASE_URL", "postgres://x:y@h:5432/db") + prisma_client._engine_confirmed_dead = False + prisma_client._engine_pid = 0 + prisma_client._start_engine_watcher = AsyncMock() + prisma_client._cleanup_engine_watcher = MagicMock() + + writer: Final = MagicMock() + writer.query_raw = AsyncMock(side_effect=[[{"transaction_read_only": "on"}], [{"?column?": 1}]]) + writer.recreate_prisma_client = AsyncMock() + prisma_client.db = writer + + await prisma_client._run_reconnect_cycle(timeout_seconds=5) + pinned = { + "recreate_called": writer.recreate_prisma_client.await_count, + "start_watcher_called": prisma_client._start_engine_watcher.await_count, + "cleanup_called": prisma_client._cleanup_engine_watcher.call_count, + } + assert pinned == { + "recreate_called": 1, + "start_watcher_called": 1, + "cleanup_called": 1, + } + + @pytest.mark.asyncio async def test_run_reconnect_cycle_force_recreate_skips_probe_and_recreates( prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch @@ -488,6 +516,93 @@ async def test_db_health_watchdog_loop_swallows_non_db_errors( assert prisma_client.attempt_db_reconnect.await_count == 0 +def _routing_db_with_writer_sessions(*transaction_read_only: str) -> tuple[RoutingPrismaWrapper, MagicMock]: + """One watchdog cycle per value, then the loop is cancelled.""" + writer: Final = MagicMock() + writer.query_raw = AsyncMock(side_effect=[[{"transaction_read_only": value}] for value in transaction_read_only]) + reader: Final = MagicMock() + reader.query_raw = AsyncMock( + side_effect=[[{"?column?": 1}] for _ in transaction_read_only] + [asyncio.CancelledError()] + ) + return RoutingPrismaWrapper(writer=writer, reader=reader), writer + + +@pytest.mark.asyncio +async def test_db_health_watchdog_loop_forces_recreate_when_writer_is_read_only( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_health_watchdog_interval_seconds = 0 + prisma_client.attempt_db_reconnect = AsyncMock(side_effect=asyncio.CancelledError()) + prisma_client.db, _ = _routing_db_with_writer_sessions("on") + + await prisma_client._db_health_watchdog_loop() + assert prisma_client.attempt_db_reconnect.await_args is not None + assert prisma_client.attempt_db_reconnect.await_args.kwargs == { + "reason": "db_health_watchdog_writer_read_only", + "timeout_seconds": prisma_client._db_watchdog_reconnect_timeout_seconds, + "force_recreate": True, + } + + +@pytest.mark.asyncio +async def test_db_health_watchdog_loop_leaves_writable_writer_alone( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_health_watchdog_interval_seconds = 0 + prisma_client.attempt_db_reconnect = AsyncMock() + prisma_client.db, writer = _routing_db_with_writer_sessions("off") + + await prisma_client._db_health_watchdog_loop() + assert (prisma_client.attempt_db_reconnect.await_count, writer.query_raw.await_count) == (0, 1) + + +@pytest.mark.asyncio +async def test_db_health_watchdog_loop_backs_off_while_database_stays_read_only( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_health_watchdog_interval_seconds = 0 + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + prisma_client.db, writer = _routing_db_with_writer_sessions("on", "on", "on") + + await prisma_client._db_health_watchdog_loop() + assert (prisma_client.attempt_db_reconnect.await_count, writer.query_raw.await_count) == (1, 3) + + +@pytest.mark.asyncio +async def test_db_health_watchdog_loop_recreates_again_once_writer_was_writable_in_between( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_health_watchdog_interval_seconds = 0 + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + prisma_client.db, _ = _routing_db_with_writer_sessions("on", "off", "on") + + await prisma_client._db_health_watchdog_loop() + assert prisma_client.attempt_db_reconnect.await_count == 2 + + +@pytest.mark.asyncio +async def test_recreate_read_only_writer_retries_after_backoff_elapses( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_reconnect_cooldown_seconds = 15 + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + + first: Final = await prisma_client.recreate_read_only_writer(reason="postgres_read_only_transaction") + within_backoff: Final = await prisma_client.recreate_read_only_writer(reason="postgres_read_only_transaction") + prisma_client._db_read_only_recreate_ts -= 30 + after_backoff: Final = await prisma_client.recreate_read_only_writer(reason="postgres_read_only_transaction") + prisma_client._db_read_only_recreate_ts -= 30 + still_within_doubled_backoff: Final = await prisma_client.recreate_read_only_writer( + reason="postgres_read_only_transaction" + ) + + assert (first, within_backoff, after_backoff, still_within_doubled_backoff) == (True, False, True, False) + assert prisma_client.attempt_db_reconnect.await_args_list == [ + call(reason="postgres_read_only_transaction", timeout_seconds=None, force_recreate=True), + call(reason="postgres_read_only_transaction", timeout_seconds=None, force_recreate=True), + ] + + @pytest.mark.asyncio async def test_iam_refresh_racing_reconnect_recreates_engine_only_once( prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch From 46a185d3cde39019a510137c5459c71f8a3cf424 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 13:56:30 -0700 Subject: [PATCH 55/60] feat(rust_bridge): count budget-check input tokens in Rust on all LLM routes (#40381) * feat(rust_bridge): count budget-check input tokens in Rust on all LLM routes Rust counts input tokens from the raw JSON body with the GIL released inside the existing budget reservation, covering every LLM route the auth dependency guards. It only fires for models on the Anthropic tokenizer when a budget is set, and Python counts whenever Rust is off, missing, or declines a body shape. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * perf(rust): count byte-level BPE tokens without the GPT-2 split regex (#40594) The oniguruma run of the ByteLevel pre-tokenizer regex is about 90% of encode_fast on a 100k token body (100 ms of the ~110 ms Rust admission count in the gateway pod). A hand-written scanner that yields the same pieces, then feeds the model directly, counts the same text in 10 ms. It only engages for tokenizers with the Anthropic shape (optional NFKC, ByteLevel without prefix space, no post-processor) and falls back to the full encoder when the text contains an added token. Parity with encode_fast is tested on random texts, the pieces are compared with the real pre-tokenizer, and the \p{L}/\p{N}/\s tables are checked against oniguruma for every code point. NFKC runs through unicode-normalization-alignments, the crate and Unicode tables NormalizedString::nfkc already uses, so the fast path normalizes exactly what the full encoder would. Using the newer unicode-normalization crate changed the count for 171 code points that gained compatibility decompositions after Unicode 9 (U+32FF, U+A7F1..). The fast normalizer is compared with the tokenizer's for every scalar value and on random texts. The scanner is built without mutable state: byte_char and mapped_len replace the const table builders and the reusable mapped buffer, and iter::successors replaces the stateful piece iterator. byte_chars_match_the_byte_level_alphabet checks the byte mapping against ByteLevel for every scalar value. Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust_bridge): bound concurrent token-count encodes and share the Anthropic tokenizer predicate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/AGENTS.md | 5 +- litellm-rust/Cargo.lock | 422 +++++++++++- litellm-rust/Cargo.toml | 3 + litellm-rust/crates/ai-gateway/README.md | 5 +- .../core/tests/workspace_crate_allowlist.rs | 7 +- litellm-rust/crates/python-bridge/Cargo.toml | 5 +- .../crates/python-bridge/src/constants.rs | 2 + .../crates/python-bridge/src/execution.rs | 39 +- litellm-rust/crates/python-bridge/src/lib.rs | 4 + .../crates/python-bridge/src/token_counter.rs | 87 +++ litellm-rust/crates/token-counter/Cargo.toml | 28 + .../token-counter/benches/allocations.rs | 128 ++++ .../token-counter/benches/token_counter.rs | 100 +++ .../crates/token-counter/src/byte_level.rs | 650 ++++++++++++++++++ .../crates/token-counter/src/counter.rs | 194 ++++++ .../crates/token-counter/src/error.rs | 31 + litellm-rust/crates/token-counter/src/lib.rs | 17 + .../crates/token-counter/src/python_json.rs | 155 +++++ .../crates/token-counter/src/tools.rs | 304 ++++++++ .../crates/token-counter/src/types.rs | 238 +++++++ .../token-counter/src/unicode_classes.rs | 73 ++ .../token-counter/tests/token_counter.rs | 176 +++++ litellm/proxy/auth/user_api_key_auth.py | 4 + .../proxy/common_utils/http_parsing_utils.py | 12 + .../spend_tracking/budget_reservation.py | 41 +- litellm/rust_bridge/token_counter.py | 79 +++ litellm/utils.py | 6 +- .../common_utils/test_http_parsing_utils.py | 49 ++ .../spend_tracking/test_budget_reservation.py | 136 +++- .../rust_bridge/test_token_counter.py | 242 +++++++ 30 files changed, 3196 insertions(+), 46 deletions(-) create mode 100644 litellm-rust/crates/python-bridge/src/constants.rs create mode 100644 litellm-rust/crates/python-bridge/src/token_counter.rs create mode 100644 litellm-rust/crates/token-counter/Cargo.toml create mode 100644 litellm-rust/crates/token-counter/benches/allocations.rs create mode 100644 litellm-rust/crates/token-counter/benches/token_counter.rs create mode 100644 litellm-rust/crates/token-counter/src/byte_level.rs create mode 100644 litellm-rust/crates/token-counter/src/counter.rs create mode 100644 litellm-rust/crates/token-counter/src/error.rs create mode 100644 litellm-rust/crates/token-counter/src/lib.rs create mode 100644 litellm-rust/crates/token-counter/src/python_json.rs create mode 100644 litellm-rust/crates/token-counter/src/tools.rs create mode 100644 litellm-rust/crates/token-counter/src/types.rs create mode 100644 litellm-rust/crates/token-counter/src/unicode_classes.rs create mode 100644 litellm-rust/crates/token-counter/tests/token_counter.rs create mode 100644 litellm/rust_bridge/token_counter.py create mode 100644 tests/test_litellm/rust_bridge/test_token_counter.py diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index b8d1f2db4d7..17856218e60 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -1,18 +1,19 @@ # AGENTS.md -litellm-rust has five crates. A crate is a layer or shared foundation, not a route. Routes (ocr, realtime, chat) and providers (mistral, openai) are modules inside the layers. +litellm-rust has six crates. A crate is a layer or shared foundation, not a route. Routes (ocr, realtime, chat) and providers (mistral, openai) are modules inside the layers. ## Crates | Crate | Role | |-------|------| | litellm-core | The LiteLLM SDK in Rust. One public entrypoint per top-level call (`messages::messages()`), owning types, transforms, provider resolution, auth, and the provider HTTP call. Call it, get a typed response. | +| litellm-token-counter | Standalone input token counting shared by host integrations without pulling in the full SDK. | | litellm-config | Config-loading boundary. Returns resolved core deployment data and optionally delegates loading to Python. | | litellm-ai-gateway | The axum server (behind the `server` feature) plus the WebSocket hosts. Translates HTTP/WS to core entrypoints; owns no provider logic and no handlers. | | litellm-python-interop | Domain-neutral PyO3 foundation for GIL handling and typed Python/Serde conversion. | | litellm-python-bridge | PyO3 cdylib exposing LiteLLM Rust APIs to the Python SDK. Owns API registration, domain wiring, and Python exception mapping. | -Dependency direction is acyclic: `litellm-config` depends on `litellm-core`, the gateway depends on both, and `litellm-python-bridge` depends on the domain layers and `litellm-python-interop`. The interop foundation depends on no LiteLLM domain crate. +Dependency direction is acyclic: `litellm-config` depends on `litellm-core`, the gateway depends on both, and `litellm-python-bridge` depends on the domain layers, `litellm-token-counter`, and `litellm-python-interop`. The token counter and interop foundations depend on no LiteLLM domain crate. ## Where a route lives diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 7addc4e3e45..0e8e6e09a21 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2,6 +2,20 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "ahash" +version = "0.8.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a15f179cd60c4584b8a8c596927aadc462e27f2ca70c04e0071964a73ba7a75" +dependencies = [ + "cfg-if", + "getrandom 0.3.4", + "once_cell", + "serde", + "version_check", + "zerocopy", +] + [[package]] name = "aho-corasick" version = "1.1.5" @@ -418,7 +432,7 @@ checksum = "edca88bc138befd0323b20752846e6587272d3b03b0343c8ea28a6f819e6e71f" dependencies = [ "async-trait", "axum-core", - "base64", + "base64 0.22.1", "bytes", "futures-util", "http 1.4.2", @@ -468,6 +482,12 @@ dependencies = [ "tracing", ] +[[package]] +name = "base64" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e1b586273c5702936fe7b7d6896644d8be71e6314cfe09d3167c95f712589e8" + [[package]] name = "base64" version = "0.22.1" @@ -542,6 +562,15 @@ version = "0.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5" +[[package]] +name = "castaway" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dec551ab6e7578819132c713a93c022a05d60159dc86e7a7050223577484c55a" +dependencies = [ + "rustversion", +] + [[package]] name = "cc" version = "1.3.0" @@ -644,6 +673,21 @@ version = "0.5.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0c9ea0ac24bc397ab3c98583a3c9ba74fa56b09a4449bbe172b9b1ddb016027a" +[[package]] +name = "compact_str" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9dfdd1c2274d9aa354115b09dc9a901d6c5576818cdf70d14cae2bdb47df00ab" +dependencies = [ + "castaway", + "cfg-if", + "itoa", + "rustversion", + "ryu", + "serde", + "static_assertions", +] + [[package]] name = "const-oid" version = "0.10.2" @@ -696,7 +740,7 @@ dependencies = [ "ciborium", "clap", "criterion-plot", - "itertools", + "itertools 0.13.0", "num-traits", "oorandom", "page_size", @@ -716,7 +760,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d8d80a2f4f5b554395e47b5d8305bc3d27813bacb73493eb1001e8f76dae29ea" dependencies = [ "cast", - "itertools", + "itertools 0.13.0", ] [[package]] @@ -778,6 +822,56 @@ dependencies = [ "cmov", ] +[[package]] +name = "daachorse" +version = "3.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5614204febbc33cc07a2806aa6440b904ac012b68eecc37f4493ea4a76455a3d" + +[[package]] +name = "darling" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc7f46116c46ff9ab3eb1597a45688b6715c6e628b5c133e288e709a29bcb4ee" +dependencies = [ + "darling_core", + "darling_macro", +] + +[[package]] +name = "darling_core" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d00b9596d185e565c2207a0b01f8bd1a135483d02d9b7b0a54b11da8d53412e" +dependencies = [ + "fnv", + "ident_case", + "proc-macro2", + "quote", + "strsim", + "syn 2.0.119", +] + +[[package]] +name = "darling_macro" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc34b93ccb385b40dc71c6fceac4b2ad23662c7eeb248cf10d529b7e055b6ead" +dependencies = [ + "darling_core", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "dary_heap" +version = "0.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b1e3a325bc115f096c8b77bbf027a7c2592230e70be2d985be950d3d5e60ebe" +dependencies = [ + "serde", +] + [[package]] name = "data-encoding" version = "2.11.0" @@ -790,6 +884,37 @@ version = "0.5.8" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7cd812cc2bc1d69d4764bd80df88b4317eaef9e773c75226407d9bc0876b211c" +[[package]] +name = "derive_builder" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "507dfb09ea8b7fa618fcf76e953f4f5e192547945816d5358edffe39f6f94947" +dependencies = [ + "derive_builder_macro", +] + +[[package]] +name = "derive_builder_core" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d5bcf7b024d6835cfb3d473887cd966994907effbe9227e8c8219824d06c4e8" +dependencies = [ + "darling", + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "derive_builder_macro" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab63b0e2bf4d5928aff72e83a7dace85d7bba5fe12dcc3c5a572d78caffd3f3c" +dependencies = [ + "derive_builder_core", + "syn 2.0.119", +] + [[package]] name = "digest" version = "0.10.7" @@ -841,6 +966,12 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" +[[package]] +name = "esaxx-rs" +version = "0.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d817e038c30374a4bcb22f94d0a8a0e216958d4c3dcde369b1439fec4bdda6e6" + [[package]] name = "fastrand" version = "2.5.0" @@ -964,6 +1095,18 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "getrandom" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" +dependencies = [ + "cfg-if", + "libc", + "r-efi 5.3.0", + "wasip2", +] + [[package]] name = "getrandom" version = "0.4.3" @@ -973,7 +1116,7 @@ dependencies = [ "cfg-if", "js-sys", "libc", - "r-efi", + "r-efi 6.0.0", "rand_core 0.10.1", "wasm-bindgen", ] @@ -1220,7 +1363,7 @@ version = "0.1.20" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0" dependencies = [ - "base64", + "base64 0.22.1", "bytes", "futures-channel", "futures-util", @@ -1319,6 +1462,12 @@ dependencies = [ "zerovec", ] +[[package]] +name = "ident_case" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9e0384b61958566e926dc50660321d12159025e767c18e043daf26b70104c39" + [[package]] name = "idna" version = "1.1.0" @@ -1348,6 +1497,8 @@ checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" dependencies = [ "equivalent", "hashbrown", + "serde", + "serde_core", ] [[package]] @@ -1365,6 +1516,15 @@ dependencies = [ "either", ] +[[package]] +name = "itertools" +version = "0.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2b192c782037fadd9cfa75548310488aabdbf3d2da73885b31bd0abd03351285" +dependencies = [ + "either", +] + [[package]] name = "itoa" version = "1.0.18" @@ -1409,7 +1569,7 @@ name = "litellm-ai-gateway" version = "0.1.0" dependencies = [ "axum", - "base64", + "base64 0.22.1", "futures-channel", "futures-util", "litellm-config", @@ -1447,7 +1607,7 @@ dependencies = [ "aws-sigv4", "aws-smithy-runtime-api", "aws-types", - "base64", + "base64 0.22.1", "rand 0.8.7", "reqwest", "rstest", @@ -1471,6 +1631,7 @@ dependencies = [ "litellm-ai-gateway", "litellm-core", "litellm-python-interop", + "litellm-token-counter", "pyo3", "pyo3-async-runtimes", "serde", @@ -1491,6 +1652,22 @@ dependencies = [ "serde_json", ] +[[package]] +name = "litellm-token-counter" +version = "0.1.0" +dependencies = [ + "criterion", + "indexmap", + "itoa", + "rand 0.8.7", + "rstest", + "serde", + "serde_json", + "thiserror 2.0.19", + "tokenizers", + "unicode-normalization-alignments", +] + [[package]] name = "litemap" version = "0.8.2" @@ -1509,6 +1686,22 @@ version = "0.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154" +[[package]] +name = "macro_rules_attribute" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b3ae8f6d608c795738406608304d30a2dfbdc8e58e44f7ba43236da5208ded3c" +dependencies = [ + "macro_rules_attribute-proc_macro", + "pastey", +] + +[[package]] +name = "macro_rules_attribute-proc_macro" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc04a4c58212d57930a24bf47d3fa87485264a3a054e9c10e042eb373573ad3c" + [[package]] name = "matchit" version = "0.7.3" @@ -1537,6 +1730,12 @@ dependencies = [ "unicase", ] +[[package]] +name = "minimal-lexical" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" + [[package]] name = "mio" version = "1.2.2" @@ -1548,6 +1747,38 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "monostate" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3341a273f6c9d5bef1908f17b7267bbab0e95c9bf69a0d4dcf8e9e1b2c76ef67" +dependencies = [ + "monostate-impl", + "serde", + "serde_core", +] + +[[package]] +name = "monostate-impl" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e4db6d5580af57bf992f59068d4ea26fd518574ff48d7639b255a36f9de6e7e9" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "nom" +version = "7.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d273983c5a657a70a3e8f2a01329822f3b8c8172b73826411a55751e404a0a4a" +dependencies = [ + "memchr", + "minimal-lexical", +] + [[package]] name = "num-conv" version = "0.2.2" @@ -1578,6 +1809,28 @@ version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" +[[package]] +name = "onig" +version = "6.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0cc3cbf698f9438986c11a880c90a6d04b9de27575afd28bbf45b154b6c709e2" +dependencies = [ + "bitflags", + "libc", + "once_cell", + "onig_sys", +] + +[[package]] +name = "onig_sys" +version = "69.9.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e68317604e77e53b85896388e1a803c1d21b74c899ec9e5e1112db90735edd7" +dependencies = [ + "cc", + "pkg-config", +] + [[package]] name = "oorandom" version = "11.1.5" @@ -1606,6 +1859,18 @@ dependencies = [ "winapi", ] +[[package]] +name = "paste" +version = "1.0.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" + +[[package]] +name = "pastey" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2ee67f1008b1ba2321834326597b8e186293b049a023cdef258527550b9935b4" + [[package]] name = "percent-encoding" version = "2.3.2" @@ -1852,6 +2117,12 @@ dependencies = [ "proc-macro2", ] +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + [[package]] name = "r-efi" version = "6.0.0" @@ -1865,10 +2136,20 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "22f6172bdec972074665ed81ed53b71da00bfc44b65a753cfde883ec4c702a1a" dependencies = [ "libc", - "rand_chacha", + "rand_chacha 0.3.1", "rand_core 0.6.4", ] +[[package]] +name = "rand" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9ef1d0d795eb7d84685bca4f72f3649f064e6641543d3a8c415898726a57b41" +dependencies = [ + "rand_chacha 0.9.0", + "rand_core 0.9.5", +] + [[package]] name = "rand" version = "0.10.2" @@ -1890,6 +2171,16 @@ dependencies = [ "rand_core 0.6.4", ] +[[package]] +name = "rand_chacha" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" +dependencies = [ + "ppv-lite86", + "rand_core 0.9.5", +] + [[package]] name = "rand_core" version = "0.6.4" @@ -1899,6 +2190,15 @@ dependencies = [ "getrandom 0.2.17", ] +[[package]] +name = "rand_core" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c" +dependencies = [ + "getrandom 0.3.4", +] + [[package]] name = "rand_core" version = "0.10.1" @@ -1924,6 +2224,17 @@ dependencies = [ "rayon-core", ] +[[package]] +name = "rayon-cond" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2964d0cf57a3e7a06e8183d14a8b527195c706b7983549cd5462d5aa3747438f" +dependencies = [ + "either", + "itertools 0.14.0", + "rayon", +] + [[package]] name = "rayon-core" version = "1.13.0" @@ -1981,7 +2292,7 @@ version = "0.12.28" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" dependencies = [ - "base64", + "base64 0.22.1", "bytes", "futures-channel", "futures-core", @@ -2363,12 +2674,36 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "spm_precompiled" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5851699c4033c63636f7ea4cf7b7c1f1bf06d0cc03cfb42e711de5a5c46cf326" +dependencies = [ + "base64 0.13.1", + "nom", + "serde", + "unicode-segmentation", +] + [[package]] name = "stable_deref_trait" version = "1.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" +[[package]] +name = "static_assertions" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2eb9349b6444b326872e140eb1cf5e7c522154d69e7a0ffb0fb81c06b37543f" + +[[package]] +name = "strsim" +version = "0.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" + [[package]] name = "subtle" version = "2.6.1" @@ -2537,6 +2872,39 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" +[[package]] +name = "tokenizers" +version = "0.23.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7afbf6e88718afcc138bad01d6ccc3051dbbc3b2ce9793d8b8a3aeb610969cfc" +dependencies = [ + "ahash", + "compact_str", + "daachorse", + "dary_heap", + "derive_builder", + "esaxx-rs", + "getrandom 0.3.4", + "itertools 0.14.0", + "log", + "macro_rules_attribute", + "monostate", + "onig", + "paste", + "rand 0.9.5", + "rayon", + "rayon-cond", + "regex", + "regex-syntax", + "serde", + "serde_json", + "spm_precompiled", + "thiserror 2.0.19", + "unicode-normalization-alignments", + "unicode-segmentation", + "unicode_categories", +] + [[package]] name = "tokio" version = "1.53.0" @@ -2775,6 +3143,27 @@ version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" +[[package]] +name = "unicode-normalization-alignments" +version = "0.1.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "43f613e4fa046e69818dd287fdc4bc78175ff20331479dab6e1b0f98d57062de" +dependencies = [ + "smallvec", +] + +[[package]] +name = "unicode-segmentation" +version = "1.13.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6f5d3c3b1bf09027a88a6bc961fc00497d651009560b5463668dc81b0fa87a8" + +[[package]] +name = "unicode_categories" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e" + [[package]] name = "untrusted" version = "0.9.0" @@ -2858,6 +3247,15 @@ version = "0.11.1+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" +[[package]] +name = "wasip2" +version = "1.0.4+wasi-0.2.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b67efb37e106e55ce722a510d6b5f9c17f083e5fc79afc2badeb12cc313d9487" +dependencies = [ + "wit-bindgen", +] + [[package]] name = "wasm-bindgen" version = "0.2.126" @@ -3083,6 +3481,12 @@ dependencies = [ "memchr", ] +[[package]] +name = "wit-bindgen" +version = "0.57.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" + [[package]] name = "writeable" version = "0.6.3" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 00fa6faaaef..f3e54e5b2aa 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -1,6 +1,7 @@ [workspace] members = [ "crates/core", + "crates/token-counter", "crates/config", "crates/ai-gateway", "crates/python-interop", @@ -18,6 +19,7 @@ repository = "https://github.com/BerriAI/litellm" tracing = "0.1" tracing-subscriber = { version = "0.3", default-features = false, features = ["registry", "std"] } litellm-core = { path = "crates/core" } +litellm-token-counter = { path = "crates/token-counter" } litellm-config = { path = "crates/config" } litellm-ai-gateway = { path = "crates/ai-gateway", default-features = false } litellm-python-interop = { path = "crates/python-interop" } @@ -40,6 +42,7 @@ tokio-tungstenite = { version = "0.24", default-features = false, features = ["c futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] } base64 = "0.22" url = "2.5.8" +criterion = "0.8.2" [profile.release] opt-level = 3 diff --git a/litellm-rust/crates/ai-gateway/README.md b/litellm-rust/crates/ai-gateway/README.md index 9fef59a277d..cbcd8119546 100644 --- a/litellm-rust/crates/ai-gateway/README.md +++ b/litellm-rust/crates/ai-gateway/README.md @@ -6,17 +6,18 @@ dials OpenAI upstream, and splices the two sockets frame-by-frame. ## Crates -`litellm-rust` has five crates. A crate is a layer or shared foundation, not a route: +`litellm-rust` has six crates. A crate is a layer or shared foundation, not a route: | Crate | Role | |-------|------| | litellm-core | The LiteLLM SDK in Rust — per-route entrypoints (`messages::messages()`) that resolve the provider, transform, and make the call; plus types, provider transforms, and the router. | +| litellm-token-counter | Standalone input token counting shared by host integrations without pulling in the full SDK. | | litellm-config | Config-loading boundary. Returns resolved deployments and optionally delegates loading to Python. | | litellm-ai-gateway | The Axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. | | litellm-python-interop | Domain-neutral PyO3 foundation for GIL handling and typed Python/Serde conversion. | | litellm-python-bridge | PyO3 cdylib exposing LiteLLM Rust APIs to the Python SDK. | -Dependency direction is acyclic: config depends on core, the gateway depends on config and core, and the Python bridge depends on the domain layers and Python interop. +Dependency direction is acyclic: config depends on core, the gateway depends on config and core, and the Python bridge depends on the domain layers, token counter, and Python interop. - **Client endpoint:** `wss:///v1/realtime?model=` (WebSocket) - **Auth:** `Authorization: Bearer $LITELLM_MASTER_KEY` (fails closed if unset) diff --git a/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs b/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs index fc0ab2b62a3..e7739fe7312 100644 --- a/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs +++ b/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs @@ -1,6 +1,7 @@ -//! Enforcement: the litellm-rust workspace has exactly five crates. +//! Enforcement: the litellm-rust workspace has exactly six crates. //! -//! `core` (the Rust SDK), `config` (the config-loading boundary), +//! `core` (the Rust SDK), `token-counter` (standalone input token counting), +//! `config` (the config-loading boundary), //! `ai-gateway` (the HTTP/WebSocket host), //! `python-interop` (domain-neutral PyO3 primitives), and `python-bridge` (the //! PyO3 cdylib). Adding or removing a crate must be a @@ -20,6 +21,7 @@ use std::path::{Path, PathBuf}; /// workspace legitimately gains or loses a crate. const EXPECTED_MEMBERS: &[&str] = &[ "crates/core", + "crates/token-counter", "crates/config", "crates/ai-gateway", "crates/python-interop", @@ -29,6 +31,7 @@ const EXPECTED_MEMBERS: &[&str] = &[ /// The crate subdirectory names that must exist under `crates/`. const EXPECTED_CRATE_DIRS: &[&str] = &[ "core", + "token-counter", "config", "ai-gateway", "python-interop", diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index bda09a7d840..337a1e8e5ac 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -24,16 +24,17 @@ trace-parity = [ futures-util.workspace = true tracing = { workspace = true, optional = true } litellm-core = { workspace = true, features = ["bedrock-auth"] } +litellm-token-counter.workspace = true litellm-ai-gateway = { workspace = true, default-features = false } litellm-python-interop.workspace = true pyo3.workspace = true pyo3-async-runtimes.workspace = true serde.workspace = true serde_json.workspace = true -tokio.workspace = true +tokio = { workspace = true, features = ["sync"] } [dev-dependencies] -criterion = "0.8.2" +criterion.workspace = true tokio-tungstenite.workspace = true tracing.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/constants.rs b/litellm-rust/crates/python-bridge/src/constants.rs new file mode 100644 index 00000000000..d5cf5749820 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/constants.rs @@ -0,0 +1,2 @@ +/// Concurrent token-count encodes allowed when the core count is unavailable. +pub(crate) const TOKEN_COUNT_FALLBACK_PARALLELISM: usize = 1; diff --git a/litellm-rust/crates/python-bridge/src/execution.rs b/litellm-rust/crates/python-bridge/src/execution.rs index f3648158cf6..b57197b9ddf 100644 --- a/litellm-rust/crates/python-bridge/src/execution.rs +++ b/litellm-rust/crates/python-bridge/src/execution.rs @@ -3,7 +3,6 @@ use std::panic::AssertUnwindSafe; use std::time::Duration; use futures_util::FutureExt; -use litellm_core::error::Error; use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil}; use pyo3::exceptions::PyRuntimeError; use pyo3::prelude::*; @@ -11,14 +10,15 @@ use serde::Serialize; use tokio::runtime::{Handle, Runtime}; use tokio::time::{self, MissedTickBehavior}; -pub(crate) fn run_sync( +pub(crate) fn run_sync( py: Python<'_>, future: F, - map_error: fn(Error) -> PyErr, + map_error: fn(E) -> PyErr, ) -> PyResult> where T: Serialize + Send + 'static, - F: Future> + Send + 'static, + E: Send + 'static, + F: Future> + Send + 'static, { run_sync_on( py, @@ -28,15 +28,16 @@ where ) } -fn run_sync_on( +fn run_sync_on( py: Python<'_>, runtime: &Runtime, future: F, - map_error: fn(Error) -> PyErr, + map_error: fn(E) -> PyErr, ) -> PyResult> where T: Serialize + Send + 'static, - F: Future> + Send + 'static, + E: Send + 'static, + F: Future> + Send + 'static, { if Handle::try_current().is_ok() { return Err(PyRuntimeError::new_err( @@ -49,14 +50,15 @@ where Pythonized(result).into_pyobject(py).map(Bound::unbind) } -pub(crate) fn run_async( +pub(crate) fn run_async( py: Python<'_>, future: F, - map_error: fn(Error) -> PyErr, + map_error: fn(E) -> PyErr, ) -> PyResult> where T: Serialize + Send + 'static, - F: Future> + Send + 'static, + E: Send + 'static, + F: Future> + Send + 'static, { pyo3_async_runtimes::tokio::future_into_py(py, async move { let result = catch_future_panic(future).await?; @@ -65,7 +67,7 @@ where }) } -fn map_core_result(result: Result, map_error: fn(Error) -> PyErr) -> PyResult { +fn map_core_result(result: Result, map_error: fn(E) -> PyErr) -> PyResult { match result { Ok(value) => Ok(value), Err(error) => Err( @@ -75,9 +77,9 @@ fn map_core_result(result: Result, map_error: fn(Error) -> PyErr) - } } -async fn catch_future_panic(future: F) -> PyResult> +async fn catch_future_panic(future: F) -> PyResult> where - F: Future>, + F: Future>, { AssertUnwindSafe(future) .catch_unwind() @@ -85,9 +87,9 @@ where .map_err(panic_to_pyerr) } -async fn wait_for_sync_result(future: F) -> PyResult> +async fn wait_for_sync_result(future: F) -> PyResult> where - F: Future>, + F: Future>, { let future = catch_future_panic(future); tokio::pin!(future); @@ -114,6 +116,7 @@ mod tests { use std::thread; use std::time::Instant; + use litellm_core::error::Error; use pyo3::panic::PanicException; use pyo3::types::{PyDict, PyModule}; use serde::Serializer; @@ -237,7 +240,7 @@ mod tests { let error = runtime.block_on(async { Python::attach(|py| { - run_sync::(py, async { Ok(true) }, runtime_error) + run_sync::(py, async { Ok(true) }, runtime_error) .expect_err("sync route should reject a nested Tokio runtime") }) }); @@ -273,7 +276,7 @@ mod tests { fn sync_runner_maps_a_panicked_future() { Python::initialize(); Python::attach(|py| { - let error = run_sync::( + let error = run_sync::( py, poll_fn(|_| -> Poll> { panic!("route future panicked") }), runtime_error, @@ -289,7 +292,7 @@ mod tests { fn sync_runner_maps_a_panicked_error_mapper() { Python::initialize(); Python::attach(|py| { - let error = run_sync::( + let error = run_sync::( py, async { Err(Error::InvalidRequest("invalid".to_string())) }, panicking_error_mapper, diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 384f0be5a1b..cf0450a1b30 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -1,3 +1,4 @@ +mod constants; mod diagnostics; mod errors; mod execution; @@ -5,6 +6,7 @@ mod execution; mod function_trace; mod marshal; mod routes; +mod token_counter; use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use pyo3::prelude::*; @@ -71,6 +73,7 @@ mod _native { super::errors::register(module)?; super::routes::register(module)?; module.add_class::()?; + super::token_counter::register(module)?; super::diagnostics::register(module) } } @@ -106,6 +109,7 @@ mod tests { "chat_completions", "achat_completions", "ResponsesWebSocketConnection", + "TokenCounter", "gil_stats", ]; diff --git a/litellm-rust/crates/python-bridge/src/token_counter.rs b/litellm-rust/crates/python-bridge/src/token_counter.rs new file mode 100644 index 00000000000..ee82c170d55 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/token_counter.rs @@ -0,0 +1,87 @@ +use std::num::NonZero; +use std::sync::Arc; +use std::thread::available_parallelism; + +use litellm_python_interop::release_gil; +use litellm_token_counter::{ + CountableRequest, Error, InputTokenCount, TokenCounter as CoreTokenCounter, +}; +use pyo3::exceptions::{PyRuntimeError, PyValueError}; +use pyo3::prelude::*; +use pyo3::types::PyAny; +use tokio::sync::Semaphore; + +use crate::constants::TOKEN_COUNT_FALLBACK_PARALLELISM; +use crate::errors::RustBridgeDeclined; +use crate::execution::run_async; + +/// Counts the input tokens of a raw request body off the Python event loop with +/// the GIL released. Python owns which requests get here and what to do with +/// the count. At most one encode per core runs at a time; the rest wait in the +/// async task, where a cancelled Python awaiter drops them before any blocking +/// work is scheduled. +#[pyclass(frozen)] +struct TokenCounter { + inner: Arc, + encode_slots: Arc, +} + +#[pymethods] +impl TokenCounter { + #[new] + fn new(py: Python<'_>, tokenizer_json: &str) -> PyResult { + let inner = release_gil(py, || CoreTokenCounter::from_json(tokenizer_json)) + .map_err(token_count_error_to_pyerr)?; + Ok(Self { + inner: Arc::new(inner), + encode_slots: Arc::new(Semaphore::new(encode_parallelism())), + }) + } + + fn acount_request<'py>(&self, py: Python<'py>, body: &[u8]) -> PyResult> { + let counter = Arc::clone(&self.inner); + let encode_slots = Arc::clone(&self.encode_slots); + let body = body.to_vec(); + run_async( + py, + async move { + let _slot = encode_slots + .acquire_owned() + .await + .map_err(|error| Error::Task(error.to_string()))?; + tokio::task::spawn_blocking(move || count_body(&counter, &body)) + .await + .map_err(|error| Error::Task(error.to_string()))? + }, + token_count_error_to_pyerr, + ) + } +} + +fn encode_parallelism() -> usize { + available_parallelism().map_or(TOKEN_COUNT_FALLBACK_PARALLELISM, NonZero::get) +} + +fn count_body(counter: &CoreTokenCounter, body: &[u8]) -> Result { + let request = CountableRequest::parse(body)?; + counter.count_request(&request) +} + +fn token_count_error_to_pyerr(error: Error) -> PyErr { + let message = error.to_string(); + match error { + Error::Load(_) => PyValueError::new_err(message), + Error::RequestParse(_) + | Error::MissingInput + | Error::FloatText + | Error::ContentBlock + | Error::ArrayItems + | Error::JsonSerialization(_) + | Error::JsonUtf8(_) => RustBridgeDeclined::new_err(message), + Error::Encode(_) | Error::Task(_) => PyRuntimeError::new_err(message), + } +} + +pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + module.add_class::() +} diff --git a/litellm-rust/crates/token-counter/Cargo.toml b/litellm-rust/crates/token-counter/Cargo.toml new file mode 100644 index 00000000000..61c9cf6e991 --- /dev/null +++ b/litellm-rust/crates/token-counter/Cargo.toml @@ -0,0 +1,28 @@ +[package] +name = "litellm-token-counter" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +indexmap = { version = "2.14.0", features = ["serde"] } +itoa = "1.0" +serde.workspace = true +serde_json.workspace = true +thiserror.workspace = true +tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] } +unicode-normalization-alignments = "0.1.12" + +[dev-dependencies] +criterion.workspace = true +rand.workspace = true +rstest.workspace = true + +[[bench]] +name = "token_counter" +harness = false + +[[bench]] +name = "allocations" +harness = false diff --git a/litellm-rust/crates/token-counter/benches/allocations.rs b/litellm-rust/crates/token-counter/benches/allocations.rs new file mode 100644 index 00000000000..343f815749d --- /dev/null +++ b/litellm-rust/crates/token-counter/benches/allocations.rs @@ -0,0 +1,128 @@ +use std::alloc::{GlobalAlloc, Layout, System}; +use std::hint::black_box; +use std::sync::atomic::{AtomicUsize, Ordering}; + +use litellm_token_counter::{CountableRequest, TokenCounter}; + +struct CountingAllocator; + +static ALLOCATIONS: AtomicUsize = AtomicUsize::new(0); +static BYTES: AtomicUsize = AtomicUsize::new(0); + +unsafe impl GlobalAlloc for CountingAllocator { + unsafe fn alloc(&self, layout: Layout) -> *mut u8 { + ALLOCATIONS.fetch_add(1, Ordering::Relaxed); + BYTES.fetch_add(layout.size(), Ordering::Relaxed); + // SAFETY: This allocator delegates the unchanged layout to `System`. + unsafe { System.alloc(layout) } + } + + unsafe fn dealloc(&self, pointer: *mut u8, layout: Layout) { + // SAFETY: The pointer and layout came from the delegated `System` allocation. + unsafe { System.dealloc(pointer, layout) } + } + + unsafe fn realloc(&self, pointer: *mut u8, layout: Layout, size: usize) -> *mut u8 { + ALLOCATIONS.fetch_add(1, Ordering::Relaxed); + BYTES.fetch_add(size, Ordering::Relaxed); + // SAFETY: The pointer and layout came from `System`; the new size is unchanged. + unsafe { System.realloc(pointer, layout, size) } + } +} + +#[global_allocator] +static ALLOCATOR: CountingAllocator = CountingAllocator; + +#[derive(Clone, Copy)] +struct AllocationCount { + allocations: usize, + bytes: usize, +} + +impl AllocationCount { + fn assert_max(self, label: &str, maximum: Self) { + eprintln!( + "{label}: {} allocations, {} bytes", + self.allocations, self.bytes + ); + assert!( + self.allocations <= maximum.allocations, + "{label} allocation count exceeded {}", + maximum.allocations + ); + assert!( + self.bytes <= maximum.bytes, + "{label} allocated bytes exceeded {}", + maximum.bytes + ); + } +} + +const TOKENIZER_JSON: &str = include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json" +)); +const OBJECT_BODY: &[u8] = br#"{"model":"claude-sonnet-4-5","input":{"text":"caf\u00e9","n":3,"ok":true,"list":[1,"a",{"z":[]}]}}"#; +const INTEGER_BODY: &[u8] = br#"{"input":[-9223372036854775808,0,18446744073709551615]}"#; + +fn measure(operation: impl FnOnce()) -> AllocationCount { + ALLOCATIONS.store(0, Ordering::Relaxed); + BYTES.store(0, Ordering::Relaxed); + operation(); + AllocationCount { + allocations: ALLOCATIONS.load(Ordering::Relaxed), + bytes: BYTES.load(Ordering::Relaxed), + } +} + +fn main() { + measure(|| { + black_box(CountableRequest::parse(OBJECT_BODY).expect("object request parses")); + }) + .assert_max( + "parse object request", + AllocationCount { + allocations: 16, + bytes: 1_900, + }, + ); + + let counter = TokenCounter::from_json(TOKENIZER_JSON).expect("tokenizer loads"); + let object = CountableRequest::parse(OBJECT_BODY).expect("object request parses"); + counter + .count_request(&object) + .expect("object warmup succeeds"); + measure(|| { + black_box( + counter + .count_request(black_box(&object)) + .expect("object counts"), + ); + }) + .assert_max( + "count object request", + AllocationCount { + allocations: 74, + bytes: 2_200, + }, + ); + + let integers = CountableRequest::parse(INTEGER_BODY).expect("integer request parses"); + counter + .count_request(&integers) + .expect("integer warmup succeeds"); + measure(|| { + black_box( + counter + .count_request(black_box(&integers)) + .expect("integers count"), + ); + }) + .assert_max( + "count integer list", + AllocationCount { + allocations: 26, + bytes: 1_050, + }, + ); +} diff --git a/litellm-rust/crates/token-counter/benches/token_counter.rs b/litellm-rust/crates/token-counter/benches/token_counter.rs new file mode 100644 index 00000000000..a7c3177b88a --- /dev/null +++ b/litellm-rust/crates/token-counter/benches/token_counter.rs @@ -0,0 +1,100 @@ +use std::hint::black_box; +use std::time::Duration; + +use criterion::{BenchmarkId, Criterion, Throughput, criterion_group, criterion_main}; +use litellm_token_counter::TokenCounter; +use tokenizers::Tokenizer; + +const TOKENIZER_JSON: &str = include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json" +)); +const FULL_CONTEXT_TOKENS: usize = 1_000_000; +const LARGE_PROMPT_UNIT: &str = "The quick brown fox jumps over the lazy dog. 0123456789\n"; + +fn full_context_input(tokenizer: &Tokenizer) -> String { + let unit_tokens = tokenizer + .encode_fast(LARGE_PROMPT_UNIT, false) + .expect("reference tokenizer should encode") + .len(); + let input = LARGE_PROMPT_UNIT.repeat(FULL_CONTEXT_TOKENS.div_ceil(unit_tokens)); + let actual_tokens = tokenizer + .encode_fast(input.as_str(), true) + .expect("reference tokenizer should encode") + .len(); + assert!((FULL_CONTEXT_TOKENS..FULL_CONTEXT_TOKENS + unit_tokens).contains(&actual_tokens)); + input +} + +fn inputs(tokenizer: &Tokenizer) -> Vec<(&'static str, String)> { + vec![ + ( + "ascii_chat", + "User: Summarize the benefits of statistical benchmarking.\nAssistant:".repeat(8), + ), + ( + "unicode_nfkc", + "A quick café résumé: مرحبا 世界 🙂 fi Ⅳ.\n".repeat(16), + ), + ("large_prompt", LARGE_PROMPT_UNIT.repeat(256)), + ("full_context_1m_tokens", full_context_input(tokenizer)), + ] +} + +fn token_counter(c: &mut Criterion) { + let counter = TokenCounter::from_json(TOKENIZER_JSON).expect("token counter should load"); + let tokenizer = TOKENIZER_JSON + .parse::() + .expect("reference tokenizer should load"); + let mut group = c.benchmark_group("anthropic_token_counter"); + + for (name, input) in inputs(&tokenizer) { + let expected = tokenizer + .encode_fast(input.as_str(), true) + .expect("reference tokenizer should encode") + .len(); + let actual = counter + .count_text(input.as_str()) + .expect("benchmark path should count"); + assert_eq!( + actual, expected, + "benchmark paths should produce the same count" + ); + group.throughput(Throughput::Bytes(input.len() as u64)); + group.bench_with_input( + BenchmarkId::new("byte_level_fast_path", name), + &input, + |b, input| { + b.iter(|| { + counter + .count_text(black_box(input.as_str())) + .expect("fast path should count") + }) + }, + ); + group.bench_with_input( + BenchmarkId::new("full_encoder", name), + &input, + |b, input| { + b.iter(|| { + tokenizer + .encode_fast(black_box(input.as_str()), true) + .expect("reference tokenizer should encode") + .len() + }) + }, + ); + } + + group.finish(); +} + +criterion_group! { + name = benches; + config = Criterion::default() + .sample_size(20) + .warm_up_time(Duration::from_secs(1)) + .measurement_time(Duration::from_secs(4)); + targets = token_counter +} +criterion_main!(benches); diff --git a/litellm-rust/crates/token-counter/src/byte_level.rs b/litellm-rust/crates/token-counter/src/byte_level.rs new file mode 100644 index 00000000000..2b435eb39bb --- /dev/null +++ b/litellm-rust/crates/token-counter/src/byte_level.rs @@ -0,0 +1,650 @@ +//! Exact token counting for a supported tokenizer configuration: optional +//! NFKC normalization, `ByteLevel` pre-tokenization with the GPT-2 split regex, +//! and no post-processing. A scanner reproduces the regex's piece boundaries +//! and hands each piece to the tokenizer's model. Unsupported configurations +//! and added-token inputs fall back to the full encoder. + +use std::borrow::Cow; +use std::iter; + +use tokenizers::normalizers::NormalizerWrapper; +use tokenizers::pre_tokenizers::PreTokenizerWrapper; +use tokenizers::{Model, Tokenizer}; +use unicode_normalization_alignments::{IsNormalized, UnicodeNormalization, is_nfkc_quick}; + +use super::unicode_classes::UnicodeClasses; + +const CONTRACTIONS: [&str; 7] = ["'s", "'t", "'re", "'ve", "'m", "'ll", "'d"]; + +pub(super) struct ByteLevelCounter { + nfkc: bool, + normalized_added_tokens: Vec, + unicode_classes: &'static UnicodeClasses, +} + +impl ByteLevelCounter { + pub(super) fn detect(tokenizer: &Tokenizer) -> Option { + let nfkc = match tokenizer.get_normalizer() { + None => false, + Some(NormalizerWrapper::NFKC(_)) => true, + Some(_) => return None, + }; + let Some(PreTokenizerWrapper::ByteLevel(byte_level)) = tokenizer.get_pre_tokenizer() else { + return None; + }; + let plain = !byte_level.add_prefix_space + && byte_level.use_regex + && tokenizer.get_post_processor().is_none() + && tokenizer.get_truncation().is_none() + && tokenizer.get_padding().is_none(); + if !plain { + return None; + } + let vocabulary = tokenizer.get_added_vocabulary(); + let normalized_added_tokens = vocabulary + .get_vocab() + .iter() + .filter_map(|(original, id)| { + vocabulary + .simple_id_to_token(*id) + .filter(|normalized| normalized != original) + }) + .collect(); + Some(Self { + nfkc, + normalized_added_tokens, + unicode_classes: UnicodeClasses::get()?, + }) + } + + /// `None` when the text contains an added token or the model rejects a + /// piece; the caller then runs the full encoder. + pub(super) fn count(&self, tokenizer: &Tokenizer, text: &str) -> Option { + let normalized = self.normalize(text); + let added_tokens = tokenizer.get_added_vocabulary().get_vocab(); + if added_tokens + .keys() + .chain(self.normalized_added_tokens.iter()) + .any(|token| text.contains(token.as_str()) || normalized.contains(token.as_str())) + { + return None; + } + let model = tokenizer.get_model(); + let mapped: String = normalized.bytes().map(byte_char).collect(); + pieces(&normalized, self.unicode_classes) + .try_fold((0, 0), |(start, total), piece| { + let end = start + mapped_len(piece); + let tokens = model.tokenize(&mapped[start..end]).ok()?; + Some((end, total + tokens.len())) + }) + .map(|(_, total)| total) + } + + /// Same crate and Unicode tables as `NormalizedString::nfkc`, so the + /// result is what the full encoder would have tokenized. + fn normalize<'a>(&self, text: &'a str) -> Cow<'a, str> { + if !self.nfkc || text.is_ascii() || is_nfkc_quick(text.chars()) == IsNormalized::Yes { + return Cow::Borrowed(text); + } + Cow::Owned(text.nfkc().map(|(character, _)| character).collect()) + } +} + +/// GPT-2 `bytes_to_unicode`: printable Latin-1 bytes map to themselves, the +/// rest to U+0100 onwards in byte order. +fn byte_char(byte: u8) -> char { + let code = match byte { + 0x21..=0x7E | 0xA1..=0xAC | 0xAE..=0xFF => u32::from(byte), + 0x00..=0x20 => 0x100 + u32::from(byte), + 0x7F..=0xA0 => 0x121 + u32::from(byte - 0x7F), + 0xAD => 0x143, + }; + char::from_u32(code).unwrap_or(char::REPLACEMENT_CHARACTER) +} + +fn mapped_len(piece: &str) -> usize { + piece.len() + + piece + .bytes() + .filter(|byte| !byte.is_ascii_graphic()) + .count() +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum Class { + Letter, + Number, + Space, + Other, +} + +fn class(character: char, unicode_classes: &UnicodeClasses) -> Class { + match character { + 'A'..='Z' | 'a'..='z' => Class::Letter, + '0'..='9' => Class::Number, + '\t'..='\r' | ' ' => Class::Space, + _ if character.is_ascii() => Class::Other, + _ if unicode_classes.is_letter(character) => Class::Letter, + _ if unicode_classes.is_number(character) => Class::Number, + _ if unicode_classes.is_space(character) => Class::Space, + _ => Class::Other, + } +} + +/// The regex matches every character, so the pieces tile the text. +fn pieces<'a>( + text: &'a str, + unicode_classes: &'static UnicodeClasses, +) -> impl Iterator { + iter::successors(split_piece(text, unicode_classes), move |(_, rest)| { + split_piece(rest, unicode_classes) + }) + .map(|(piece, _)| piece) +} + +fn split_piece<'a>(text: &'a str, unicode_classes: &UnicodeClasses) -> Option<(&'a str, &'a str)> { + let first = text.chars().next()?; + Some(text.split_at(piece_len(text, first, unicode_classes))) +} + +fn piece_len(text: &str, first: char, unicode_classes: &UnicodeClasses) -> usize { + if let Some(contraction) = CONTRACTIONS.iter().find(|word| text.starts_with(**word)) { + return contraction.len(); + } + let first_class = class(first, unicode_classes); + if first_class != Class::Space { + return run_len(text, first_class, unicode_classes); + } + if first != ' ' { + return space_run_len(text, unicode_classes); + } + let after_space = &text[1..]; + match after_space + .chars() + .next() + .map(|character| class(character, unicode_classes)) + { + None | Some(Class::Space) => space_run_len(text, unicode_classes), + Some(run_class) => 1 + run_len(after_space, run_class, unicode_classes), + } +} + +fn run_len(text: &str, run_class: Class, unicode_classes: &UnicodeClasses) -> usize { + text.char_indices() + .find(|(_, character)| class(*character, unicode_classes) != run_class) + .map_or(text.len(), |(index, _)| index) +} + +/// `\s+(?!\S)|\s+`: whitespace followed by a non-space leaves its last +/// character to start the next piece (` ?` on the following alternatives). +fn space_run_len(text: &str, unicode_classes: &UnicodeClasses) -> usize { + let run = run_len(text, Class::Space, unicode_classes); + if run == text.len() { + return run; + } + let last = text[..run].chars().next_back().map_or(0, char::len_utf8); + match run - last { + 0 => run, + shorter => shorter, + } +} + +#[cfg(test)] +mod tests { + use rand::rngs::StdRng; + use rand::seq::SliceRandom; + use rand::{Rng, SeedableRng}; + use rstest::{fixture, rstest}; + use tokenizers::normalizers::NFKC; + use tokenizers::pre_tokenizers::byte_level::ByteLevel; + use tokenizers::utils::SysRegex; + use tokenizers::{ + NormalizedString, Normalizer, OffsetReferential, OffsetType, PreTokenizedString, + PreTokenizer, + }; + + use super::*; + + #[fixture] + fn anthropic_tokenizer() -> Tokenizer { + let path = concat!( + env!("CARGO_MANIFEST_DIR"), + "/../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json" + ); + std::fs::read_to_string(path) + .expect("anthropic tokenizer json is in the repo") + .parse() + .expect("anthropic tokenizer loads") + } + + fn reference_count(tokenizer: &Tokenizer, text: &str) -> usize { + tokenizer.encode_fast(text, true).expect("encode").len() + } + + fn byte_level_counter(nfkc: bool) -> ByteLevelCounter { + ByteLevelCounter { + nfkc, + normalized_added_tokens: Vec::new(), + unicode_classes: UnicodeClasses::get().expect("Oniguruma exposes Unicode classes"), + } + } + + const ALPHABET: &[&str] = &[ + "a", + "Z", + "e", + "s", + "t", + "d", + "m", + "'", + "'s", + "'re", + "'ll", + "'S", + "0", + "9", + " ", + " ", + "\t", + "\n", + "\r\n", + "\u{b}", + ".", + ",", + "!", + "-", + "(", + "\"", + "\u{a0}", + "\u{85}", + "\u{2028}", + "\u{3000}", + "\u{200b}", + "\u{200d}", + "é", + "e\u{301}", + "ß", + "漢", + "字", + "ع", + "३", + "½", + "Ⅳ", + "🙂", + "👍🏽", + "A", + "fi", + "㍿", + "㋿", + "ꟲ", + "𐞁", + "a\u{30a}", + "\u{1e0b}\u{323}", + "<", + ">", + "EOT", + "", + "", + ]; + + fn random_text(rng: &mut StdRng) -> String { + let pieces = rng.gen_range(0..40); + (0..pieces) + .map(|_| *ALPHABET.choose(rng).expect("alphabet is not empty")) + .collect() + } + + #[rstest] + #[case::plain_text("Hello, how are you today?", true)] + #[case::added_token("stop here", false)] + #[case::normalized_added_token("stop <EOT> here", false)] + fn anthropic_tokenizer_takes_the_fast_path( + anthropic_tokenizer: Tokenizer, + #[case] text: &str, + #[case] supported: bool, + ) { + let fast = + ByteLevelCounter::detect(&anthropic_tokenizer).expect("anthropic shape is supported"); + assert!(fast.nfkc); + let count = fast.count(&anthropic_tokenizer, text); + if supported { + assert_eq!(count, Some(reference_count(&anthropic_tokenizer, text))); + } else { + assert_eq!(count, None); + } + } + + #[rstest] + fn counts_match_the_full_encoder(anthropic_tokenizer: Tokenizer) { + let fast = ByteLevelCounter::detect(&anthropic_tokenizer).expect("supported"); + let mut rng = StdRng::seed_from_u64(2026); + for _ in 0..4000 { + let text = random_text(&mut rng).replace('<', "("); + let expected = reference_count(&anthropic_tokenizer, &text); + assert_eq!( + fast.count(&anthropic_tokenizer, &text), + Some(expected), + "text {text:?}" + ); + } + } + + #[rstest] + fn nfkc_matches_the_tokenizer_normalizer_for_every_scalar_value() { + let fast = byte_level_counter(true); + let mut text = String::new(); + for character in (0..=0x10FFFFu32).filter_map(char::from_u32) { + text.clear(); + text.push(character); + let mut expected = NormalizedString::from(text.as_str()); + NFKC.normalize(&mut expected).expect("nfkc"); + assert_eq!( + fast.normalize(&text), + expected.get(), + "U+{:04X}", + u32::from(character) + ); + } + } + + #[rstest] + fn nfkc_matches_the_tokenizer_normalizer_on_random_texts() { + let fast = byte_level_counter(true); + let mut rng = StdRng::seed_from_u64(11); + for _ in 0..4000 { + let text = random_text(&mut rng); + let mut expected = NormalizedString::from(text.as_str()); + NFKC.normalize(&mut expected).expect("nfkc"); + assert_eq!(fast.normalize(&text), expected.get(), "text {text:?}"); + } + } + + #[rstest] + fn pieces_match_the_byte_level_pre_tokenizer() { + let byte_level = ByteLevel::new(false, true, true); + let mut rng = StdRng::seed_from_u64(7); + for _ in 0..4000 { + let text = byte_level_counter(true) + .normalize(&random_text(&mut rng)) + .into_owned(); + let mut pre_tokenized = PreTokenizedString::from(text.as_str()); + byte_level + .pre_tokenize(&mut pre_tokenized) + .expect("pre-tokenize"); + let expected: Vec<(String, (usize, usize))> = pre_tokenized + .get_splits(OffsetReferential::Original, OffsetType::Byte) + .into_iter() + .map(|(mapped, offsets, _)| (mapped.to_string(), offsets)) + .collect(); + let actual: Vec<(String, (usize, usize))> = pieces( + &text, + UnicodeClasses::get().expect("Oniguruma exposes Unicode classes"), + ) + .map(|piece| { + let start = piece.as_ptr() as usize - text.as_ptr() as usize; + let mapped: String = piece.bytes().map(byte_char).collect(); + (mapped, (start, start + piece.len())) + }) + .collect(); + assert_eq!(actual, expected, "text {text:?}"); + } + } + + #[rstest] + fn byte_chars_match_the_byte_level_alphabet() { + let byte_level = ByteLevel::new(false, false, false); + let characters: Vec = (0..=0x10FFFFu32).filter_map(char::from_u32).collect(); + for chunk in characters.chunks(1024) { + let text: String = chunk.iter().collect(); + let mut pre_tokenized = PreTokenizedString::from(text.as_str()); + byte_level + .pre_tokenize(&mut pre_tokenized) + .expect("pre-tokenize"); + let expected: String = pre_tokenized + .get_splits(OffsetReferential::Original, OffsetType::Byte) + .into_iter() + .map(|(mapped, _, _)| mapped) + .collect(); + let actual: String = text.bytes().map(byte_char).collect(); + assert_eq!(actual.len(), mapped_len(&text)); + assert_eq!( + actual, + expected, + "chunk starting at U+{:04X}", + u32::from(chunk[0]) + ); + } + } + + #[rstest] + fn classes_match_oniguruma() { + let unicode_classes = UnicodeClasses::get().expect("Oniguruma exposes Unicode classes"); + let letter = SysRegex::new(r"\p{L}").expect("regex"); + let number = SysRegex::new(r"\p{N}").expect("regex"); + let space = SysRegex::new(r"\s").expect("regex"); + let whole = + |regex: &SysRegex, text: &str| regex.find_iter(text).next() == Some((0, text.len())); + let mut text = String::new(); + for character in (0..=0x10FFFFu32).filter_map(char::from_u32) { + text.clear(); + text.push(character); + let expected = if whole(&letter, &text) { + Class::Letter + } else if whole(&number, &text) { + Class::Number + } else if whole(&space, &text) { + Class::Space + } else { + Class::Other + }; + assert_eq!( + class(character, unicode_classes), + expected, + "U+{:04X}", + u32::from(character) + ); + } + } + + #[rstest] + #[case("prefix")] + #[case("regex")] + #[case("normalizer")] + #[case("no_pre_tokenizer")] + #[case("other_pre_tokenizer")] + #[case("post_processor")] + #[case("truncation")] + #[case("padding")] + fn other_tokenizer_shapes_are_declined( + mut anthropic_tokenizer: Tokenizer, + #[case] shape: &str, + ) { + use tokenizers::{PaddingParams, PaddingStrategy, TruncationParams}; + match shape { + "prefix" => { + anthropic_tokenizer.with_pre_tokenizer(Some(ByteLevel::new(true, true, true))); + } + "regex" => { + anthropic_tokenizer.with_pre_tokenizer(Some(ByteLevel::new(false, true, false))); + } + "normalizer" => { + anthropic_tokenizer + .with_normalizer(Some(tokenizers::normalizers::Lowercase)) + .expect("normalizer"); + } + "no_pre_tokenizer" => { + anthropic_tokenizer.with_pre_tokenizer(None::); + } + "other_pre_tokenizer" => { + anthropic_tokenizer + .with_pre_tokenizer(Some(tokenizers::pre_tokenizers::whitespace::Whitespace)); + } + "post_processor" => { + anthropic_tokenizer.with_post_processor(Some(ByteLevel::default())); + } + "truncation" => { + anthropic_tokenizer + .with_truncation(Some(TruncationParams { + max_length: 2, + ..Default::default() + })) + .expect("truncation"); + } + "padding" => { + anthropic_tokenizer.with_padding(Some(PaddingParams { + strategy: PaddingStrategy::Fixed(32), + ..Default::default() + })); + } + _ => unreachable!(), + } + assert!(ByteLevelCounter::detect(&anthropic_tokenizer).is_none()); + let counter = crate::TokenCounter::from_json( + &anthropic_tokenizer.to_string(false).expect("serialize"), + ) + .expect("load"); + for text in ["", "Hello WORLD! AB fi Ⅳ", " stop"] { + assert_eq!( + counter.count_text(text).expect("count"), + reference_count(&anthropic_tokenizer, text) + ); + } + } + + #[rstest] + #[case(false)] + #[case(true)] + fn arbitrary_unicode_and_long_inputs_use_fast_path( + mut anthropic_tokenizer: Tokenizer, + #[case] nfkc: bool, + ) { + if !nfkc { + anthropic_tokenizer + .with_normalizer(None::) + .expect("normalizer"); + } + let fast = ByteLevelCounter::detect(&anthropic_tokenizer).expect("supported"); + let mut rng = StdRng::seed_from_u64(314159); + for _ in 0..1000 { + let text: String = (0..64) + .filter_map(|_| char::from_u32(rng.gen_range(0..=0x10ffff))) + .collect(); + assert_eq!( + fast.count(&anthropic_tokenizer, &text), + Some(reference_count(&anthropic_tokenizer, &text)), + "text {text:?}" + ); + } + for text in [ + "", + "'s't're've'm'll'd'S'RE", + " a \t\r\n b\u{85}\u{a0}c ", + "\0é漢🙂", + "a\u{30a}\u{301}", + "AfiⅣ", + ] { + let text = text.repeat(2048); + assert_eq!( + fast.count(&anthropic_tokenizer, &text), + Some(reference_count(&anthropic_tokenizer, &text)) + ); + } + } + + #[rstest] + #[case(false, false, false, false)] + #[case(true, false, false, false)] + #[case(false, true, false, false)] + #[case(false, false, true, false)] + #[case(false, false, false, true)] + fn added_token_options_fall_back( + mut anthropic_tokenizer: Tokenizer, + #[case] special: bool, + #[case] single_word: bool, + #[case] lstrip: bool, + #[case] rstrip: bool, + ) { + anthropic_tokenizer + .add_tokens([tokenizers::AddedToken::from("custom token", special) + .single_word(single_word) + .lstrip(lstrip) + .rstrip(rstrip)]) + .expect("add token"); + let fast = ByteLevelCounter::detect(&anthropic_tokenizer).expect("supported"); + let counter = crate::TokenCounter::from_json( + &anthropic_tokenizer.to_string(false).expect("serialize"), + ) + .expect("load"); + for text in [ + "custom token", + "a custom token b", + "acustom tokenb", + " custom token ", + ] { + assert_eq!(fast.count(&anthropic_tokenizer, text), None); + assert_eq!( + counter.count_text(text).expect("count"), + reference_count(&anthropic_tokenizer, text) + ); + } + } + + #[test] + fn model_errors_reach_public_caller() { + let mut tokenizer = Tokenizer::new(tokenizers::models::wordpiece::WordPiece::default()); + tokenizer.with_pre_tokenizer(Some(ByteLevel::new(false, true, true))); + let fast = ByteLevelCounter::detect(&tokenizer).expect("supported"); + assert_eq!(fast.count(&tokenizer, "hello"), None); + assert!(tokenizer.encode_fast("hello", true).is_err()); + let counter = + crate::TokenCounter::from_json(&tokenizer.to_string(false).expect("serialize")) + .expect("load"); + assert!(matches!( + counter.count_text("hello"), + Err(crate::Error::Encode(_)) + )); + } + + #[rstest] + fn shared_counter_matches_encoder_across_threads(anthropic_tokenizer: Tokenizer) { + let counter = crate::TokenCounter::from_json( + &anthropic_tokenizer.to_string(false).expect("serialize"), + ) + .expect("load"); + let inputs = [ + "hello world", + "AB fi\n漢字🙂", + " stop", + "\t 're \r\n", + ]; + let expected = inputs.map(|text| reference_count(&anthropic_tokenizer, text)); + std::thread::scope(|scope| { + for _ in 0..8 { + let counter = &counter; + scope.spawn(move || { + for _ in 0..100 { + for (text, count) in inputs.iter().zip(expected) { + assert_eq!(counter.count_text(text).expect("count"), count); + } + } + }); + } + }); + } + + #[rstest] + fn normalized_added_token_spelling_declines_fast_path(mut anthropic_tokenizer: Tokenizer) { + anthropic_tokenizer + .add_tokens([tokenizers::AddedToken::from("ABCD EFGH", false)]) + .expect("add token"); + let fast = ByteLevelCounter::detect(&anthropic_tokenizer).expect("supported"); + assert_eq!(reference_count(&anthropic_tokenizer, "ABCD EFGH"), 1); + assert_eq!(fast.count(&anthropic_tokenizer, "ABCD EFGH"), None); + let counter = crate::TokenCounter::from_json( + &anthropic_tokenizer.to_string(false).expect("serialize"), + ) + .expect("load"); + assert_eq!(counter.count_text("ABCD EFGH").expect("count"), 1); + } +} diff --git a/litellm-rust/crates/token-counter/src/counter.rs b/litellm-rust/crates/token-counter/src/counter.rs new file mode 100644 index 00000000000..a3943515a64 --- /dev/null +++ b/litellm-rust/crates/token-counter/src/counter.rs @@ -0,0 +1,194 @@ +use serde::Serialize; + +use crate::Error; +use crate::byte_level::ByteLevelCounter; +use crate::python_json; +use crate::tools::format_function_definitions; +use crate::types::{ + ContentBlock, ContentItem, CountableRequest, Message, MessageContent, TextValue, ToolChoice, + ToolDefinition, +}; + +const TOKENS_PER_MESSAGE: usize = 3; +const TOKENS_PER_NAME: usize = 1; +const REPLY_PRIMING_TOKENS: usize = 3; +const TOOL_DEFINITIONS_TOKENS: usize = 9; +const TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT: usize = 4; +const TOOL_CHOICE_NONE_TOKENS: usize = 1; +const NAMED_TOOL_CHOICE_TOKENS: usize = 7; + +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +pub struct InputTokenCount { + pub model: Option, + pub input_tokens: usize, +} + +/// A loaded HuggingFace tokenizer plus the message accounting Python applies on +/// top of it. Encoding is CPU-bound and synchronous; hosts run it off their +/// event loop. +pub struct TokenCounter { + tokenizer: tokenizers::Tokenizer, + byte_level: Option, +} + +impl TokenCounter { + /// Load a HuggingFace `tokenizer.json` document. The host reads the file. + pub fn from_json(tokenizer_json: &str) -> Result { + let tokenizer = tokenizer_json + .parse::() + .map_err(Error::Load)?; + let byte_level = ByteLevelCounter::detect(&tokenizer); + Ok(Self { + tokenizer, + byte_level, + }) + } + + pub fn count_text(&self, text: &str) -> Result { + if let Some(count) = self + .byte_level + .as_ref() + .and_then(|counter| counter.count(&self.tokenizer, text)) + { + return Ok(count); + } + self.tokenizer + .encode_fast(text, true) + .map(|encoding| encoding.len()) + .map_err(Error::Encode) + } + + /// Mirrors the host's key precedence: `messages`, then `prompt`, then + /// `input`, then `query` plus `documents`. + pub fn count_request(&self, request: &CountableRequest) -> Result { + let input_tokens = if let Some(messages) = &request.messages { + self.count_messages(request, messages)? + } else if let Some(prompt) = &request.prompt { + self.count_text_value(prompt)? + } else if let Some(input) = &request.input { + self.count_text_value(input)? + } else if request.query.is_some() || request.documents.is_some() { + self.count_optional_text_value(request.query.as_ref())? + + self.count_optional_text_value(request.documents.as_ref())? + } else { + return Err(Error::MissingInput); + }; + Ok(InputTokenCount { + model: request.model.clone(), + input_tokens, + }) + } + + fn count_messages( + &self, + request: &CountableRequest, + messages: &[Message], + ) -> Result { + let message_tokens = messages + .iter() + .map(|message| self.count_message(message)) + .sum::>()?; + let includes_system_message = messages + .iter() + .any(|message| message.role.as_deref() == Some("system")); + let extra_tokens = self.count_extra( + request.tools.as_deref().unwrap_or_default(), + request.tool_choice.as_ref(), + includes_system_message, + )?; + Ok(message_tokens + extra_tokens) + } + + fn count_optional_text_value(&self, value: Option<&TextValue>) -> Result { + value.map_or(Ok(0), |value| self.count_text_value(value)) + } + + /// `str()` for scalars, `json.dumps()` for objects, lists flattened, nulls + /// skipped. Floats are declined because Python's `repr` and Rust's float + /// formatting disagree on exponents. + fn count_text_value(&self, value: &TextValue) -> Result { + match value { + TextValue::Null => Ok(0), + TextValue::Bool(true) => self.count_text("True"), + TextValue::Bool(false) => self.count_text("False"), + TextValue::Number(number) => match (number.as_i64(), number.as_u64()) { + (Some(number), _) => self.count_text(itoa::Buffer::new().format(number)), + (_, Some(number)) => self.count_text(itoa::Buffer::new().format(number)), + _ => Err(Error::FloatText), + }, + TextValue::Text(text) => self.count_text(text), + TextValue::List(items) => items + .iter() + .map(|item| self.count_text_value(item)) + .sum::>(), + TextValue::Object(_) => self.count_text(&python_json::dumps(value)?), + } + } + + fn count_message(&self, message: &Message) -> Result { + let role_tokens = match &message.role { + Some(role) => self.count_text(role)?, + None => 0, + }; + let name_tokens = match &message.name { + Some(name) => self.count_text(name)? + TOKENS_PER_NAME, + None => 0, + }; + let content_tokens = match &message.content { + Some(MessageContent::Text(text)) => self.count_text(text)?, + Some(MessageContent::Blocks(items)) => items + .iter() + .map(|item| self.count_content_item(item)) + .sum::>()?, + None => 0, + }; + Ok(TOKENS_PER_MESSAGE + role_tokens + name_tokens + content_tokens) + } + + fn count_content_item(&self, item: &ContentItem) -> Result { + match item { + ContentItem::Text(text) => self.count_text(text), + ContentItem::Block(ContentBlock::Text { text }) => self.count_text(text), + ContentItem::Block(ContentBlock::Thinking { thinking }) => { + if thinking.is_empty() { + return Ok(0); + } + self.count_text(thinking) + } + ContentItem::Block(ContentBlock::ToolReference { tool_name }) => { + match tool_name.as_deref().filter(|name| !name.is_empty()) { + Some(name) => self.count_text(name), + None => Ok(0), + } + } + ContentItem::Block(ContentBlock::Unsupported) => Err(Error::ContentBlock), + } + } + + fn count_extra( + &self, + tools: &[ToolDefinition], + tool_choice: Option<&ToolChoice>, + includes_system_message: bool, + ) -> Result { + let tool_tokens = if tools.is_empty() { + 0 + } else { + let definitions = self.count_text(&format_function_definitions(tools)?)?; + let discount = if includes_system_message { + TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT + } else { + 0 + }; + definitions + TOOL_DEFINITIONS_TOKENS - discount + }; + let choice_tokens = match tool_choice { + Some(ToolChoice::Mode(mode)) if mode == "none" => TOOL_CHOICE_NONE_TOKENS, + Some(ToolChoice::Mode(_)) | None => 0, + Some(ToolChoice::Named(named)) => { + NAMED_TOOL_CHOICE_TOKENS + self.count_text(&named.function.name)? + } + }; + Ok(REPLY_PRIMING_TOKENS + tool_tokens + choice_tokens) + } +} diff --git a/litellm-rust/crates/token-counter/src/error.rs b/litellm-rust/crates/token-counter/src/error.rs new file mode 100644 index 00000000000..dc590093f64 --- /dev/null +++ b/litellm-rust/crates/token-counter/src/error.rs @@ -0,0 +1,31 @@ +use std::string::FromUtf8Error; + +use thiserror::Error as ThisError; + +#[derive(Debug, ThisError)] +pub enum Error { + #[error("failed to load tokenizer: {0}")] + Load(#[source] tokenizers::Error), + #[error("unsupported by the rust token counter: request body could not be parsed: {0}")] + RequestParse(#[source] serde_json::Error), + #[error("unsupported by the rust token counter: request has no countable input")] + MissingInput, + #[error( + "unsupported by the rust token counter: float text values are counted by the python path" + )] + FloatText, + #[error( + "unsupported by the rust token counter: content block type is counted by the python path" + )] + ContentBlock, + #[error("unsupported by the rust token counter: array parameter without items")] + ArrayItems, + #[error("unsupported by the rust token counter: text value could not be serialized: {0}")] + JsonSerialization(#[source] serde_json::Error), + #[error("unsupported by the rust token counter: serialized text value is not UTF-8: {0}")] + JsonUtf8(#[source] FromUtf8Error), + #[error("tokenization failed: {0}")] + Encode(#[source] tokenizers::Error), + #[error("token counting task failed: {0}")] + Task(String), +} diff --git a/litellm-rust/crates/token-counter/src/lib.rs b/litellm-rust/crates/token-counter/src/lib.rs new file mode 100644 index 00000000000..eb602b3cade --- /dev/null +++ b/litellm-rust/crates/token-counter/src/lib.rs @@ -0,0 +1,17 @@ +//! Input token counting for a request body, mirroring `litellm.token_counter` +//! for the shapes it can count exactly. Everything else is declined so the host +//! keeps its own counter as the reference. + +#![forbid(unsafe_code)] + +mod byte_level; +mod counter; +mod error; +mod python_json; +mod tools; +mod types; +mod unicode_classes; + +pub use counter::{InputTokenCount, TokenCounter}; +pub use error::Error; +pub use types::CountableRequest; diff --git a/litellm-rust/crates/token-counter/src/python_json.rs b/litellm-rust/crates/token-counter/src/python_json.rs new file mode 100644 index 00000000000..00a55ad300c --- /dev/null +++ b/litellm-rust/crates/token-counter/src/python_json.rs @@ -0,0 +1,155 @@ +//! `json.dumps(value)` with Python's default arguments: `", "` and `": "` +//! separators, `ensure_ascii=True`, and keys in insertion order. + +use std::io::{self, Write}; + +use serde::Serialize; +use serde_json::ser::{Formatter, Serializer}; + +use super::Error; +use super::types::TextValue; + +pub(super) fn dumps(value: &TextValue) -> Result { + let mut output = Vec::with_capacity(serialized_len(value)?); + value + .serialize(&mut Serializer::with_formatter( + &mut output, + PythonFormatter, + )) + .map_err(Error::JsonSerialization)?; + debug_assert_eq!(output.len(), output.capacity()); + String::from_utf8(output).map_err(Error::JsonUtf8) +} + +fn serialized_len(value: &TextValue) -> Result { + match value { + TextValue::Null => Ok(4), + TextValue::Bool(true) => Ok(4), + TextValue::Bool(false) => Ok(5), + TextValue::Number(number) => match (number.as_i64(), number.as_u64()) { + (Some(number), _) => Ok(unsigned_len(number.unsigned_abs()) + usize::from(number < 0)), + (_, Some(number)) => Ok(unsigned_len(number)), + _ => Err(Error::FloatText), + }, + TextValue::Text(text) => Ok(quoted_len(text)), + TextValue::List(items) => items + .iter() + .try_fold(2 + items.len().saturating_sub(1) * 2, |len, item| { + Ok(len + serialized_len(item)?) + }), + TextValue::Object(entries) => entries.iter().try_fold( + 2 + entries.len().saturating_sub(1) * 2, + |len, (key, value)| Ok(len + quoted_len(key) + 2 + serialized_len(value)?), + ), + } +} + +fn unsigned_len(number: u64) -> usize { + if number == 0 { + 1 + } else { + number.ilog10() as usize + 1 + } +} + +fn quoted_len(value: &str) -> usize { + value.chars().fold(2, |len, character| { + len + match character { + '"' | '\\' | '\u{0008}' | '\u{000c}' | '\n' | '\r' | '\t' => 2, + '\u{0000}'..='\u{001f}' => 6, + ' '..='~' => 1, + _ => character.len_utf16() * 6, + } + }) +} + +struct PythonFormatter; + +impl Formatter for PythonFormatter { + fn begin_array_value(&mut self, writer: &mut W, first: bool) -> io::Result<()> + where + W: ?Sized + Write, + { + if !first { + writer.write_all(b", ")?; + } + Ok(()) + } + + fn begin_object_key(&mut self, writer: &mut W, first: bool) -> io::Result<()> + where + W: ?Sized + Write, + { + if !first { + writer.write_all(b", ")?; + } + Ok(()) + } + + fn begin_object_value(&mut self, writer: &mut W) -> io::Result<()> + where + W: ?Sized + Write, + { + writer.write_all(b": ") + } + + fn write_string_fragment(&mut self, writer: &mut W, fragment: &str) -> io::Result<()> + where + W: ?Sized + Write, + { + for character in fragment.chars() { + if (' '..='~').contains(&character) { + write!(writer, "{character}")?; + continue; + } + let mut units = [0u16; 2]; + for unit in character.encode_utf16(&mut units) { + write!(writer, "\\u{unit:04x}")?; + } + } + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + + use super::dumps; + + #[rstest] + #[case::null("null", "null")] + #[case::boolean("true", "true")] + #[case::signed_integer("-3", "-3")] + #[case::unsigned_integer("18446744073709551615", "18446744073709551615")] + #[case::empty_array("[]", "[]")] + #[case::empty_object("{}", "{}")] + #[case::nested( + r#"{"first":1,"second":{"ok":true,"none":null},"third":[false,2]}"#, + r#"{"first": 1, "second": {"ok": true, "none": null}, "third": [false, 2]}"# + )] + #[case::string_escaping( + r#""caf\u00e9 \u2014 \ud83d\ude00 \"q\" \\ \n\t\u0001\u007f ~ /""#, + r#""caf\u00e9 \u2014 \ud83d\ude00 \"q\" \\ \n\t\u0001\u007f ~ /""# + )] + #[case::short_control_escapes(r#""\b\f\r""#, r#""\b\f\r""#)] + fn matches_python_json_dumps(#[case] input: &str, #[case] expected: &str) { + let value = serde_json::from_str(input).expect("fixture parses"); + + assert_eq!(dumps(&value).expect("fixture dumps"), expected); + } + + #[rstest] + #[case::top_level("1.5")] + #[case::array("[1,2.5]")] + #[case::object(r#"{"nested":{"value":-0.25}}"#)] + fn rejects_floats(#[case] input: &str) { + let value = serde_json::from_str(input).expect("fixture parses"); + let error = dumps(&value).expect_err("floats are declined"); + + assert_eq!( + error.to_string(), + "unsupported by the rust token counter: float text values are counted by the python path" + ); + } +} diff --git a/litellm-rust/crates/token-counter/src/tools.rs b/litellm-rust/crates/token-counter/src/tools.rs new file mode 100644 index 00000000000..22150a4e4b1 --- /dev/null +++ b/litellm-rust/crates/token-counter/src/tools.rs @@ -0,0 +1,304 @@ +//! Renders tool definitions the way `litellm.token_counter` does before +//! tokenizing them (the TypeScript-like namespace OpenAI appears to use). + +use std::fmt::Write; + +use super::Error; +use super::types::{EnumValue, Schema, SchemaType, ToolDefinition}; + +pub(super) fn format_function_definitions(tools: &[ToolDefinition]) -> Result { + ToolFormatter::new(tools.len()).format(tools) +} + +struct ToolFormatter { + output: String, +} + +impl ToolFormatter { + fn new(tool_count: usize) -> Self { + Self { + output: String::with_capacity(tool_count.saturating_mul(128).saturating_add(48)), + } + } + + fn format(mut self, tools: &[ToolDefinition]) -> Result { + self.output.push_str("namespace functions {\n\n"); + + for tool in tools { + self.write_function(tool)?; + } + + self.output.push_str("} // namespace functions"); + Ok(self.output) + } + + fn write_function(&mut self, tool: &ToolDefinition) -> Result<(), Error> { + let (name, description, parameters) = resolve_function(tool); + let Some(name) = name.filter(|name| !name.is_empty()) else { + return Ok(()); + }; + + if let Some(description) = description.filter(|description| !description.is_empty()) { + self.output.push_str("// "); + self.output.push_str(description); + self.output.push('\n'); + } + + match parameters.filter(|parameters| { + parameters + .properties + .as_ref() + .is_some_and(|properties| !properties.is_empty()) + }) { + Some(parameters) => { + self.output.push_str("type "); + self.output.push_str(name); + self.output.push_str(" = (_: {\n"); + self.write_object_parameters(parameters, 0)?; + self.output.push_str("\n}) => any;\n\n"); + } + _ => { + self.output.push_str("type "); + self.output.push_str(name); + self.output.push_str(" = () => any;\n\n"); + } + } + + Ok(()) + } + + fn write_object_parameters(&mut self, parameters: &Schema, indent: usize) -> Result<(), Error> { + let Some(properties) = parameters + .properties + .as_ref() + .filter(|properties| !properties.is_empty()) + else { + return Ok(()); + }; + let required = parameters.required.as_deref().unwrap_or_default(); + for (index, (key, props)) in properties.iter().enumerate() { + if index > 0 { + self.output.push('\n'); + } + if let Some(description) = props + .description + .as_deref() + .filter(|description| !description.is_empty()) + { + self.write_indent(indent); + self.output.push_str("// "); + self.output.push_str(description); + self.output.push('\n'); + } + + self.write_indent(indent); + self.output.push_str(key); + if !required.iter().any(|required| required == key) { + self.output.push('?'); + } + self.output.push_str(": "); + self.write_type(props, indent)?; + self.output.push(','); + } + + Ok(()) + } + + fn write_type(&mut self, props: &Schema, indent: usize) -> Result<(), Error> { + let Some(SchemaType::Name(schema_type)) = &props.schema_type else { + self.output.push_str("any"); + return Ok(()); + }; + + match schema_type.as_str() { + "string" | "integer" | "number" => match &props.enum_values { + Some(values) => self.write_enum(values), + None if schema_type == "string" => self.output.push_str("string"), + None => self.output.push_str("number"), + }, + "array" => { + let items = props.items.as_deref().ok_or(Error::ArrayItems)?; + self.write_type(items, indent)?; + self.output.push_str("[]"); + } + "object" => { + self.output.push_str("{\n"); + self.write_object_parameters(props, indent + 2)?; + self.output.push_str("\n}"); + } + "boolean" => self.output.push_str("boolean"), + "null" => self.output.push_str("null"), + _ => self.output.push_str("any"), + } + + Ok(()) + } + + fn write_enum(&mut self, values: &[EnumValue]) { + for (index, value) in values.iter().enumerate() { + if index > 0 { + self.output.push_str(" | "); + } + self.output.push('"'); + match value { + EnumValue::Text(text) => self.output.push_str(text), + EnumValue::Integer(number) => { + write!(self.output, "{number}").expect("writing to a String cannot fail"); + } + } + self.output.push('"'); + } + } + + fn write_indent(&mut self, indent: usize) { + for _ in 0..indent { + self.output.push(' '); + } + } +} + +fn resolve_function(tool: &ToolDefinition) -> (Option<&str>, Option<&str>, Option<&Schema>) { + match &tool.function { + Some(function) => ( + function.name.as_deref(), + function.description.as_deref(), + function.parameters.as_ref(), + ), + None => ( + tool.name.as_deref(), + tool.description.as_deref(), + tool.input_schema.as_ref().or(tool.parameters.as_ref()), + ), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn parse_tools(json: &str) -> Vec { + serde_json::from_str(json).expect("tool fixture parses") + } + + #[test] + fn empty_tool_list_renders_an_empty_namespace() { + assert_eq!( + format_function_definitions(&[]).expect("empty tool list renders"), + "namespace functions {\n\n} // namespace functions" + ); + } + + #[test] + fn unnamed_tools_are_skipped_and_function_shape_takes_precedence() { + let tools = parse_tools( + r#"[ + {}, + {"name":""}, + {"name":"ignored","function":{"description":"missing a function name"}}, + {"name":"ping","description":""} + ]"#, + ); + + assert_eq!( + format_function_definitions(&tools).expect("tools render"), + "namespace functions {\n\ntype ping = () => any;\n\n} // namespace functions" + ); + } + + #[test] + fn empty_or_missing_properties_render_a_no_argument_function() { + let tools = parse_tools( + r#"[ + {"name":"missing","input_schema":{"type":"object"}}, + {"name":"empty","input_schema":{"type":"object","properties":{}}} + ]"#, + ); + + assert_eq!( + format_function_definitions(&tools).expect("tools render"), + concat!( + "namespace functions {\n\n", + "type missing = () => any;\n\n", + "type empty = () => any;\n\n", + "} // namespace functions" + ) + ); + } + + #[test] + fn anthropic_parameters_render_all_supported_types() { + let tools = parse_tools( + r#"[{ + "name":"inspect", + "description":"Inspect a value", + "input_schema":{ + "type":"object", + "properties":{ + "text":{"type":"string"}, + "count":{"type":"integer","description":"Number of attempts"}, + "ratio":{"type":"number"}, + "enabled":{"type":"boolean"}, + "nothing":{"type":"null"}, + "unknown":{"type":"custom"}, + "union":{"type":["string","null"]}, + "labels":{"type":"array","items":{"type":"string"}}, + "config":{"type":"object","properties":{"retries":{"type":"integer"}},"required":["retries"]}, + "mode":{"type":"string","enum":["fast",2]} + }, + "required":["text"] + } + }]"#, + ); + + assert_eq!( + format_function_definitions(&tools).expect("tool renders"), + concat!( + "namespace functions {\n\n", + "// Inspect a value\n", + "type inspect = (_: {\n", + "text: string,\n", + "// Number of attempts\n", + "count?: number,\n", + "ratio?: number,\n", + "enabled?: boolean,\n", + "nothing?: null,\n", + "unknown?: any,\n", + "union?: any,\n", + "labels?: string[],\n", + "config?: {\n", + " retries: number,\n", + "},\n", + "mode?: \"fast\" | \"2\",\n", + "}) => any;\n\n", + "} // namespace functions" + ) + ); + } + + #[test] + fn input_schema_takes_precedence_over_legacy_parameters() { + let tools = parse_tools( + r#"[{ + "name":"choose", + "input_schema":{"type":"object","properties":{"current":{"type":"string"}}}, + "parameters":{"type":"object","properties":{"legacy":{"type":"string"}}} + }]"#, + ); + + let rendered = format_function_definitions(&tools).expect("tool renders"); + assert!(rendered.contains("current?: string,")); + assert!(!rendered.contains("legacy")); + } + + #[test] + fn array_without_items_returns_an_error() { + let tools = parse_tools( + r#"[{"name":"broken","parameters":{"type":"object","properties":{"values":{"type":"array"}}}}]"#, + ); + + assert!(matches!( + format_function_definitions(&tools), + Err(Error::ArrayItems) + )); + } +} diff --git a/litellm-rust/crates/token-counter/src/types.rs b/litellm-rust/crates/token-counter/src/types.rs new file mode 100644 index 00000000000..d25554beaac --- /dev/null +++ b/litellm-rust/crates/token-counter/src/types.rs @@ -0,0 +1,238 @@ +use std::fmt; + +use indexmap::IndexMap; +use serde::de::{MapAccess, SeqAccess, Visitor}; +use serde::{Deserialize, Deserializer, Serialize}; +use serde_json::Number; + +use super::Error; + +/// The parts of a request body the host's budget counter reads. Chat and +/// Anthropic Messages bodies carry `messages`; completions carry `prompt`; +/// Responses and embeddings carry `input`; rerank carries `query` and +/// `documents`. The host checks key presence, not nullness, so an explicit +/// `null` is kept distinct from an absent key. Anything outside this shape is +/// declined so the host can fall back to its own counter instead of silently +/// miscounting. +#[derive(Clone, Debug, Deserialize, PartialEq)] +pub struct CountableRequest { + pub(crate) model: Option, + #[serde(default, deserialize_with = "present_messages")] + pub(crate) messages: Option>, + pub(crate) tools: Option>, + pub(crate) tool_choice: Option, + #[serde(default, deserialize_with = "present_text")] + pub(crate) prompt: Option, + #[serde(default, deserialize_with = "present_text")] + pub(crate) input: Option, + #[serde(default, deserialize_with = "present_text")] + pub(crate) query: Option, + #[serde(default, deserialize_with = "present_text")] + pub(crate) documents: Option, +} + +impl CountableRequest { + pub fn parse(body: &[u8]) -> Result { + serde_json::from_slice(body).map_err(Error::RequestParse) + } +} + +fn present_messages<'de, D: Deserializer<'de>>( + deserializer: D, +) -> Result>, D::Error> { + Option::>::deserialize(deserializer) + .map(|messages| Some(messages.unwrap_or_default())) +} + +fn present_text<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + TextValue::deserialize(deserializer).map(Some) +} + +/// Free-form JSON the host counts as text: strings and integers via `str()`, +/// objects via `json.dumps()`, lists flattened. Objects keep document order so +/// the dumped text matches Python byte for byte. +#[derive(Clone, Debug, PartialEq, Serialize)] +#[serde(untagged)] +pub(crate) enum TextValue { + Null, + Bool(bool), + Number(Number), + Text(String), + List(Vec), + Object(IndexMap), +} + +impl<'de> Deserialize<'de> for TextValue { + fn deserialize>(deserializer: D) -> Result { + deserializer.deserialize_any(TextValueVisitor) + } +} + +struct TextValueVisitor; + +impl<'de> Visitor<'de> for TextValueVisitor { + type Value = TextValue; + + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("a JSON value") + } + + fn visit_unit(self) -> Result { + Ok(TextValue::Null) + } + + fn visit_none(self) -> Result { + Ok(TextValue::Null) + } + + fn visit_bool(self, value: bool) -> Result { + Ok(TextValue::Bool(value)) + } + + fn visit_i64(self, value: i64) -> Result { + Ok(TextValue::Number(value.into())) + } + + fn visit_u64(self, value: u64) -> Result { + Ok(TextValue::Number(value.into())) + } + + fn visit_f64(self, value: f64) -> Result + where + E: serde::de::Error, + { + Number::from_f64(value) + .map(TextValue::Number) + .ok_or_else(|| E::custom("non-finite JSON number")) + } + + fn visit_str(self, value: &str) -> Result { + Ok(TextValue::Text(value.to_owned())) + } + + fn visit_string(self, value: String) -> Result { + Ok(TextValue::Text(value)) + } + + fn visit_seq(self, mut sequence: A) -> Result + where + A: SeqAccess<'de>, + { + let mut items = Vec::with_capacity(sequence.size_hint().unwrap_or(0)); + while let Some(item) = sequence.next_element()? { + items.push(item); + } + Ok(TextValue::List(items)) + } + + fn visit_map(self, mut map: A) -> Result + where + A: MapAccess<'de>, + { + let mut entries = IndexMap::with_capacity(map.size_hint().unwrap_or(0)); + while let Some((key, value)) = map.next_entry()? { + entries.insert(key, value); + } + Ok(TextValue::Object(entries)) + } +} + +/// Python counts every string-valued key of a message, so any key beyond these +/// makes the shape unsupported rather than silently uncounted. +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(deny_unknown_fields)] +pub(crate) struct Message { + pub(crate) role: Option, + pub(crate) name: Option, + pub(crate) content: Option, +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub(crate) enum MessageContent { + Text(String), + Blocks(Vec), +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub(crate) enum ContentItem { + Text(String), + Block(ContentBlock), +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(tag = "type")] +pub(crate) enum ContentBlock { + #[serde(rename = "text")] + Text { text: String }, + #[serde(rename = "thinking")] + Thinking { thinking: String }, + #[serde(rename = "tool_reference")] + ToolReference { tool_name: Option }, + /// Images, documents, files and tool use/result blocks price through + /// Python-only helpers, so they stay on the Python counter. + #[serde(other)] + Unsupported, +} + +/// Either the OpenAI `{"type": "function", "function": {...}}` shape or the +/// Anthropic `{"name", "description", "input_schema"}` shape. +#[derive(Clone, Debug, Deserialize, PartialEq)] +pub(crate) struct ToolDefinition { + pub(crate) function: Option, + pub(crate) name: Option, + pub(crate) description: Option, + pub(crate) input_schema: Option, + pub(crate) parameters: Option, +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +pub(crate) struct FunctionDefinition { + pub(crate) name: Option, + pub(crate) description: Option, + pub(crate) parameters: Option, +} + +#[derive(Clone, Debug, Default, Deserialize, PartialEq)] +pub(crate) struct Schema { + #[serde(rename = "type")] + pub(crate) schema_type: Option, + pub(crate) description: Option, + #[serde(rename = "enum")] + pub(crate) enum_values: Option>, + pub(crate) items: Option>, + pub(crate) properties: Option>, + pub(crate) required: Option>, +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub(crate) enum SchemaType { + Name(String), + Union(Vec), +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub(crate) enum EnumValue { + Text(String), + Integer(i64), +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub(crate) enum ToolChoice { + Mode(String), + Named(NamedToolChoice), +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +pub(crate) struct NamedToolChoice { + pub(crate) function: NamedFunction, +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +pub(crate) struct NamedFunction { + pub(crate) name: String, +} diff --git a/litellm-rust/crates/token-counter/src/unicode_classes.rs b/litellm-rust/crates/token-counter/src/unicode_classes.rs new file mode 100644 index 00000000000..405cee34949 --- /dev/null +++ b/litellm-rust/crates/token-counter/src/unicode_classes.rs @@ -0,0 +1,73 @@ +use std::cmp::Ordering; +use std::sync::LazyLock; + +use tokenizers::utils::SysRegex; + +struct Ranges(Box<[(u32, u32)]>); + +pub(super) struct UnicodeClasses { + letters: Ranges, + numbers: Ranges, + spaces: Ranges, +} + +static CLASSES: LazyLock> = LazyLock::new(|| { + let scalars: String = (0..=u32::from(char::MAX)) + .filter_map(char::from_u32) + .collect(); + Some(UnicodeClasses { + letters: Ranges::load(r"\p{L}+", &scalars)?, + numbers: Ranges::load(r"\p{N}+", &scalars)?, + spaces: Ranges::load(r"\s+", &scalars)?, + }) +}); + +impl Ranges { + fn load(pattern: &str, scalars: &str) -> Option { + let regex = SysRegex::new(pattern).ok()?; + let ranges = regex + .find_iter(scalars) + .map(|(start, end)| { + let matched = scalars.get(start..end)?; + Some(( + u32::from(matched.chars().next()?), + u32::from(matched.chars().next_back()?), + )) + }) + .collect::>>()?; + Some(Self(ranges)) + } + + fn contains(&self, character: char) -> bool { + let code = u32::from(character); + self.0 + .binary_search_by(|(low, high)| { + if *high < code { + Ordering::Less + } else if *low > code { + Ordering::Greater + } else { + Ordering::Equal + } + }) + .is_ok() + } +} + +impl UnicodeClasses { + pub(super) fn get() -> Option<&'static Self> { + CLASSES.as_ref() + } + + pub(super) fn is_letter(&self, character: char) -> bool { + self.letters.contains(character) + } + + pub(super) fn is_number(&self, character: char) -> bool { + self.numbers.contains(character) + } + + pub(super) fn is_space(&self, character: char) -> bool { + self.spaces.contains(character) + } +} diff --git a/litellm-rust/crates/token-counter/tests/token_counter.rs b/litellm-rust/crates/token-counter/tests/token_counter.rs new file mode 100644 index 00000000000..0e71a5c03cf --- /dev/null +++ b/litellm-rust/crates/token-counter/tests/token_counter.rs @@ -0,0 +1,176 @@ +use rstest::rstest; + +use litellm_token_counter::{CountableRequest, Error, InputTokenCount, TokenCounter}; + +/// Expected counts are pinned from `litellm.token_counter(model="claude-sonnet-4-5", ...)` +/// so this test also guards Python parity. +fn counter() -> TokenCounter { + let path = concat!( + env!("CARGO_MANIFEST_DIR"), + "/../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json" + ); + let json = std::fs::read_to_string(path).expect("anthropic tokenizer json is in the repo"); + TokenCounter::from_json(&json).expect("anthropic tokenizer loads") +} + +const SIMPLE: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"Hello, how are you today?"}]}"#; + +const BLOCKS_AND_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5","messages":[ + {"role":"system","content":"You are a terse assistant."}, + {"role":"user","name":"alice","content":[ + {"type":"text","text":"Summarise this paragraph about ships and harbours."}, + "plain string item", + {"type":"thinking","thinking":"pondering"}, + {"type":"tool_reference","tool_name":"get_weather"}]}, + {"role":"assistant","content":[{"type":"text","text":"Sure.","cache_control":{"type":"ephemeral"}}]}]}"#; + +const TOOLS_OPENAI: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"weather?"}], + "tools":[ + {"type":"function","function":{"name":"get_weather","description":"Get weather","parameters":{ + "type":"object", + "properties":{ + "location":{"type":"string","description":"City name"}, + "unit":{"type":"string","enum":["celsius","fahrenheit"]}, + "days":{"type":"integer"}, + "tags":{"type":"array","items":{"type":"string"}}, + "opts":{"type":"object","properties":{"verbose":{"type":"boolean"},"level":{"type":"integer","enum":[1,2]}},"required":["verbose"]}, + "anything":{}}, + "required":["location"]}}}, + {"type":"function","function":{"name":"noop"}}], + "tool_choice":{"type":"function","function":{"name":"get_weather"}}}"#; + +const TOOLS_ANTHROPIC_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5", + "messages":[{"role":"system","content":"sys"},{"role":"user","content":"weather?"}], + "tools":[{"name":"get_weather","description":"Get weather","input_schema":{ + "type":"object","properties":{"location":{"type":["string","null"]}},"required":["location"]}}], + "tool_choice":"none"}"#; + +const COMPLETIONS_PROMPT: &str = + r#"{"model":"claude-sonnet-4-5","prompt":"Write a haiku about ships."}"#; + +const COMPLETIONS_PROMPT_LIST: &str = + r#"{"model":"claude-sonnet-4-5","prompt":["first prompt","second prompt"]}"#; + +const RESPONSES_INPUT: &str = r#"{"model":"claude-sonnet-4-5","input":[ + {"role":"user","content":[{"type":"input_text","text":"Summarise caf\u00e9 menus, na\u00efve \u2014 ok? \"quoted\"\n"}]}, + {"role":"assistant","content":"Sure."}],"instructions":"be terse"}"#; + +const EMBEDDINGS_TOKEN_IDS: &str = + r#"{"model":"claude-sonnet-4-5","input":[[101,2023,5],[7]],"encoding_format":"float"}"#; + +const RERANK: &str = r#"{"model":"claude-sonnet-4-5","query":"best harbour", + "documents":["doc one",{"text":"doc two","title":"T","n":3,"ok":true,"none":null,"tags":["a","b"]}]}"#; + +/// Expected counts are pinned from +/// `litellm.proxy.spend_tracking.budget_reservation._count_input_tokens(body, "claude-sonnet-4-5")`. +#[rstest] +#[case::text_only(SIMPLE, 14)] +#[case::content_blocks_name_and_system(BLOCKS_AND_SYSTEM, 45)] +#[case::openai_tools_named_choice(TOOLS_OPENAI, 123)] +#[case::anthropic_tools_system_discount_choice_none(TOOLS_ANTHROPIC_SYSTEM, 53)] +#[case::completions_prompt(COMPLETIONS_PROMPT, 7)] +#[case::completions_prompt_list(COMPLETIONS_PROMPT_LIST, 4)] +#[case::responses_input_items(RESPONSES_INPUT, 62)] +#[case::embeddings_token_ids(EMBEDDINGS_TOKEN_IDS, 5)] +#[case::rerank_query_and_documents(RERANK, 41)] +fn count_request_matches_python_token_counter(#[case] body: &str, #[case] expected: usize) { + let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses"); + let count = counter().count_request(&request).expect("fixture counts"); + assert_eq!( + count, + InputTokenCount { + model: Some("claude-sonnet-4-5".to_string()), + input_tokens: expected, + } + ); +} + +#[rstest] +#[case::null_messages_win_over_prompt(r#"{"model":"m","messages":null,"prompt":"ignored"}"#, 3)] +#[case::model_from_route(r#"{"prompt":"hi"}"#, 1)] +#[case::bools_and_ints_use_python_str(r#"{"model":"m","prompt":[true,false,42]}"#, 3)] +#[case::null_prompt_counts_zero(r#"{"model":"m","prompt":null}"#, 0)] +fn key_presence_follows_python(#[case] body: &str, #[case] expected: usize) { + let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses"); + let count = counter().count_request(&request).expect("fixture counts"); + assert_eq!(count.input_tokens, expected); +} + +#[rstest] +#[case::not_json(b"not json" as &[u8])] +#[case::messages_not_a_list(br#"{"model":"m","messages":"hi"}"#)] +#[case::message_with_tool_calls( + br#"{"model":"m","messages":[{"role":"assistant","tool_calls":[{"id":"1","type":"function","function":{"name":"f","arguments":"{}"}}]}]}"# +)] +#[case::dict_content( + br#"{"model":"m","messages":[{"role":"user","content":{"type":"text","text":"x"}}]}"# +)] +#[case::float_enum( + br#"{"model":"m","messages":[],"tools":[{"name":"f","input_schema":{"type":"object","properties":{"x":{"type":"number","enum":[1.5]}}}}]}"# +)] +#[case::anthropic_tool_choice_without_function( + br#"{"model":"m","messages":[],"tool_choice":{"type":"auto"}}"# +)] +fn shapes_outside_the_mirror_are_declined_at_parse(#[case] body: &[u8]) { + assert!(matches!( + CountableRequest::parse(body), + Err(Error::RequestParse(_)) + )); +} + +#[rstest] +#[case::no_countable_input(br#"{"model":"m","instructions":"hi"}"# as &[u8])] +#[case::float_prompt(br#"{"model":"m","prompt":1.5}"#)] +#[case::float_inside_document(br#"{"model":"m","documents":[{"score":0.5}]}"#)] +#[case::image_block( + br#"{"model":"m","messages":[{"role":"user","content":[{"type":"image","source":{"type":"base64","media_type":"image/png","data":"AA=="}}]}]}"# +)] +#[case::tool_result_block( + br#"{"model":"m","messages":[{"role":"user","content":[{"type":"tool_result","tool_use_id":"1","content":"ok"}]}]}"# +)] +#[case::array_without_items( + br#"{"model":"m","messages":[],"tools":[{"name":"f","input_schema":{"type":"object","properties":{"x":{"type":"array"}}}}]}"# +)] +fn shapes_outside_the_mirror_are_declined_at_count(#[case] body: &[u8]) { + let request = CountableRequest::parse(body).expect("shape parses"); + assert!(matches!( + counter().count_request(&request), + Err(Error::MissingInput | Error::FloatText | Error::ContentBlock | Error::ArrayItems) + )); +} + +#[test] +fn tool_choice_and_system_discount_change_the_count() { + let counter = counter(); + let count = |body: &str| { + counter + .count_request(&CountableRequest::parse(body.as_bytes()).expect("parses")) + .expect("counts") + .input_tokens + }; + let base = count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#); + assert_eq!( + count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"none"}"#), + base + 1 + ); + assert_eq!( + count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"auto"}"#), + base + ); + let with_tools = count( + r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[{"name":"f"}]}"#, + ); + let with_tools_and_system = count( + r#"{"model":"m","messages":[{"role":"system","content":"hi"}],"tools":[{"name":"f"}]}"#, + ); + assert_eq!(with_tools - with_tools_and_system, 4); + assert_eq!( + count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[]}"#), + base + ); +} + +#[test] +fn loading_a_bad_tokenizer_is_a_load_error() { + assert!(matches!(TokenCounter::from_json("{}"), Err(Error::Load(_)))); +} diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 64c0da1c28f..20ab9904f46 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -91,6 +91,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_query_params, _safe_set_request_parsed_body, populate_request_with_path_params, + read_raw_json_body, ) from litellm.proxy.common_utils.model_listing_utils import claude_code_requested_group from litellm.proxy.common_utils.realtime_utils import _realtime_request_body @@ -2689,6 +2690,7 @@ async def _run_centralized_common_checks( await _reserve_budget_after_common_checks( user_api_key_auth_obj=user_api_key_auth_obj, + request=request, request_data=request_data, route=route, llm_router=llm_router, @@ -2724,6 +2726,7 @@ async def _reserve_budget_after_common_checks( general_settings: dict, end_user_id: str | None = None, end_user_object: LiteLLM_EndUserTable | None = None, + request: Request | None = None, ) -> None: user_api_key_auth_obj.budget_reservation = None if skip_budget_checks: @@ -2749,6 +2752,7 @@ async def _reserve_budget_after_common_checks( end_user_object=end_user_object, apply_user_budget_to_team_keys=general_settings.get("apply_user_budget_to_team_keys") is True, fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True, + raw_body=await read_raw_json_body(request=request), ) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 635cbf3ba9b..845589aee7a 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -213,6 +213,18 @@ async def _read_request_body(request: Request | None) -> dict: return {} +async def read_raw_json_body(request: Request | None) -> bytes | None: + if request is None or _safe_get_request_parsed_body(request=request) is None: + return None + content_type: Final = _safe_get_request_headers(request=request).get("content-type", "") + if _is_form_content_type(content_type): + return None + try: + return await request.body() + except RuntimeError: + return None + + def _safe_get_request_parsed_body(request: Request | None) -> dict | None: if request is None: return None diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 8d4d72d9974..ef21551ca93 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -34,6 +34,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( ) from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.router import Router +from litellm.rust_bridge.token_counter import count_anthropic_input_tokens, uses_anthropic_tokenizer from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget from litellm.types.router import DeploymentTypedDict @@ -210,6 +211,7 @@ async def reserve_budget_for_request( end_user_object: object = None, apply_user_budget_to_team_keys: bool = False, fail_closed_budget_enforcement: bool = False, + raw_body: bytes | None = None, ) -> dict | None: if valid_token is None or not RouteChecks.is_llm_api_route(route=route): return None @@ -237,6 +239,7 @@ async def reserve_budget_for_request( request_body=request_body, route=route, llm_router=llm_router, + raw_body=raw_body, ) current_spend_by_counter_key: Final[dict[str, float]] = {} @@ -1356,24 +1359,46 @@ async def count_request_input_tokens( request_body: dict, route: str, llm_router: Router | None, + raw_body: bytes | None = None, ) -> Mapping[str, int]: """Input-token count per candidate model, counted once per request. Tokenizing is the reservation path's dominant CPU cost and is O(prompt), so counting a large prompt inline stalls every other request on the worker. - Large prompts are counted in a worker thread, and the counts are reused by - both the max-cost and the input-cost estimate. + Models on the Anthropic tokenizer are counted from the raw body by the Rust + bridge when it is enabled, which parses and tokenizes with the GIL released. + Everything it declines is counted in Python, large prompts in a worker + thread. The counts are reused by both the max-cost and the input-cost + estimate. """ models: Final = _get_request_models(request_body=request_body, route=route, llm_router=llm_router) if not models: return MappingProxyType({}) - if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS: - return _count_input_tokens_for_models(request_body=request_body, models=models) - return await asyncio.to_thread( - _count_input_tokens_for_models, - request_body=request_body, - models=models, + rust_count: Final = ( + await count_anthropic_input_tokens(raw_body) + if raw_body is not None and any(uses_anthropic_tokenizer(model) for model in models) + else None ) + rust_counts: Final = MappingProxyType( + { + model: rust_count.input_tokens + for model in models + if rust_count is not None and uses_anthropic_tokenizer(model) + } + ) + python_models: Final = tuple(model for model in models if model not in rust_counts) + if not python_models: + return rust_counts + python_counts: Final = ( + _count_input_tokens_for_models(request_body=request_body, models=python_models) + if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS + else await asyncio.to_thread( + _count_input_tokens_for_models, + request_body=request_body, + models=python_models, + ) + ) + return MappingProxyType({**rust_counts, **python_counts}) def _count_input_tokens_for_models( diff --git a/litellm/rust_bridge/token_counter.py b/litellm/rust_bridge/token_counter.py new file mode 100644 index 00000000000..ee755b317d8 --- /dev/null +++ b/litellm/rust_bridge/token_counter.py @@ -0,0 +1,79 @@ +"""Thin Python wrapper for the native Rust input token counter.""" + +from __future__ import annotations + +from collections.abc import Awaitable +from dataclasses import dataclass +from functools import lru_cache +from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables + +from pydantic import TypeAdapter + +import litellm +from litellm._logging import verbose_logger +from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.configuration import rust_enabled +from litellm.rust_bridge.runtime import BridgeErrorContext, RustHandled, aattempt +from litellm.utils import claude_json_str +from litellm.utils import uses_anthropic_tokenizer as _python_uses_anthropic_tokenizer + + +class RustTokenCounter(Protocol): + def acount_request(self, body: bytes) -> Awaitable[object]: + raise NotImplementedError + + +class RustTokenCounterFactory(Protocol): + def __call__(self, tokenizer_json: str) -> RustTokenCounter: + raise NotImplementedError + + +@dataclass(frozen=True, slots=True) +class InputTokenCount: + model: str | None + input_tokens: int + + +_INPUT_TOKEN_COUNT: Final = TypeAdapter(InputTokenCount) + + +def _as_factory(value: object) -> RustTokenCounterFactory | None: + return ( + cast( # cast-ok: native extension protocol is runtime-defined + RustTokenCounterFactory, value + ) + if callable(value) + else None + ) + + +TOKEN_COUNTER: Final = NativeBinding("TokenCounter", validate=_as_factory) + + +def uses_anthropic_tokenizer(model: str) -> bool: + if litellm.disable_token_counter is True or litellm.disable_hf_tokenizer_download is True: + return False + return _python_uses_anthropic_tokenizer(model) + + +@lru_cache(maxsize=4) +def _anthropic_counter(factory: RustTokenCounterFactory) -> RustTokenCounter: + return factory(claude_json_str) + + +async def count_anthropic_input_tokens(body: bytes) -> InputTokenCount | None: + if not rust_enabled(): + return None + factory: Final = TOKEN_COUNTER.load() + if factory is None: + return None + try: + attempt: Final = await aattempt( + native_call=lambda: _anthropic_counter(factory).acount_request(body), + adapt=_INPUT_TOKEN_COUNT.validate_python, + context=BridgeErrorContext(route="token_counter", provider="anthropic", model=""), + ) + except (RuntimeError, ValueError) as error: + verbose_logger.debug("Rust token counter failed, counting in Python: %s", error) + return None + return attempt.value if isinstance(attempt, RustHandled) else None diff --git a/litellm/utils.py b/litellm/utils.py index 917af2b89d4..504de08fa91 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2197,13 +2197,17 @@ def _return_openai_tokenizer(model: str) -> SelectTokenizerResponse: return {"type": "openai_tokenizer", "tokenizer": _get_default_encoding()} +def uses_anthropic_tokenizer(model: str) -> bool: + return model in litellm.anthropic_models and "claude-3" not in model + + def _return_huggingface_tokenizer(model: str) -> SelectTokenizerResponse | None: if model in litellm.cohere_models and "command-r" in model: # cohere cohere_tokenizer: Final = Tokenizer.from_pretrained("Xenova/c4ai-command-r-v01-tokenizer") return {"type": "huggingface_tokenizer", "tokenizer": cohere_tokenizer} # anthropic - elif model in litellm.anthropic_models and "claude-3" not in model: + elif uses_anthropic_tokenizer(model): claude_tokenizer: Final = Tokenizer.from_str(claude_json_str) return {"type": "huggingface_tokenizer", "tokenizer": claude_tokenizer} # llama2 diff --git a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py index 011571a37e0..bc4e756eb65 100644 --- a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py @@ -26,9 +26,58 @@ from litellm.proxy.common_utils.http_parsing_utils import ( get_tags_from_request_body, numeric_form_fields, populate_request_with_path_params, + read_raw_json_body, ) +def _starlette_request(body: bytes, content_type: str) -> Request: + scope = { + "type": "http", + "method": "POST", + "path": "/v1/messages", + "headers": [(b"content-type", content_type.encode())], + "query_string": b"", + } + chunks = iter((body,)) + + async def receive(): + return {"type": "http.request", "body": next(chunks, b""), "more_body": False} + + return Request(scope, receive) + + +@pytest.mark.asyncio +async def test_read_raw_json_body_returns_the_bytes_the_parsed_body_came_from(): + body = b'{"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "hi"}]}' + request = _starlette_request(body, "application/json") + + assert await _read_request_body(request) == orjson.loads(body) + assert await read_raw_json_body(request) == body + + +@pytest.mark.asyncio +async def test_read_raw_json_body_is_none_until_the_body_has_been_parsed(): + request = _starlette_request(b'{"model": "claude-sonnet-4-5"}', "application/json") + + assert await read_raw_json_body(request) is None + assert await read_raw_json_body(None) is None + + +@pytest.mark.asyncio +async def test_read_raw_json_body_is_none_for_form_bodies(): + request = _starlette_request(b"model=claude-sonnet-4-5", "application/x-www-form-urlencoded") + + assert await _read_request_body(request) == {"model": "claude-sonnet-4-5"} + assert await read_raw_json_body(request) is None + + +@pytest.mark.asyncio +async def test_read_raw_json_body_is_none_for_a_request_that_only_mocks_the_parsed_body_path(): + mock_request = MagicMock() + + assert await read_raw_json_body(mock_request) is None + + @pytest.mark.asyncio async def test_request_body_caching(): """ diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py index e7dc5af9bfd..ee062b14b96 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py @@ -1,16 +1,23 @@ +import json import math from typing import Final import pytest import litellm -import litellm.proxy.proxy_server as proxy_server from litellm.caching import DualCache +from litellm.proxy import proxy_server from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache -from litellm.proxy.spend_tracking.budget_reservation import estimate_request_max_cost, reserve_budget_for_request +from litellm.proxy.spend_tracking.budget_reservation import ( + count_request_input_tokens, + estimate_request_max_cost, + reserve_budget_for_request, +) from litellm.proxy.utils import ProxyLogging from litellm.router import Router +from litellm.rust_bridge import bindings, configuration +from litellm.rust_bridge import token_counter as rust_token_counter from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo TOKEN_COUNTING_ROUTES: Final = ( @@ -197,3 +204,128 @@ def test_deployment_pricing_update_invalidates_cached_estimate() -> None: after: Final = estimate_request_max_cost(request_body=TIERED_BODY, route="/chat/completions", llm_router=router) assert after is not None assert math.isclose(after, before * 1000) + + +ANTHROPIC_TOKENIZER_MODEL: Final = "claude-sonnet-4-5-20250929" +RUST_COUNTED_BODY: Final = {"model": ANTHROPIC_TOKENIZER_MODEL, "max_tokens": 16, "messages": ANTHROPIC_MESSAGES} +RUST_INPUT_TOKENS: Final = 4_321 + + +class _FakeDeclined(Exception): + pass + + +class _FakeUpstream(Exception): + pass + + +class _FakeNative: + RustBridgeDeclined = _FakeDeclined + RustUpstreamError = _FakeUpstream + + +class _RecordingCounter: + bodies: Final[list[bytes]] = [] + + def __init__(self, tokenizer_json: str) -> None: + pass + + async def acount_request(self, body: bytes) -> object: + self.bodies.append(body) + return {"model": ANTHROPIC_TOKENIZER_MODEL, "input_tokens": RUST_INPUT_TOKENS} + + +class _DecliningCounter: + def __init__(self, tokenizer_json: str) -> None: + pass + + async def acount_request(self, body: bytes) -> object: + raise _FakeDeclined("unsupported content block") + + +@pytest.fixture +def rust_counter(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative()) + rust_token_counter._anthropic_counter.cache_clear() + configuration.reset_rust_configuration() + _RecordingCounter.bodies.clear() + yield + rust_token_counter.TOKEN_COUNTER.reset() + rust_token_counter._anthropic_counter.cache_clear() + configuration.reset_rust_configuration() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("route", "request_body"), + ( + ("/v1/messages", RUST_COUNTED_BODY), + ("/v1/chat/completions", {"model": ANTHROPIC_TOKENIZER_MODEL, "messages": ANTHROPIC_MESSAGES}), + ("/v1/completions", {"model": ANTHROPIC_TOKENIZER_MODEL, "prompt": "hi"}), + ("/v1/responses", {"model": ANTHROPIC_TOKENIZER_MODEL, "input": "hi"}), + ("/v1/embeddings", {"model": ANTHROPIC_TOKENIZER_MODEL, "input": ["hi"]}), + ("/v1/rerank", {"model": ANTHROPIC_TOKENIZER_MODEL, "query": "hi", "documents": ["a"]}), + ), +) +async def test_rust_count_replaces_python_tokenizing_on_every_llm_route( + rust_counter: None, route: str, request_body: dict +) -> None: + litellm.rust(True) + rust_token_counter.TOKEN_COUNTER.override(_RecordingCounter) + raw_body: Final = json.dumps(request_body).encode() + + counts: Final = await count_request_input_tokens( + request_body=request_body, route=route, llm_router=None, raw_body=raw_body + ) + + assert dict(counts) == {ANTHROPIC_TOKENIZER_MODEL: RUST_INPUT_TOKENS} + assert _RecordingCounter.bodies == [raw_body] + + +@pytest.mark.asyncio +async def test_rust_decline_falls_back_to_python_count(rust_counter: None) -> None: + litellm.rust(True) + rust_token_counter.TOKEN_COUNTER.override(_DecliningCounter) + python_counts: Final = await count_request_input_tokens( + request_body=RUST_COUNTED_BODY, route="/v1/messages", llm_router=None + ) + + counts: Final = await count_request_input_tokens( + request_body=RUST_COUNTED_BODY, + route="/v1/messages", + llm_router=None, + raw_body=json.dumps(RUST_COUNTED_BODY).encode(), + ) + + assert dict(counts) == dict(python_counts) + assert counts[ANTHROPIC_TOKENIZER_MODEL] != RUST_INPUT_TOKENS + + +@pytest.mark.asyncio +async def test_disabled_rust_never_sees_the_raw_body(rust_counter: None) -> None: + litellm.rust(False) + rust_token_counter.TOKEN_COUNTER.override(_RecordingCounter) + + counts: Final = await count_request_input_tokens( + request_body=RUST_COUNTED_BODY, + route="/v1/messages", + llm_router=None, + raw_body=json.dumps(RUST_COUNTED_BODY).encode(), + ) + + assert _RecordingCounter.bodies == [] + assert counts[ANTHROPIC_TOKENIZER_MODEL] != RUST_INPUT_TOKENS + + +@pytest.mark.asyncio +async def test_non_anthropic_tokenizer_models_stay_in_python(rust_counter: None) -> None: + litellm.rust(True) + rust_token_counter.TOKEN_COUNTER.override(_RecordingCounter) + body: Final = {"model": "gpt-4o", "messages": ANTHROPIC_MESSAGES} + + counts: Final = await count_request_input_tokens( + request_body=body, route="/v1/chat/completions", llm_router=None, raw_body=json.dumps(body).encode() + ) + + assert _RecordingCounter.bodies == [] + assert counts["gpt-4o"] != RUST_INPUT_TOKENS diff --git a/tests/test_litellm/rust_bridge/test_token_counter.py b/tests/test_litellm/rust_bridge/test_token_counter.py new file mode 100644 index 00000000000..2df91204390 --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_token_counter.py @@ -0,0 +1,242 @@ +"""Tests for the Rust input token counter bridge. + +The native factory is dependency-injected through ``TOKEN_COUNTER.override`` +so the fallback cases run without the compiled extension present. The parity +cases need the extension and are skipped when it is not built. +""" + +from __future__ import annotations + +import json +from typing import Final + +import pytest + +import litellm +from litellm.proxy.spend_tracking.budget_reservation import _count_input_tokens +from litellm.rust_bridge import bindings, configuration +from litellm.rust_bridge import token_counter as bridge + +MODEL: Final = "claude-sonnet-4-5-20250929" +BODY: Final = json.dumps({"model": MODEL, "messages": [{"role": "user", "content": "hello"}]}).encode() + + +class _FakeDeclined(Exception): + pass + + +class _FakeUpstream(Exception): + pass + + +class _FakeNative: + RustBridgeDeclined = _FakeDeclined + RustUpstreamError = _FakeUpstream + + +class _RecordingCounter: + def __init__(self, tokenizer_json: str) -> None: + self.tokenizer_json = tokenizer_json + self.bodies: list[bytes] = [] + + async def acount_request(self, body: bytes) -> object: + self.bodies.append(body) + return {"model": MODEL, "input_tokens": 42} + + +class _DecliningCounter: + def __init__(self, tokenizer_json: str) -> None: + pass + + async def acount_request(self, body: bytes) -> object: + raise _FakeDeclined("request has no messages") + + +class _FailingCounter: + def __init__(self, tokenizer_json: str) -> None: + pass + + async def acount_request(self, body: bytes) -> object: + raise RuntimeError("encode failed") + + +@pytest.fixture(autouse=True) +def _reset_bridge(monkeypatch: pytest.MonkeyPatch): + bridge.TOKEN_COUNTER.reset() + bridge._anthropic_counter.cache_clear() + configuration.reset_rust_configuration() + monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative()) + yield + bridge.TOKEN_COUNTER.reset() + bridge._anthropic_counter.cache_clear() + configuration.reset_rust_configuration() + + +@pytest.mark.asyncio +async def test_disabled_bridge_never_constructs_a_counter() -> None: + constructed: list[str] = [] + + def factory(tokenizer_json: str) -> _RecordingCounter: + constructed.append(tokenizer_json) + return _RecordingCounter(tokenizer_json) + + litellm.rust(False) + bridge.TOKEN_COUNTER.override(factory) + + assert await bridge.count_anthropic_input_tokens(BODY) is None + assert constructed == [] + + +@pytest.mark.asyncio +async def test_enabled_bridge_returns_typed_count_and_reuses_one_counter() -> None: + counters: list[_RecordingCounter] = [] + + def factory(tokenizer_json: str) -> _RecordingCounter: + counter = _RecordingCounter(tokenizer_json) + counters.append(counter) + return counter + + litellm.rust(True) + bridge.TOKEN_COUNTER.override(factory) + + first: Final = await bridge.count_anthropic_input_tokens(BODY) + second: Final = await bridge.count_anthropic_input_tokens(BODY) + + assert first == bridge.InputTokenCount(model=MODEL, input_tokens=42) + assert second == first + assert len(counters) == 1 + assert counters[0].bodies == [BODY, BODY] + assert json.loads(counters[0].tokenizer_json)["model"]["type"] == "BPE" + + +@pytest.mark.asyncio +async def test_missing_native_module_falls_back(monkeypatch: pytest.MonkeyPatch) -> None: + litellm.rust(True) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) + + assert await bridge.count_anthropic_input_tokens(BODY) is None + + +@pytest.mark.asyncio +async def test_declined_request_falls_back() -> None: + litellm.rust(True) + bridge.TOKEN_COUNTER.override(_DecliningCounter) + + assert await bridge.count_anthropic_input_tokens(BODY) is None + + +@pytest.mark.asyncio +async def test_runtime_failure_falls_back() -> None: + litellm.rust(True) + bridge.TOKEN_COUNTER.override(_FailingCounter) + + assert await bridge.count_anthropic_input_tokens(BODY) is None + + +@pytest.mark.parametrize( + ("model", "expected"), + ((MODEL, True), ("claude-3-5-sonnet-20241022", False), ("gpt-4o", False), ("my-router-alias", False)), +) +def test_uses_anthropic_tokenizer_mirrors_python_tokenizer_selection(model: str, expected: bool) -> None: + assert bridge.uses_anthropic_tokenizer(model) is expected + + +@pytest.mark.parametrize("flag", ("disable_hf_tokenizer_download", "disable_token_counter")) +def test_uses_anthropic_tokenizer_respects_python_opt_outs(monkeypatch: pytest.MonkeyPatch, flag: str) -> None: + monkeypatch.setattr(litellm, flag, True) + + assert bridge.uses_anthropic_tokenizer(MODEL) is False + + +PARITY_REQUESTS: Final[tuple[dict[str, object], ...]] = ( + {"model": MODEL, "messages": [{"role": "user", "content": "Hello, how are you today?"}]}, + { + "model": MODEL, + "messages": [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "name": "bob", "content": [{"type": "text", "text": "Summarize this."}]}, + {"role": "assistant", "content": "Sure."}, + ], + }, + { + "model": MODEL, + "messages": [{"role": "user", "content": "weather in sf?"}], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": { + "type": "object", + "properties": { + "city": {"type": "string", "description": "City"}, + "unit": {"type": "string", "enum": ["c", "f"]}, + }, + "required": ["city"], + }, + }, + } + ], + "tool_choice": {"type": "function", "function": {"name": "get_weather"}}, + }, + { + "model": MODEL, + "messages": [{"role": "user", "content": "x " * 20_000}], + }, + {"model": MODEL, "prompt": "Write a haiku about ships.", "max_tokens": 20}, + {"model": MODEL, "prompt": ["first prompt", "second prompt"]}, + { + "model": MODEL, + "instructions": "be terse", + "input": [ + {"role": "user", "content": [{"type": "input_text", "text": "Summarise caf\u00e9 menus \u2014 \"ok\"?\n"}]}, + {"role": "assistant", "content": "Sure."}, + ], + }, + {"model": MODEL, "input": "a single embedding string"}, + {"model": MODEL, "input": [[101, 2023, 5], [7]], "encoding_format": "float"}, + {"model": MODEL, "query": "best harbour", "documents": ["doc one", {"text": "doc two", "title": "T", "n": 3}]}, + {"model": MODEL, "messages": None, "prompt": "messages key wins even when null"}, + {"prompt": "model comes from the route"}, +) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("request_body", PARITY_REQUESTS) +async def test_native_count_matches_python_budget_counter( + monkeypatch: pytest.MonkeyPatch, request_body: dict[str, object] +) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) + litellm.rust(True) + + rust_count: Final = await bridge.count_anthropic_input_tokens(json.dumps(request_body).encode()) + python_count: Final = _count_input_tokens(request_body=request_body, model=MODEL) + + assert rust_count is not None + assert rust_count.model == request_body.get("model") + assert rust_count.input_tokens == python_count + + +DECLINED_REQUESTS: Final[tuple[dict[str, object], ...]] = ( + { + "model": MODEL, + "messages": [{"role": "user", "content": [{"type": "image_url", "image_url": {"url": "data:image/png;base64,AA"}}]}], + }, + {"model": MODEL, "prompt": 1.5}, + {"model": MODEL, "documents": [{"score": 0.5}]}, + {"model": MODEL, "file": "audio.mp3"}, +) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("request_body", DECLINED_REQUESTS) +async def test_native_declines_shapes_python_prices_differently( + monkeypatch: pytest.MonkeyPatch, request_body: dict[str, object] +) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) + litellm.rust(True) + + assert await bridge.count_anthropic_input_tokens(json.dumps(request_body).encode()) is None From c59fc6dc28cfe2ede38603f6c6f9935049acf1e2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 14:25:30 -0700 Subject: [PATCH 56/60] fix(mcp): reject initialize with 403 when the key grants no MCP servers (#40616) * fix(mcp): reject initialize with 403 when the key grants no MCP servers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): e2e expects 403 initialize for a key with no MCP servers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): mention IP filtering in the no-servers initialize denial and keep zero-grant tool coverage Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: mateo Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/server.py | 28 +++++ tests/mcp_tests/test_proxy_mcp_e2e.py | 113 ++++++++++++------ .../mcp_server/test_mcp_server.py | 87 ++++++++++++++ .../mcp_server/test_mcp_stale_session.py | 8 +- 4 files changed, 195 insertions(+), 41 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index d259fa32cd9..95129d9aeed 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2048,18 +2048,44 @@ if MCP_AVAILABLE: return texts[0][1] return "\n\n---\n\n".join(f"[{lbl}]\n{txt}" for lbl, txt in texts) + async def _raise_if_initialize_grants_no_mcp_servers( + allowed: Sequence[MCPServer], + user_api_key_auth: UserAPIKeyAuth | None, + mcp_servers: Sequence[str] | None, + client_ip: str | None, + ) -> None: + if allowed or user_api_key_auth is None or not user_api_key_auth.api_key: + return + if mcp_servers: + await raise_denied_scoped_mcp_access( + requested_names=mcp_servers, + user_api_key_auth=user_api_key_auth, + client_ip=client_ip, + ) + no_servers_denial: Final[_McpDeniedDetail] = { + "error": ( + "The key has no MCP servers granted, or none of its granted servers is loaded and allowed for " + "this client IP. Grant servers or access groups to the key, its team, or its organization " + "(object_permission.mcp_servers), check the server's allowed IPs, and reconnect." + ) + } + raise HTTPException(status_code=403, detail=no_servers_denial) + @contextlib.asynccontextmanager async def _gateway_initialize_instructions_request_scope( user_api_key_auth: UserAPIKeyAuth | None, mcp_servers: list[str] | None, client_ip: str | None, scoped_server_endpoint: bool = False, + is_initialize: bool = False, ) -> AsyncIterator[None]: allowed: Final = await _get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip, ) + if is_initialize: + await _raise_if_initialize_grants_no_mcp_servers(allowed, user_api_key_auth, mcp_servers, client_ip) if allowed: # return_exceptions=True: a per-server probe failure (incl. CancelledError # bubbled from anyio task group teardown on connection refused) must not @@ -4683,6 +4709,7 @@ if MCP_AVAILABLE: mcp_servers, _client_ip, scoped_server_endpoint=scoped_server_endpoint, + is_initialize=is_initialize, ): await target_manager.handle_request(scope, receive, local_send) if use_stateful and session_id and scope.get("method") == "DELETE": @@ -4819,6 +4846,7 @@ if MCP_AVAILABLE: mcp_servers, _sse_client_ip, scoped_server_endpoint=scoped_server_endpoint, + is_initialize=scope.get("method") == "GET", ): await sse_session_manager.handle_request(scope, receive, send) except MCPUpstreamAuthError as e: diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index f0fc3d892e2..e1099fe0a62 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -510,6 +510,37 @@ async def _call(session: ClientSession, tool_id: str, a: int = 3, b: int = 4) -> return await session.call_tool("call_tool", arguments={"tool_id": tool_id, "arguments": {"a": a, "b": b}}) +async def _raw_rpc( + proxy_server_url: str, key: str | None, method: str, params: dict[str, object], **headers: str +) -> httpx.Response: + async with httpx.AsyncClient() as client: + return await client.post( + f"{proxy_server_url}/mcp/proxy", + headers={ + "Accept": "application/json, text/event-stream", + **({"Authorization": f"Bearer {key}"} if key else {}), + **headers, + }, + json={"jsonrpc": "2.0", "id": 1, "method": method, "params": params}, + ) + + +async def _raw_initialize(proxy_server_url: str, key: str | None) -> httpx.Response: + return await _raw_rpc( + proxy_server_url, + key, + "initialize", + {"protocolVersion": "2025-03-26", "capabilities": {}, "clientInfo": {"name": "auth-test", "version": "1"}}, + ) + + +def _rpc_result(response: httpx.Response) -> dict[str, typing.Any]: + if response.headers["content-type"].startswith("text/event-stream"): + data_line = next(line for line in response.text.splitlines() if line.startswith("data:")) + return json.loads(data_line.removeprefix("data:"))["result"] + return response.json()["result"] + + def _assert_unauthorized(result: CallToolResult) -> None: assert result.isError is True assert result.content[0].text == "Unknown or unauthorized tool_id" @@ -527,13 +558,31 @@ class TestProxyMcpAuthorizationScope: _assert_unauthorized(await _call(ungranted, restricted_id)) @pytest.mark.asyncio - async def test_no_mcp_servers_sentinel_hides_every_tool(self, proxy_server_url: str) -> None: + async def test_no_mcp_servers_sentinel_rejects_initialize_and_hides_every_tool(self, proxy_server_url: str) -> None: async with _scoped_session(proxy_server_url) as granted: tool_id = (await _search(granted, "add"))["math_stdio-add"] - async with _scoped_session(proxy_server_url, "sk-none") as session: - assert await _search(session, "add") == {} - _assert_unauthorized(await session.call_tool("get_tool_schema", {"tool_id": tool_id})) - _assert_unauthorized(await _call(session, tool_id)) + response = await _raw_initialize(proxy_server_url, "sk-none") + assert response.status_code == 403, response.text + assert "no MCP servers granted" in response.json()["detail"]["error"] + + async def raw_call(name: str, arguments: dict[str, object]) -> dict[str, typing.Any]: + call = await _raw_rpc(proxy_server_url, "sk-none", "tools/call", {"name": name, "arguments": arguments}) + assert call.status_code == 200, call.text + return _rpc_result(call) + + listed = await _raw_rpc(proxy_server_url, "sk-none", "tools/list", {}) + assert listed.status_code == 200, listed.text + assert {tool["name"] for tool in _rpc_result(listed)["tools"]} == {"search_tools", "get_tool_schema", "call_tool"} + search = await raw_call("search_tools", {"query": "add"}) + assert search["isError"] is False, search + assert json.loads(search["content"][0]["text"]) == [] + for name, arguments in ( + ("get_tool_schema", {"tool_id": tool_id}), + ("call_tool", {"tool_id": tool_id, "arguments": {"a": 3, "b": 4}}), + ): + denied = await raw_call(name, arguments) + assert denied["isError"] is True, denied + assert denied["content"][0]["text"] == "Unknown or unauthorized tool_id" @pytest.mark.asyncio async def test_tool_grant_hides_ungranted_tools_and_blocks_their_ids(self, proxy_server_url: str) -> None: @@ -581,24 +630,7 @@ class TestProxyMcpAuthorizationScope: @pytest.mark.asyncio @pytest.mark.parametrize("key", [None, "sk-invalid"]) async def test_missing_or_invalid_key_cannot_initialize(self, proxy_server_url: str, key: str | None) -> None: - async with httpx.AsyncClient() as client: - response = await client.post( - f"{proxy_server_url}/mcp/proxy", - headers={ - "Accept": "application/json, text/event-stream", - **({"Authorization": f"Bearer {key}"} if key else {}), - }, - json={ - "jsonrpc": "2.0", - "id": 1, - "method": "initialize", - "params": { - "protocolVersion": "2025-03-26", - "capabilities": {}, - "clientInfo": {"name": "auth-test", "version": "1"}, - }, - }, - ) + response = await _raw_initialize(proxy_server_url, key) assert response.status_code == 401, response.text @pytest.mark.asyncio @@ -646,25 +678,28 @@ class TestProxyMcpAuthorizationScope: @pytest.mark.asyncio async def test_proxy_scope_exception_returns_iserror_and_emits_failure_log(self, proxy_server_url: str) -> None: - async with _scoped_session( + response = await _raw_rpc( proxy_server_url, "sk-none", + "tools/call", + {"name": "call_tool", "arguments": {"tool_id": "denied-scope", "arguments": {}}}, **{"x-mcp-servers": "math_restricted", "x-litellm-call-id": "proxy-scope-denial"}, - ) as session: - result = await session.call_tool("call_tool", {"tool_id": "denied-scope", "arguments": {}}) - assert result.isError is True - assert result.content[0].text == ( - "Error: The key is not allowed to access the requested MCP servers: math_restricted" - ) - async with asyncio.timeout(10): - while True: - payload = json.loads(await asyncio.to_thread(proxy_call_recorder.failures.get, True, 5)) - if payload["id"] == "proxy-scope-denial": - break - assert payload["call_type"] == "call_mcp_tool" - assert payload["status"] == "failure" - assert payload["response_cost"] == 0 - assert "math_restricted" in payload["error_str"] + ) + assert response.status_code == 200, response.text + result = _rpc_result(response) + assert result["isError"] is True + assert result["content"][0]["text"] == ( + "Error: The key is not allowed to access the requested MCP servers: math_restricted" + ) + async with asyncio.timeout(10): + while True: + payload = json.loads(await asyncio.to_thread(proxy_call_recorder.failures.get, True, 5)) + if payload["id"] == "proxy-scope-denial": + break + assert payload["call_type"] == "call_mcp_tool" + assert payload["status"] == "failure" + assert payload["response_cost"] == 0 + assert "math_restricted" in payload["error_str"] @pytest.mark.parametrize("arguments", ["wrong", False, None, [], 0]) def test_handler_rejects_non_object_arguments( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 12805e355f9..7eb4396e60d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -2003,6 +2003,11 @@ async def test_mcp_routing_chunked_initialize_to_stateful(): patch( "litellm.proxy._experimental.mcp_server.server.set_auth_context", ), + patch( # test-quality-ok: registry is empty in unit tests; key owns one server + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[MagicMock()], + ), patch( "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True, @@ -2488,6 +2493,11 @@ async def test_initialize_request_tracks_active_session_after_response_header(): new_callable=AsyncMock, return_value=(owner_auth, None, None, None, None, None), ), + patch( # test-quality-ok: registry is empty in unit tests; key owns one server + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[MagicMock()], + ), patch( "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True, @@ -2600,6 +2610,11 @@ async def test_initialize_request_with_existing_session_tracks_new_session(): {"x-new-header": "new"}, ), ), + patch( # test-quality-ok: registry is empty in unit tests; key owns one server + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[MagicMock()], + ), patch( "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True, @@ -5614,6 +5629,78 @@ class TestGatewayCreateInitializationOptions: assert server.create_initialization_options().server_name == "litellm-mcp-server" + @pytest.mark.asyncio + async def test_initialize_with_no_granted_servers_returns_403(self): + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server.server import ( + _gateway_initialize_instructions_request_scope, + ) + from litellm.proxy._types import UserAPIKeyAuth + + with patch( # test-quality-ok: grant resolution is the input under test + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[], + ): + with pytest.raises(HTTPException) as exc_info: + async with _gateway_initialize_instructions_request_scope( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-no-mcp"), + mcp_servers=None, + client_ip=None, + is_initialize=True, + ): + pytest.fail("initialize must not proceed when the key grants no MCP servers") + + assert exc_info.value.status_code == 403 + assert "no MCP servers granted" in exc_info.value.detail["error"] + + @pytest.mark.asyncio + async def test_initialize_with_no_granted_scoped_servers_returns_scoped_denial(self): + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server.server import ( + _gateway_initialize_instructions_request_scope, + ) + from litellm.proxy._types import UserAPIKeyAuth + + with patch( # test-quality-ok: grant resolution is the input under test + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[], + ): + with pytest.raises(HTTPException) as exc_info: + async with _gateway_initialize_instructions_request_scope( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-no-mcp"), + mcp_servers=["grafana"], + client_ip=None, + is_initialize=True, + ): + pytest.fail("scoped initialize must not proceed when nothing resolves") + + assert exc_info.value.status_code == 403 + assert "grafana" in exc_info.value.detail["error"] + + @pytest.mark.asyncio + async def test_non_initialize_request_with_no_granted_servers_is_not_rejected_here(self): + from litellm.proxy._experimental.mcp_server.server import ( + _gateway_initialize_instructions_request_scope, + _mcp_gateway_initialize_instructions, + ) + from litellm.proxy._types import UserAPIKeyAuth + + with patch( # test-quality-ok: grant resolution is the input under test + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[], + ): + async with _gateway_initialize_instructions_request_scope( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-no-mcp"), + mcp_servers=None, + client_ip=None, + ): + assert _mcp_gateway_initialize_instructions.get() is None + @pytest.mark.asyncio async def test_sse_handler_scopes_server_name_from_single_server_path(self): try: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index feab179570b..9420eecd222 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -750,8 +750,7 @@ async def test_admitted_subject_missing_stored_token_challenged_with_resource_me challenge = exc_info.value.headers["www-authenticate"] assert "authorization_uri=" not in challenge assert challenge == ( - 'Bearer resource_metadata="http://localhost:8000' - '/.well-known/oauth-protected-resource/mcp/repro_oauth_server"' + 'Bearer resource_metadata="http://localhost:8000/.well-known/oauth-protected-resource/mcp/repro_oauth_server"' ) @@ -938,6 +937,11 @@ async def test_handle_streamable_http_mcp_delegated_server_surfaces_upstream_cha "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", return_value=delegated_server, ), + patch( # test-quality-ok: registry is empty in unit tests; key owns the delegated server + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[delegated_server], + ), patch.object( session_manager_stateful, "handle_request", From 03815cf9f61fa1780cda35927b5fa9450439b1b2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 14:25:42 -0700 Subject: [PATCH 57/60] feat(claude-code): accept https zip archive plugin sources for skills (#40496) Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_lazy_openapi_snapshot.json | 12 +-- .../claude_code_marketplace.py | 39 ++++++-- litellm/types/proxy/claude_code_endpoints.py | 10 ++- .../test_claude_code_marketplace.py | 77 +++++++++++++++- .../_components/add_plugin_form.test.tsx | 89 ++++++++++++++++++- .../skills/_components/add_plugin_form.tsx | 71 +++++++++++---- .../components/AIHub/SkillHubTableColumns.tsx | 2 +- .../claude_code_plugins/helpers.test.ts | 60 +++++++++++++ .../components/claude_code_plugins/helpers.ts | 30 +++++-- .../claude_code_plugins/skill_detail.tsx | 2 +- .../components/claude_code_plugins/types.ts | 3 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 18 ++-- 12 files changed, 356 insertions(+), 57 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 4d5cdcc8003..3f060f1f5f6 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -5384,7 +5384,7 @@ "additionalProperties": { "type": "string" }, - "description": "Git source reference", + "description": "Plugin source reference", "title": "Source", "type": "object" }, @@ -5411,7 +5411,7 @@ "type": "object" }, "RegisterPluginRequest": { - "description": "Request body for registering a plugin in the marketplace.\n\nLiteLLM acts as a registry/discovery layer. Plugins are hosted on\nGitHub/GitLab/Bitbucket and referenced by their git source.", + "description": "Request body for registering a plugin in the marketplace.\n\nLiteLLM acts as a registry/discovery layer. Plugins are hosted on\nGitHub/GitLab/Bitbucket or as a zip archive on any https host and referenced by their source.", "properties": { "author": { "anyOf": [ @@ -5509,7 +5509,7 @@ "additionalProperties": { "type": "string" }, - "description": "Git source reference. Supported formats:\n- GitHub: {'source': 'github', 'repo': 'org/repo'}\n- Git URL: {'source': 'url', 'url': 'https://github.com/org/repo.git'}\n- Git Subdir: {'source': 'git-subdir', 'url': 'https://github.com/org/repo.git', 'path': 'plugins/plugin-name'}", + "description": "Plugin source reference. Supported formats:\n- GitHub: {'source': 'github', 'repo': 'org/repo'}\n- Git URL: {'source': 'url', 'url': 'https://github.com/org/repo.git'}\n- Git Subdir: {'source': 'git-subdir', 'url': 'https://github.com/org/repo.git', 'path': 'plugins/plugin-name'}\n- Zip archive on any https host (e.g. S3): {'source': 'archive', 'url': 'https://bucket.s3.amazonaws.com/plugin.zip', 'sha256': ''}", "title": "Source", "type": "object" }, @@ -5653,7 +5653,7 @@ "additionalProperties": { "type": "string" }, - "description": "Git source reference. Supported formats:\n- GitHub: {'source': 'github', 'repo': 'org/repo'}\n- Git URL: {'source': 'url', 'url': 'https://github.com/org/repo.git'}\n- Git Subdir: {'source': 'git-subdir', 'url': 'https://github.com/org/repo.git', 'path': 'plugins/plugin-name'}", + "description": "Plugin source reference. Supported formats:\n- GitHub: {'source': 'github', 'repo': 'org/repo'}\n- Git URL: {'source': 'url', 'url': 'https://github.com/org/repo.git'}\n- Git Subdir: {'source': 'git-subdir', 'url': 'https://github.com/org/repo.git', 'path': 'plugins/plugin-name'}\n- Zip archive on any https host (e.g. S3): {'source': 'archive', 'url': 'https://bucket.s3.amazonaws.com/plugin.zip', 'sha256': ''}", "title": "Source", "type": "object" }, @@ -5788,7 +5788,7 @@ ] }, "post": { - "description": "Register a new plugin in the LiteLLM marketplace.\n\nLiteLLM acts as a registry/discovery layer. Plugins are hosted on\nGitHub/GitLab/Bitbucket. Claude Code will clone from the git source\nwhen users install.\n\nThis endpoint is create-only and never overwrites. If a plugin with\nthe same name already exists it returns 409 Conflict; use\nPUT /claude-code/plugins/{plugin_name} to update an existing plugin.\n\nRequires a proxy admin API key.\n\nParameters:\n - name: Plugin name (kebab-case)\n - source: Git source reference (github, url, or git-subdir format)\n - version: Semantic version (optional)\n - description: Plugin description (optional)\n - author: Author information (optional)\n - homepage: Plugin homepage URL (optional)\n - keywords: Search keywords (optional)\n - category: Plugin category (optional)\n\nReturns:\n Registration status (action is always \"created\") and plugin information.\n\nExample:\n ```bash\n curl -X POST http://localhost:4000/claude-code/plugins \\\n -H \"Authorization: Bearer sk-...\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"name\": \"my-plugin\",\n \"source\": {\"source\": \"github\", \"repo\": \"org/my-plugin\"},\n \"version\": \"1.0.0\",\n \"description\": \"My awesome plugin\"\n }'\n ```", + "description": "Register a new plugin in the LiteLLM marketplace.\n\nLiteLLM acts as a registry/discovery layer. Plugins are hosted on\nGitHub/GitLab/Bitbucket or as a zip archive on any https host (e.g. S3).\nClaude Code clones the git source or downloads the archive when users install.\n\nThis endpoint is create-only and never overwrites. If a plugin with\nthe same name already exists it returns 409 Conflict; use\nPUT /claude-code/plugins/{plugin_name} to update an existing plugin.\n\nRequires a proxy admin API key.\n\nParameters:\n - name: Plugin name (kebab-case)\n - source: Plugin source reference (github, url, git-subdir, or archive format)\n - version: Semantic version (optional)\n - description: Plugin description (optional)\n - author: Author information (optional)\n - homepage: Plugin homepage URL (optional)\n - keywords: Search keywords (optional)\n - category: Plugin category (optional)\n\nReturns:\n Registration status (action is always \"created\") and plugin information.\n\nExample:\n ```bash\n curl -X POST http://localhost:4000/claude-code/plugins \\\n -H \"Authorization: Bearer sk-...\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"name\": \"my-plugin\",\n \"source\": {\"source\": \"github\", \"repo\": \"org/my-plugin\"},\n \"version\": \"1.0.0\",\n \"description\": \"My awesome plugin\"\n }'\n ```", "operationId": "register_plugin_claude_code_plugins_post", "requestBody": { "content": { @@ -5923,7 +5923,7 @@ ] }, "put": { - "description": "Update an existing plugin in the LiteLLM marketplace.\n\nThe plugin is identified by its name in the path, which is the resource\nidentity and cannot be changed here. This is a full replace, not a merge:\nthe manifest is rebuilt from the request body, so any optional field left\nout is reset to its default (e.g. an omitted version is cleared, not kept).\nSend the full desired state.\n\nReturns 404 if no plugin with the given name exists; use\nPOST /claude-code/plugins to create a new plugin.\n\nRequires a proxy admin API key.\n\nParameters:\n - plugin_name: Name of the plugin to update (path parameter)\n - source: Git source reference (github, url, or git-subdir format)\n - version: Semantic version (optional)\n - description: Plugin description (optional)\n - author: Author information (optional)\n - homepage: Plugin homepage URL (optional)\n - keywords: Search keywords (optional)\n - category: Plugin category (optional)\n\nReturns:\n Update status (action is always \"updated\") and plugin information.\n\nExample:\n ```bash\n curl -X PUT http://localhost:4000/claude-code/plugins/my-plugin \\\n -H \"Authorization: Bearer sk-...\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"source\": {\"source\": \"github\", \"repo\": \"org/my-plugin\"},\n \"version\": \"2.0.0\",\n \"description\": \"My awesome plugin\"\n }'\n ```", + "description": "Update an existing plugin in the LiteLLM marketplace.\n\nThe plugin is identified by its name in the path, which is the resource\nidentity and cannot be changed here. This is a full replace, not a merge:\nthe manifest is rebuilt from the request body, so any optional field left\nout is reset to its default (e.g. an omitted version is cleared, not kept).\nSend the full desired state.\n\nReturns 404 if no plugin with the given name exists; use\nPOST /claude-code/plugins to create a new plugin.\n\nRequires a proxy admin API key.\n\nParameters:\n - plugin_name: Name of the plugin to update (path parameter)\n - source: Plugin source reference (github, url, git-subdir, or archive format)\n - version: Semantic version (optional)\n - description: Plugin description (optional)\n - author: Author information (optional)\n - homepage: Plugin homepage URL (optional)\n - keywords: Search keywords (optional)\n - category: Plugin category (optional)\n\nReturns:\n Update status (action is always \"updated\") and plugin information.\n\nExample:\n ```bash\n curl -X PUT http://localhost:4000/claude-code/plugins/my-plugin \\\n -H \"Authorization: Bearer sk-...\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"source\": {\"source\": \"github\", \"repo\": \"org/my-plugin\"},\n \"version\": \"2.0.0\",\n \"description\": \"My awesome plugin\"\n }'\n ```", "operationId": "update_plugin_claude_code_plugins__plugin_name__put", "parameters": [ { diff --git a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py index 65bc46edfaf..cdeb1c6a111 100644 --- a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py +++ b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py @@ -2,8 +2,9 @@ CLAUDE CODE MARKETPLACE Provides a registry/discovery layer for Claude Code plugins. -Plugins are stored as metadata + git source references in LiteLLM database. -Actual plugin files are hosted on GitHub/GitLab/Bitbucket. +Plugins are stored as metadata + source references in LiteLLM database. +Actual plugin files are hosted on GitHub/GitLab/Bitbucket or as a zip archive on +any HTTPS host (S3, Artifactory, a static file server). Endpoints: /claude-code/marketplace.json - GET - List plugins for Claude Code discovery (unauthenticated) @@ -21,6 +22,7 @@ import re from collections.abc import Mapping, Sequence from datetime import datetime, timezone from typing import Annotated, Final, Protocol, TypedDict +from urllib.parse import urlsplit from fastapi import APIRouter, Depends, HTTPException from fastapi.responses import JSONResponse @@ -162,6 +164,15 @@ async def get_marketplace(): # alphanumeric characters, dots, hyphens, and underscores. # This implicitly blocks '..', leading '/', backslashes, and percent-encoded sequences. _VALID_GIT_SUBDIR_PATH_RE: Final = re.compile(r"^[a-zA-Z0-9][a-zA-Z0-9._-]*(/[a-zA-Z0-9][a-zA-Z0-9._-]*)*$") +_VALID_SHA256_RE: Final = re.compile(r"^[0-9a-fA-F]{64}$") + + +def _is_https_url_with_host(url: str) -> bool: + try: + parts: Final = urlsplit(url) + except ValueError: + return False + return parts.scheme == "https" and bool(parts.hostname) def _validate_plugin_source(source: Mapping[str, str]) -> None: @@ -199,10 +210,24 @@ def _validate_plugin_source(source: Mapping[str, str]) -> None: "error": "git-subdir 'path' must be a relative path of the form 'segment/segment' (alphanumeric, dots, hyphens, underscores only)" }, ) + elif source_type == "archive": + if not _is_https_url_with_host(source.get("url", "")): + raise HTTPException( + status_code=400, + detail={ + "error": "archive source must include an https 'url' field " + "(e.g., 'https://bucket.s3.amazonaws.com/plugins/plugin-name.zip')" + }, + ) + if "sha256" in source and not _VALID_SHA256_RE.match(source["sha256"]): + raise HTTPException( + status_code=400, + detail={"error": "archive 'sha256' must be a 64-character hex digest"}, + ) else: raise HTTPException( status_code=400, - detail={"error": "source.source must be 'github', 'url', or 'git-subdir'"}, + detail={"error": "source.source must be 'github', 'url', 'git-subdir', or 'archive'"}, ) @@ -248,8 +273,8 @@ async def register_plugin( Register a new plugin in the LiteLLM marketplace. LiteLLM acts as a registry/discovery layer. Plugins are hosted on - GitHub/GitLab/Bitbucket. Claude Code will clone from the git source - when users install. + GitHub/GitLab/Bitbucket or as a zip archive on any https host (e.g. S3). + Claude Code clones the git source or downloads the archive when users install. This endpoint is create-only and never overwrites. If a plugin with the same name already exists it returns 409 Conflict; use @@ -259,7 +284,7 @@ async def register_plugin( Parameters: - name: Plugin name (kebab-case) - - source: Git source reference (github, url, or git-subdir format) + - source: Plugin source reference (github, url, git-subdir, or archive format) - version: Semantic version (optional) - description: Plugin description (optional) - author: Author information (optional) @@ -503,7 +528,7 @@ async def update_plugin( Parameters: - plugin_name: Name of the plugin to update (path parameter) - - source: Git source reference (github, url, or git-subdir format) + - source: Plugin source reference (github, url, git-subdir, or archive format) - version: Semantic version (optional) - description: Plugin description (optional) - author: Author information (optional) diff --git a/litellm/types/proxy/claude_code_endpoints.py b/litellm/types/proxy/claude_code_endpoints.py index 2ee1bbbbb98..dcb5561cfeb 100644 --- a/litellm/types/proxy/claude_code_endpoints.py +++ b/litellm/types/proxy/claude_code_endpoints.py @@ -25,10 +25,12 @@ class PluginSpec(BaseModel): source: dict[str, str] = Field( ..., description=( - "Git source reference. Supported formats:\n" + "Plugin source reference. Supported formats:\n" "- GitHub: {'source': 'github', 'repo': 'org/repo'}\n" "- Git URL: {'source': 'url', 'url': 'https://github.com/org/repo.git'}\n" - "- Git Subdir: {'source': 'git-subdir', 'url': 'https://github.com/org/repo.git', 'path': 'plugins/plugin-name'}" + "- Git Subdir: {'source': 'git-subdir', 'url': 'https://github.com/org/repo.git', 'path': 'plugins/plugin-name'}\n" + "- Zip archive on any https host (e.g. S3): " + "{'source': 'archive', 'url': 'https://bucket.s3.amazonaws.com/plugin.zip', 'sha256': ''}" ), ) version: str | None = Field("1.0.0", description="Semantic version") @@ -46,7 +48,7 @@ class RegisterPluginRequest(PluginSpec): Request body for registering a plugin in the marketplace. LiteLLM acts as a registry/discovery layer. Plugins are hosted on - GitHub/GitLab/Bitbucket and referenced by their git source. + GitHub/GitLab/Bitbucket or as a zip archive on any https host and referenced by their source. """ name: str = Field( @@ -76,7 +78,7 @@ class PluginResponse(BaseModel): name: str = Field(..., description="Plugin name") version: str | None = Field(None, description="Plugin version") description: str | None = Field(None, description="Plugin description") - source: dict[str, str] = Field(..., description="Git source reference") + source: dict[str, str] = Field(..., description="Plugin source reference") enabled: bool = Field(..., description="Whether plugin is enabled") diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py index 99b09f48fe5..aa2da262bae 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py @@ -1,7 +1,7 @@ """ Unit tests for claude_code_marketplace.py source validation. -Covers the git-subdir source type added alongside the existing github and url types. +Covers the git-subdir and archive source types added alongside the existing github and url types. """ import json @@ -87,6 +87,12 @@ _GIT_SUBDIR_SOURCE = { "path": "plugins/my-plugin", } +_ARCHIVE_SOURCE = { + "source": "archive", + "url": "https://skills-bucket.s3.us-east-1.amazonaws.com/plugins/s3-skill-1.0.0.zip", + "sha256": "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", +} + @pytest.fixture(autouse=True) def _patch_proxy_globals(monkeypatch): @@ -377,6 +383,75 @@ async def test_register_plugin_unknown_source_type(): assert exc_info.value.status_code == 400 assert "git-subdir" in exc_info.value.detail["error"] + assert "archive" in exc_info.value.detail["error"] + + +@pytest.mark.asyncio +async def test_archive_source_registers_and_is_served_verbatim_in_marketplace(): + response = await register_plugin( + request=RegisterPluginRequest(name="s3-skill", source=_ARCHIVE_SOURCE), + user_api_key_dict=_USER, + ) + + assert response.action == "created" + assert response.plugin.source == _ARCHIVE_SOURCE + + marketplace = json.loads((await get_marketplace()).body) + assert marketplace["plugins"] == [{"name": "s3-skill", "source": _ARCHIVE_SOURCE, "version": "1.0.0"}] + + +@pytest.mark.asyncio +async def test_archive_source_without_sha256_is_accepted(): + source = {"source": "archive", "url": "https://artifacts.example.com/plugin.zip"} + + response = await register_plugin( + request=RegisterPluginRequest(name="unpinned-skill", source=source), + user_api_key_dict=_USER, + ) + + assert response.plugin.source == source + + +@pytest.mark.asyncio +async def test_update_plugin_to_archive_source(): + name = "my-monorepo-plugin" + await register_plugin( + request=RegisterPluginRequest(name=name, source=_GIT_SUBDIR_SOURCE), + user_api_key_dict=_USER, + ) + + response = await update_plugin( + plugin_name=name, + request=UpdatePluginRequest(source=_ARCHIVE_SOURCE), + user_api_key_dict=_USER, + ) + + assert response.action == "updated" + assert (await _read_stored_manifest(name))["source"] == _ARCHIVE_SOURCE + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "source, expected_fragment", + [ + ({"source": "archive"}, "url"), + ({"source": "archive", "url": ""}, "url"), + ({"source": "archive", "url": "http://artifacts.example.com/plugin.zip"}, "https"), + ({"source": "archive", "url": "s3://skills-bucket/plugin.zip"}, "https"), + ({"source": "archive", "url": "https://"}, "https"), + ({"source": "archive", "url": "https:///plugin.zip"}, "https"), + ({"source": "archive", "url": "https://[::1/plugin.zip"}, "https"), + ({"source": "archive", "url": "https://artifacts.example.com/plugin.zip", "sha256": "a" * 63}, "sha256"), + ({"source": "archive", "url": "https://artifacts.example.com/plugin.zip", "sha256": "a" * 65}, "sha256"), + ({"source": "archive", "url": "https://artifacts.example.com/plugin.zip", "sha256": "g" * 64}, "sha256"), + ], +) +async def test_register_plugin_archive_rejects_malformed_source(source, expected_fragment): + with pytest.raises(HTTPException) as exc_info: + await register_plugin(request=RegisterPluginRequest(name="bad-plugin", source=source), user_api_key_dict=_USER) + + assert exc_info.value.status_code == 400 + assert expected_fragment in exc_info.value.detail["error"] @pytest.mark.asyncio diff --git a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/add_plugin_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/add_plugin_form.test.tsx index 1058d4d0466..6507855d0c7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/add_plugin_form.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/add_plugin_form.test.tsx @@ -20,21 +20,25 @@ const DEFAULT_PROPS = { onSuccess: vi.fn(), }; -const URL_PLACEHOLDER = "https://github.com/org/repo or https://gitlab.com/org/repo"; +const URL_PLACEHOLDER = "https://github.com/org/repo or https://bucket.s3.amazonaws.com/my-skill.zip"; const SUBPATH_PLACEHOLDER = "plugins/my-skill"; +const SHA256_PLACEHOLDER = "64 hex characters"; +const S3_ZIP_URL = "https://skills-bucket.s3.us-east-1.amazonaws.com/plugins/s3-skill-1.0.0.zip"; +const DIGEST = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"; describe("AddPluginForm", () => { beforeEach(() => { vi.clearAllMocks(); }); - it("renders the host-agnostic repository URL input and subfolder field", () => { + it("renders the host-agnostic source URL input and subfolder field, hiding the digest until a zip is entered", () => { renderWithProviders(); - expect(screen.getByText("Repository URL")).toBeInTheDocument(); + expect(screen.getByText("Source URL")).toBeInTheDocument(); expect(screen.getByPlaceholderText(URL_PLACEHOLDER)).toBeInTheDocument(); expect(screen.getByText("Subfolder path (Optional)")).toBeInTheDocument(); expect(screen.getByPlaceholderText(SUBPATH_PLACEHOLDER)).toBeInTheDocument(); + expect(screen.queryByPlaceholderText(SHA256_PLACEHOLDER)).not.toBeInTheDocument(); }); it("shows GitHub repo preview for a plain repo URL", async () => { @@ -255,6 +259,85 @@ describe("AddPluginForm", () => { }); }); + it("shows a zip archive preview, disables the subfolder field, and reveals the digest field", async () => { + renderWithProviders(); + + await typeUrl(S3_ZIP_URL); + + expect(await screen.findByText(/Zip archive/)).toBeInTheDocument(); + expect(screen.getByPlaceholderText(SUBPATH_PLACEHOLDER)).toBeDisabled(); + expect(screen.getByText("A zip archive is installed as a whole, so this field is disabled")).toBeInTheDocument(); + expect(screen.getByPlaceholderText(SHA256_PLACEHOLDER)).toBeInTheDocument(); + expect((screen.getByPlaceholderText("my-skill") as HTMLInputElement).value).toBe("s3-skill-1-0-0"); + }); + + it("submits an archive source without a digest when the field is left empty", async () => { + renderWithProviders(); + + await typeUrl(S3_ZIP_URL); + await submit(); + + await waitFor(() => { + expect(mockRegister).toHaveBeenCalledWith( + "sk-test", + expect.objectContaining({ source: { source: "archive", url: S3_ZIP_URL } }), + ); + }); + }); + + it("submits an archive source pinned to the lowercased digest", async () => { + renderWithProviders(); + + await typeUrl(S3_ZIP_URL); + await act(async () => { + fireEvent.change(screen.getByPlaceholderText(SHA256_PLACEHOLDER), { + target: { value: ` ${DIGEST.toUpperCase()} ` }, + }); + }); + await submit(); + + await waitFor(() => { + expect(mockRegister).toHaveBeenCalledWith( + "sk-test", + expect.objectContaining({ source: { source: "archive", url: S3_ZIP_URL, sha256: DIGEST } }), + ); + }); + }); + + it("drops the digest once the archive URL changes so a stale checksum is never sent for a new file", async () => { + renderWithProviders(); + + await typeUrl(S3_ZIP_URL); + await act(async () => { + fireEvent.change(screen.getByPlaceholderText(SHA256_PLACEHOLDER), { target: { value: DIGEST } }); + }); + const otherZipUrl = S3_ZIP_URL.replace("1.0.0", "1.1.0"); + await typeUrl(otherZipUrl); + + expect(screen.getByPlaceholderText(SHA256_PLACEHOLDER)).toHaveValue(""); + await submit(); + + await waitFor(() => { + expect(mockRegister).toHaveBeenCalledWith( + "sk-test", + expect.objectContaining({ source: { source: "archive", url: otherZipUrl } }), + ); + }); + }); + + it("blocks submission and shows the digest error for a malformed sha256", async () => { + renderWithProviders(); + + await typeUrl(S3_ZIP_URL); + await act(async () => { + fireEvent.change(screen.getByPlaceholderText(SHA256_PLACEHOLDER), { target: { value: "not-a-digest" } }); + }); + await submit(); + + expect(await screen.findByText("SHA-256 must be a 64-character hex digest")).toBeInTheDocument(); + expect(mockRegister).not.toHaveBeenCalled(); + }); + it("surfaces the backend error message when registration fails", async () => { mockRegister.mockRejectedValueOnce(new Error("Plugin 'claude-code' already exists")); renderWithProviders(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/add_plugin_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/add_plugin_form.tsx index 70669ffc0cc..e420038b124 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/add_plugin_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/add_plugin_form.tsx @@ -26,6 +26,7 @@ import { parseKeywords, parseSkillSource, isValidSubPath, + isValidSha256, SkillSourcePreview, } from "@/components/claude_code_plugins/helpers"; import { PluginAuthor, PluginSource, SkillRegisterRequest } from "@/components/claude_code_plugins/types"; @@ -39,13 +40,14 @@ interface AddPluginFormProps { } const addPluginShape = { - skillUrl: z.string().min(1, "Please enter a repository URL"), + skillUrl: z.string().min(1, "Please enter a repository or zip archive URL"), subPath: z .string() .refine( (value) => !value || isValidSubPath(value), "Subfolder must be a relative path like plugins/my-skill (letters, numbers, dots, hyphens, underscores)", ), + sha256: z.string().refine(isValidSha256, "SHA-256 must be a 64-character hex digest"), name: z .string() .min(1, "Please enter skill name") @@ -69,6 +71,7 @@ type AddPluginFormValues = z.infer; const EMPTY_VALUES: AddPluginFormValues = { skillUrl: "", subPath: "", + sha256: "", name: "", domain: "", namespace: "", @@ -89,11 +92,19 @@ const buildAuthor = (values: AddPluginFormValues): PluginAuthor | undefined => { return email ? { name, email } : { name }; }; +const archiveUrlOf = (preview: SkillSourcePreview | null): string | undefined => + preview?.parsed.source === "archive" ? preview.parsed.url : undefined; + +const withArchiveDigest = (source: PluginSource, sha256: string): PluginSource => { + const digest = sha256.trim(); + return source.source === "archive" && digest ? { ...source, sha256: digest.toLowerCase() } : source; +}; + const buildRegisterRequest = (values: AddPluginFormValues, source: PluginSource): SkillRegisterRequest => { const author = buildAuthor(values); return { name: values.name.trim(), - source, + source: withArchiveDigest(source, values.sha256), ...(values.version ? { version: values.version.trim() } : {}), ...(values.description ? { description: values.description.trim() } : {}), ...(author ? { author } : {}), @@ -115,6 +126,16 @@ const PREDEFINED_CATEGORIES = [ "Documentation", ]; +const SUB_PATH_LOCK_REASON = { + "git-subdir": "The URL already points to a subfolder, so this field is disabled", + archive: "A zip archive is installed as a whole, so this field is disabled", +} as const; + +type SubPathLock = keyof typeof SUB_PATH_LOCK_REASON; + +const subPathLockFor = (source: PluginSource["source"] | undefined): SubPathLock | null => + source === "git-subdir" || source === "archive" ? source : null; + const labelWithHint = (label: string, hint: string): React.ReactNode => ( <> {label} @@ -129,15 +150,18 @@ const AddPluginForm: React.FC = ({ visible, onClose, accessT const form = useZodForm(addPluginSchema, { defaultValues: EMPTY_VALUES }); const [isSubmitting, setIsSubmitting] = useState(false); const [urlPreview, setUrlPreview] = useState(null); - const [urlEncodesSubdir, setUrlEncodesSubdir] = useState(false); + const [subPathLock, setSubPathLock] = useState(null); const recomputePreview = (skillUrl: string, subPath: string) => { - const encodesSubdir = parseSkillSource(skillUrl)?.parsed.source === "git-subdir"; - setUrlEncodesSubdir(encodesSubdir); - if (encodesSubdir && form.getValues("subPath")) { + const lock = subPathLockFor(parseSkillSource(skillUrl)?.parsed.source); + setSubPathLock(lock); + if (lock && form.getValues("subPath")) { form.setValue("subPath", ""); } - const preview = parseSkillSource(skillUrl, encodesSubdir ? undefined : subPath); + const preview = parseSkillSource(skillUrl, lock ? undefined : subPath); + if (archiveUrlOf(preview) !== archiveUrlOf(urlPreview) && form.getValues("sha256")) { + form.setValue("sha256", ""); + } setUrlPreview(preview); if (preview && !form.getValues("name")) { form.setValue("name", preview.suggestedName); @@ -151,7 +175,7 @@ const AddPluginForm: React.FC = ({ visible, onClose, accessT } if (!urlPreview) { - toast.error("Please enter a valid repository URL"); + toast.error("Please enter a valid repository or zip archive URL"); return; } @@ -176,7 +200,7 @@ const AddPluginForm: React.FC = ({ visible, onClose, accessT toast.success("Skill registered successfully"); form.reset(EMPTY_VALUES); setUrlPreview(null); - setUrlEncodesSubdir(false); + setSubPathLock(null); onSuccess(); onClose(); } catch (error) { @@ -190,7 +214,7 @@ const AddPluginForm: React.FC = ({ visible, onClose, accessT const handleCancel = () => { form.reset(EMPTY_VALUES); setUrlPreview(null); - setUrlEncodesSubdir(false); + setSubPathLock(null); onClose(); }; @@ -207,15 +231,15 @@ const AddPluginForm: React.FC = ({ visible, onClose, accessT control={form.control} name="skillUrl" label={labelWithHint( - "Repository URL", - "Paste an HTTPS git repository URL from GitHub, GitLab, Bitbucket, or a self-hosted host. E.g. github.com/org/repo, gitlab.com/org/repo, or github.com/org/repo/tree/main/my-skill", + "Source URL", + "Paste an HTTPS git repository URL from GitHub, GitLab, Bitbucket, or a self-hosted host (e.g. github.com/org/repo or github.com/org/repo/tree/main/my-skill), or an HTTPS link to a .zip archive of the skill hosted on S3 or any static file server.", )} > {({ ref, onChange, ...field }) => ( { onChange(event); @@ -232,9 +256,7 @@ const AddPluginForm: React.FC = ({ visible, onClose, accessT "Subfolder path (Optional)", "Path within the repository where the skill lives (e.g., plugins/my-skill). Leave empty if the skill is at the repo root.", )} - description={ - urlEncodesSubdir ? "The URL already points to a subfolder, so this field is disabled" : undefined - } + description={subPathLock ? SUB_PATH_LOCK_REASON[subPathLock] : undefined} > {({ ref, onChange, ...field }) => ( = ({ visible, onClose, accessT onChange(event); recomputePreview(form.getValues("skillUrl"), event.target.value); }} - disabled={urlEncodesSubdir} + disabled={subPathLock !== null} /> )} + {urlPreview?.parsed.source === "archive" && ( + + {({ ref, ...field }) => ( + + )} + + )} + {urlPreview && (
Detected: {urlPreview.label} diff --git a/ui/litellm-dashboard/src/components/AIHub/SkillHubTableColumns.tsx b/ui/litellm-dashboard/src/components/AIHub/SkillHubTableColumns.tsx index 2a1530cc352..1104e2c7485 100644 --- a/ui/litellm-dashboard/src/components/AIHub/SkillHubTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/SkillHubTableColumns.tsx @@ -26,7 +26,7 @@ function getSkillSourceLink(skill: Plugin): { url: string; label: string } | nul const url = src.path ? `${src.url}/tree/main/${src.path}` : src.url; return { url, label: url.replace("https://github.com/", "") }; } - if (src?.source === "url" && src.url) { + if ((src?.source === "url" || src?.source === "archive") && src.url) { return { url: src.url, label: src.url.replace(/^https?:\/\//, "") }; } return null; diff --git a/ui/litellm-dashboard/src/components/claude_code_plugins/helpers.test.ts b/ui/litellm-dashboard/src/components/claude_code_plugins/helpers.test.ts index 1d07ca09df5..4713fcbc55b 100644 --- a/ui/litellm-dashboard/src/components/claude_code_plugins/helpers.test.ts +++ b/ui/litellm-dashboard/src/components/claude_code_plugins/helpers.test.ts @@ -17,6 +17,7 @@ import { formatKeywords, parseSkillSource, isValidSubPath, + isValidSha256, buildMarketplaceSettingsSnippet, } from "./helpers"; import { MarketplacePluginEntry } from "./types"; @@ -114,6 +115,12 @@ describe("getSourceDisplayText", () => { ); }); + it("shows the archive url for an archive source", () => { + expect(getSourceDisplayText({ source: "archive", url: "https://bucket.s3.amazonaws.com/skill.zip" })).toBe( + "https://bucket.s3.amazonaws.com/skill.zip", + ); + }); + it("returns unknown for missing data", () => { expect(getSourceDisplayText({ source: "github" })).toBe("Unknown source"); }); @@ -140,6 +147,12 @@ describe("getSourceLink", () => { ); }); + it("returns the archive url for an archive source", () => { + expect(getSourceLink({ source: "archive", url: "https://bucket.s3.amazonaws.com/skill.zip" })).toBe( + "https://bucket.s3.amazonaws.com/skill.zip", + ); + }); + it("returns null when no repo or url", () => { expect(getSourceLink({ source: "github" })).toBeNull(); }); @@ -329,6 +342,20 @@ describe("isValidUrl", () => { }); }); +describe("isValidSha256", () => { + it("accepts an empty digest and a 64-character hex digest in either case", () => { + expect(isValidSha256("")).toBe(true); + expect(isValidSha256("a".repeat(64))).toBe(true); + expect(isValidSha256(" " + "ABCDEF0123456789".repeat(4) + " ")).toBe(true); + }); + + it("rejects wrong length and non-hex digests", () => { + expect(isValidSha256("a".repeat(63))).toBe(false); + expect(isValidSha256("a".repeat(65))).toBe(false); + expect(isValidSha256("g".repeat(64))).toBe(false); + }); +}); + describe("parseKeywords", () => { it("splits comma-separated keywords", () => { expect(parseKeywords("a, b, c")).toEqual(["a", "b", "c"]); @@ -445,6 +472,39 @@ describe("parseSkillSource", () => { expect(parseSkillSource("not a url")).toBeNull(); }); + it("parses an S3 zip URL into an archive source and names the skill after the file", () => { + expect(parseSkillSource("https://skills-bucket.s3.us-east-1.amazonaws.com/plugins/My_Skill-1.0.0.zip")).toEqual({ + parsed: { source: "archive", url: "https://skills-bucket.s3.us-east-1.amazonaws.com/plugins/My_Skill-1.0.0.zip" }, + label: "Zip archive — skills-bucket.s3.us-east-1.amazonaws.com/plugins/My_Skill-1.0.0.zip", + suggestedName: "my-skill-1-0-0", + }); + }); + + it("keeps the query string of a zip URL so versioned or signed object links still resolve", () => { + expect(parseSkillSource("https://bucket.s3.amazonaws.com/skill.ZIP?versionId=abc")?.parsed).toEqual({ + source: "archive", + url: "https://bucket.s3.amazonaws.com/skill.ZIP?versionId=abc", + }); + }); + + it("ignores the subfolder for a zip URL since the archive is installed whole", () => { + expect(parseSkillSource("https://artifacts.example.com/skill.zip", "plugins/x")?.parsed).toEqual({ + source: "archive", + url: "https://artifacts.example.com/skill.zip", + }); + }); + + it("rejects a plain http zip URL", () => { + expect(parseSkillSource("http://artifacts.example.com/skill.zip")).toBeNull(); + }); + + it("treats a github zip download URL as an archive rather than a repo path", () => { + expect(parseSkillSource("https://github.com/org/repo/releases/download/v1/skill.zip")?.parsed).toEqual({ + source: "archive", + url: "https://github.com/org/repo/releases/download/v1/skill.zip", + }); + }); + it("suggests a kebab-friendly name from the last path segment", () => { expect(parseSkillSource("github.com/org/my-awesome-skill")?.suggestedName).toBe("my-awesome-skill"); expect(parseSkillSource("github.com/org/repo/tree/main/plugins/cool-skill")?.suggestedName).toBe("cool-skill"); diff --git a/ui/litellm-dashboard/src/components/claude_code_plugins/helpers.ts b/ui/litellm-dashboard/src/components/claude_code_plugins/helpers.ts index 5b3d4f20c1e..6f673d67f87 100644 --- a/ui/litellm-dashboard/src/components/claude_code_plugins/helpers.ts +++ b/ui/litellm-dashboard/src/components/claude_code_plugins/helpers.ts @@ -23,6 +23,12 @@ const GITHUB_HOST = "github.com"; const SKILL_FILE_EXTENSION_REGEX = /\.(md|markdown|txt|json|ya?ml|toml)$/i; +const ZIP_ARCHIVE_REGEX = /\.zip$/i; + +export const SHA256_REGEX = /^[0-9a-fA-F]{64}$/; + +export const isValidSha256 = (digest: string): boolean => digest.trim() === "" || SHA256_REGEX.test(digest.trim()); + // WHATWG normalizes obfuscated IPv4 (e.g. 2130706433, 0x7f.0.0.1) to dotted-decimal, so this // catches every IPv4 form; bracketed IPv6 is rejected separately. const IPV4_HOST_REGEX = /^\d{1,3}(\.\d{1,3}){3}$/; @@ -160,16 +166,26 @@ const parseRawGitSource = (url: URL, subPath?: string): SkillSourcePreview | nul }; }; +const parseArchiveSource = (url: URL): SkillSourcePreview => ({ + parsed: { source: "archive", url: url.href }, + label: `Zip archive — ${url.host}${url.pathname}`, + suggestedName: toKebabCase(lastSegment(url.pathname).replace(ZIP_ARCHIVE_REGEX, "")), +}); + /** - * Parse any git-accessible repository URL into a registerable skill source. - * GitHub URLs keep their `github`/`git-subdir` shorthand; every other host is - * treated as a raw repo URL, with an optional subfolder turning it into git-subdir. + * Parse any git-accessible repository URL or https zip archive URL into a registerable skill + * source. A `.zip` path is an `archive` source (S3, Artifactory, any static host). GitHub URLs + * keep their `github`/`git-subdir` shorthand; every other host is treated as a raw repo URL, + * with an optional subfolder turning it into git-subdir. */ export const parseSkillSource = (rawUrl: string, subPath?: string): SkillSourcePreview | null => { const url = parseRepoUrl(rawUrl); if (!url) { return null; } + if (ZIP_ARCHIVE_REGEX.test(url.pathname)) { + return parseArchiveSource(url); + } if (url.hostname.replace(/^www\./, "") === GITHUB_HOST) { return parseGitHubSource(url, subPath); } @@ -245,7 +261,7 @@ export const getSourceDisplayText = (source: PluginSource): string => { if (source.source === "git-subdir" && source.url && source.path) { return `${source.url} @ ${source.path}`; } - if (source.source === "url" && source.url) { + if ((source.source === "url" || source.source === "archive") && source.url) { return source.url; } return "Unknown source"; @@ -258,10 +274,8 @@ export const getSourceLink = (source: PluginSource): string | null => { if (source.source === "github" && source.repo) { return `https://github.com/${source.repo}`; } - if ((source.source === "url" || source.source === "git-subdir") && source.url) { - return source.url; - } - return null; + const linksToUrl = source.source === "url" || source.source === "git-subdir" || source.source === "archive"; + return linksToUrl && source.url ? source.url : null; }; /** diff --git a/ui/litellm-dashboard/src/components/claude_code_plugins/skill_detail.tsx b/ui/litellm-dashboard/src/components/claude_code_plugins/skill_detail.tsx index 7cbbb08b623..b37204c6073 100644 --- a/ui/litellm-dashboard/src/components/claude_code_plugins/skill_detail.tsx +++ b/ui/litellm-dashboard/src/components/claude_code_plugins/skill_detail.tsx @@ -26,7 +26,7 @@ const SkillDetail: React.FC = ({ skill, onBack }) => { const src = skill.source; if (src.source === "github" && src.repo) return `https://github.com/${src.repo}`; if (src.source === "git-subdir" && src.url) return src.path ? `${src.url}/tree/main/${src.path}` : src.url; - if (src.source === "url" && src.url) return src.url; + if ((src.source === "url" || src.source === "archive") && src.url) return src.url; return null; })(); diff --git a/ui/litellm-dashboard/src/components/claude_code_plugins/types.ts b/ui/litellm-dashboard/src/components/claude_code_plugins/types.ts index d16c880749b..98f80fd2959 100644 --- a/ui/litellm-dashboard/src/components/claude_code_plugins/types.ts +++ b/ui/litellm-dashboard/src/components/claude_code_plugins/types.ts @@ -8,10 +8,11 @@ import type { components } from "@/lib/http/schema"; // Kept hand-written: the backend types `source` as Dict[str, str], so the generated type is a // loose string map; this discriminant union is what the parser and display helpers rely on. export interface PluginSource { - source: "github" | "url" | "git-subdir"; + source: "github" | "url" | "git-subdir" | "archive"; repo?: string; // Format: "org/repo" for GitHub url?: string; // Full URL for other sources path?: string; // Subdirectory path for git-subdir + sha256?: string; } export type PluginAuthor = components["schemas"]["PluginAuthor"]; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index ab2c1f27a59..c7a760e1eee 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -2152,8 +2152,8 @@ export interface paths { * @description Register a new plugin in the LiteLLM marketplace. * * LiteLLM acts as a registry/discovery layer. Plugins are hosted on - * GitHub/GitLab/Bitbucket. Claude Code will clone from the git source - * when users install. + * GitHub/GitLab/Bitbucket or as a zip archive on any https host (e.g. S3). + * Claude Code clones the git source or downloads the archive when users install. * * This endpoint is create-only and never overwrites. If a plugin with * the same name already exists it returns 409 Conflict; use @@ -2163,7 +2163,7 @@ export interface paths { * * Parameters: * - name: Plugin name (kebab-case) - * - source: Git source reference (github, url, or git-subdir format) + * - source: Plugin source reference (github, url, git-subdir, or archive format) * - version: Semantic version (optional) * - description: Plugin description (optional) * - author: Author information (optional) @@ -2229,7 +2229,7 @@ export interface paths { * * Parameters: * - plugin_name: Name of the plugin to update (path parameter) - * - source: Git source reference (github, url, or git-subdir format) + * - source: Plugin source reference (github, url, git-subdir, or archive format) * - version: Semantic version (optional) * - description: Plugin description (optional) * - author: Author information (optional) @@ -33638,7 +33638,7 @@ export interface components { name: string; /** * Source - * @description Git source reference + * @description Plugin source reference */ source: { [key: string]: string; @@ -34887,7 +34887,7 @@ export interface components { * @description Request body for registering a plugin in the marketplace. * * LiteLLM acts as a registry/discovery layer. Plugins are hosted on - * GitHub/GitLab/Bitbucket and referenced by their git source. + * GitHub/GitLab/Bitbucket or as a zip archive on any https host and referenced by their source. */ RegisterPluginRequest: { /** @description Plugin author */ @@ -34929,10 +34929,11 @@ export interface components { namespace?: string | null; /** * Source - * @description Git source reference. Supported formats: + * @description Plugin source reference. Supported formats: * - GitHub: {'source': 'github', 'repo': 'org/repo'} * - Git URL: {'source': 'url', 'url': 'https://github.com/org/repo.git'} * - Git Subdir: {'source': 'git-subdir', 'url': 'https://github.com/org/repo.git', 'path': 'plugins/plugin-name'} + * - Zip archive on any https host (e.g. S3): {'source': 'archive', 'url': 'https://bucket.s3.amazonaws.com/plugin.zip', 'sha256': ''} */ source: { [key: string]: string; @@ -38211,10 +38212,11 @@ export interface components { namespace?: string | null; /** * Source - * @description Git source reference. Supported formats: + * @description Plugin source reference. Supported formats: * - GitHub: {'source': 'github', 'repo': 'org/repo'} * - Git URL: {'source': 'url', 'url': 'https://github.com/org/repo.git'} * - Git Subdir: {'source': 'git-subdir', 'url': 'https://github.com/org/repo.git', 'path': 'plugins/plugin-name'} + * - Zip archive on any https host (e.g. S3): {'source': 'archive', 'url': 'https://bucket.s3.amazonaws.com/plugin.zip', 'sha256': ''} */ source: { [key: string]: string; From 6cc13e07a68f1e822c1acbb7bec808c45baec8e2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 14:26:09 -0700 Subject: [PATCH 58/60] feat(proxy): granular key/team access control for Claude Code marketplace plugins (#40518) Adds object_permission.skills to keys and teams, enforces it on /claude-code/marketplace.json?key=, /claude-code/plugins and /claude-code/plugins/{name}, and exposes an Allowed Skills selector in the key and team create/edit forms of the Admin UI Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../migration.sql | 1 + .../litellm_proxy_extras/schema.prisma | 1 + litellm/models/object_permission.py | 1 + litellm/proxy/_lazy_openapi_snapshot.json | 30 ++++- litellm/proxy/_types.py | 4 + .../claude_code_marketplace.py | 43 +++++-- .../claude_code_skill_access.py | 67 +++++++++++ litellm/proxy/schema.prisma | 1 + .../object_permission_repository.py | 6 + litellm/types/object_permission.py | 3 +- schema.prisma | 1 + .../test_claude_code_marketplace.py | 2 +- .../test_claude_code_marketplace.py | 97 +++++++++++++++- .../test_claude_code_skill_access.py | 79 +++++++++++++ .../proxy/auth/test_route_checks.py | 14 +++ .../test_customer_endpoints.py | 1 + .../test_object_permission_utils.py | 23 ++++ .../src/components/Teams.test.tsx | 29 +++++ ui/litellm-dashboard/src/components/Teams.tsx | 46 ++++++++ .../components/object_permissions_view.tsx | 11 ++ .../organisms/createKeyPayload.test.ts | 12 ++ .../components/organisms/createKeyPayload.ts | 5 + .../create_key_button.integration.test.tsx | 24 ++++ .../organisms/create_key_button.tsx | 31 +++++ .../components/skills/SkillSelector.test.tsx | 54 +++++++++ .../src/components/skills/SkillSelector.tsx | 109 ++++++++++++++++++ .../src/components/team/TeamInfo.test.tsx | 43 +++++++ .../src/components/team/TeamInfo.tsx | 26 +++++ .../templates/KeyEditViewControls.tsx | 34 +++++- .../KeyInfoView.handleKeyUpdate.test.tsx | 39 +++++++ .../components/templates/keyEditFormValues.ts | 4 + .../templates/key_edit_view.test.tsx | 40 +++++++ .../components/templates/key_edit_view.tsx | 15 +-- .../components/templates/key_info_view.tsx | 8 ++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 22 +++- 35 files changed, 892 insertions(+), 34 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260910000000_skills_on_object_permission/migration.sql create mode 100644 litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_skill_access.py create mode 100644 tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_skill_access.py create mode 100644 ui/litellm-dashboard/src/components/skills/SkillSelector.test.tsx create mode 100644 ui/litellm-dashboard/src/components/skills/SkillSelector.tsx diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260910000000_skills_on_object_permission/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260910000000_skills_on_object_permission/migration.sql new file mode 100644 index 00000000000..c982bc38a69 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260910000000_skills_on_object_permission/migration.sql @@ -0,0 +1 @@ +ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN IF NOT EXISTS "skills" TEXT[] DEFAULT ARRAY[]::TEXT[]; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 05c5aad9303..817df082d8c 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -282,6 +282,7 @@ model LiteLLM_ObjectPermissionTable { mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user search_tools String[] @default([]) // search_tool_name values this key/team/user may call mcp_tool_search_enabled Boolean? + skills String[] @default([]) // Claude Code plugin names granted to this key/team beyond the public (enabled) set teams LiteLLM_TeamTable[] projects LiteLLM_ProjectTable[] verification_tokens LiteLLM_VerificationToken[] diff --git a/litellm/models/object_permission.py b/litellm/models/object_permission.py index a09d50ddc33..f178c5ad47f 100644 --- a/litellm/models/object_permission.py +++ b/litellm/models/object_permission.py @@ -23,3 +23,4 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase): blocked_tools: list[str] | None = [] search_tools: list[str] | None = [] mcp_tool_search_enabled: bool | None = None + skills: list[str] | None = None diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 3f060f1f5f6..f0af17ab818 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -5721,8 +5721,26 @@ "paths": { "/claude-code/marketplace.json": { "get": { - "description": "Serve marketplace.json for Claude Code plugin discovery.\n\nThis endpoint is accessed by Claude Code CLI when users run:\n- claude plugin marketplace add \n- claude plugin install @\n\nReturns:\n Marketplace catalog with list of available plugins and their git sources.\n\nExample:\n ```bash\n claude plugin marketplace add http://localhost:4000/claude-code/marketplace.json\n claude plugin install my-plugin@litellm\n ```", + "description": "Serve marketplace.json for Claude Code plugin discovery.\n\nThis endpoint is accessed by Claude Code CLI when users run:\n- claude plugin marketplace add \n- claude plugin install @\n\nWithout `key` the catalog holds the enabled (public) plugins. With `?key=sk-...`\nthe key is authenticated and the catalog also holds the disabled plugins granted\nto it through `object_permission.skills` on the key or its team.\n\nReturns:\n Marketplace catalog with list of available plugins and their git sources.\n\nExample:\n ```bash\n claude plugin marketplace add http://localhost:4000/claude-code/marketplace.json\n claude plugin marketplace add \"http://localhost:4000/claude-code/marketplace.json?key=sk-...\"\n claude plugin install my-plugin@litellm\n ```", "operationId": "get_marketplace_claude_code_marketplace_json_get", + "parameters": [ + { + "in": "query", + "name": "key", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Key" + } + } + ], "responses": { "200": { "content": { @@ -5731,6 +5749,16 @@ } }, "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" } }, "summary": "Get Marketplace", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index f496140e905..23dc8237160 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -497,6 +497,9 @@ class LiteLLMRoutes(enum.Enum): "/v1/messages/count_tokens", "/v1/skills", "/v1/skills/{skill_id}", + "/claude-code/marketplace.json", + "/claude-code/plugins", + "/claude-code/plugins/{plugin_name}", ] # MCP tool-call / passthrough routes — data-plane. Gated by DISABLE_LLM_API_ENDPOINTS. @@ -1124,6 +1127,7 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase): models: list[str] | None = None search_tools: list[str] | None = None mcp_tool_search_enabled: bool | None = None + skills: list[str] | None = None from litellm.models.team import BudgetLimitEntry as BudgetLimitEntry # noqa: E402 diff --git a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py index cdeb1c6a111..7c6a4571948 100644 --- a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py +++ b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py @@ -7,10 +7,10 @@ Actual plugin files are hosted on GitHub/GitLab/Bitbucket or as a zip archive on any HTTPS host (S3, Artifactory, a static file server). Endpoints: -/claude-code/marketplace.json - GET - List plugins for Claude Code discovery (unauthenticated) +/claude-code/marketplace.json - GET - List plugins for Claude Code discovery (unauthenticated; `?key=` adds the key's granted skills) /claude-code/plugins - POST - Register a new plugin (create-only, proxy admin only) -/claude-code/plugins - GET - List plugins (any authenticated key) -/claude-code/plugins/{name} - GET - Get plugin details (any authenticated key) +/claude-code/plugins - GET - List plugins visible to the key (enabled, plus granted disabled ones) +/claude-code/plugins/{name} - GET - Get plugin details (403 on a disabled plugin the key is not granted) /claude-code/plugins/{name} - PUT - Update an existing plugin (proxy admin only) /claude-code/plugins/{name}/enable - POST - Enable a plugin (proxy admin only) /claude-code/plugins/{name}/disable - POST - Disable a plugin (proxy admin only) @@ -24,11 +24,15 @@ from datetime import datetime, timezone from typing import Annotated, Final, Protocol, TypedDict from urllib.parse import urlsplit -from fastapi import APIRouter, Depends, HTTPException +from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import JSONResponse from litellm._logging import verbose_proxy_logger -from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth +from litellm.proxy._types import CommonProxyErrors, ProxyException, UserAPIKeyAuth +from litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_skill_access import ( + SkillVisibility, + skill_visibility, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.resource_ownership import is_proxy_admin from litellm.repositories.table_repositories import ClaudeCodePluginRepository @@ -84,7 +88,7 @@ async def _get_prisma_client() -> object: "/claude-code/marketplace.json", tags=["Claude Code Marketplace"], ) -async def get_marketplace(): +async def get_marketplace(request: Request, key: str | None = None): """ Serve marketplace.json for Claude Code plugin discovery. @@ -92,24 +96,35 @@ async def get_marketplace(): - claude plugin marketplace add - claude plugin install @ + Without `key` the catalog holds the enabled (public) plugins. With `?key=sk-...` + the key is authenticated and the catalog also holds the disabled plugins granted + to it through `object_permission.skills` on the key or its team. + Returns: Marketplace catalog with list of available plugins and their git sources. Example: ```bash claude plugin marketplace add http://localhost:4000/claude-code/marketplace.json + claude plugin marketplace add "http://localhost:4000/claude-code/marketplace.json?key=sk-..." claude plugin install my-plugin@litellm ``` """ try: prisma_client: Final = await _get_prisma_client() + caller: Final[UserAPIKeyAuth | None] = ( + await user_api_key_auth(request=request, api_key=f"Bearer {key}") if key else None + ) + visibility: Final[SkillVisibility] = skill_visibility(caller) plugins: Final[Sequence[_PluginRecord]] = await ClaudeCodePluginRepository(prisma_client).table.find_many( - where={"enabled": True} + where=visibility.where() ) plugin_list: Final = [] for plugin in plugins: + if not visibility.allows(plugin): + continue try: manifest: Mapping[str, object] = json.loads(plugin.manifest_json or "{}") except json.JSONDecodeError: @@ -149,7 +164,7 @@ async def get_marketplace(): return JSONResponse(content=marketplace) - except HTTPException: + except (HTTPException, ProxyException): raise except Exception as e: verbose_proxy_logger.exception("Error generating marketplace: %s", e) @@ -395,13 +410,15 @@ async def list_plugins( try: prisma_client: Final = await _get_prisma_client() - where: Final = {"enabled": True} if enabled_only else {} + visibility: Final[SkillVisibility] = skill_visibility(user_api_key_dict) plugins: Final[Sequence[_PluginRecord]] = await ClaudeCodePluginRepository(prisma_client).table.find_many( - where=where + where={"enabled": True} if enabled_only else visibility.where() ) plugin_list: Final = [] for p in plugins: + if not visibility.allows(p): + continue # Parse manifest to get additional fields manifest = json.loads(p.manifest_json) if p.manifest_json else {} @@ -473,6 +490,12 @@ async def get_plugin( detail={"error": f"Plugin '{plugin_name}' not found"}, ) + if not skill_visibility(user_api_key_dict).allows(plugin): + raise HTTPException( + status_code=403, + detail={"error": f"Plugin '{plugin_name}' is not granted to this key"}, + ) + manifest: Final[Mapping[str, object]] = json.loads(plugin.manifest_json or "{}") if plugin.manifest_json else {} return { diff --git a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_skill_access.py b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_skill_access.py new file mode 100644 index 00000000000..e1e1ee6f160 --- /dev/null +++ b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_skill_access.py @@ -0,0 +1,67 @@ +""" +Claude Code marketplace visibility: enabled plugins are public, disabled plugins +are private and resolve only for proxy admins or keys granted them via +``object_permission.skills``. +""" + +from dataclasses import dataclass +from typing import TYPE_CHECKING, Final, Protocol + +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth +from litellm.proxy.common_utils.resource_ownership import is_proxy_admin + +if TYPE_CHECKING: + from prisma.types import LiteLLM_ClaudeCodePluginTableWhereInput + + +class _SkillRecord(Protocol): + name: str + enabled: bool + + +def _skills_of(permission: LiteLLM_ObjectPermissionTable | None) -> frozenset[str]: + return frozenset(permission.skills or ()) if permission is not None else frozenset() + + +def granted_skills(user_api_key_dict: UserAPIKeyAuth) -> frozenset[str]: + """Key grant intersected with the team grant when both are non-empty; either alone applies as is. + + An empty list is the Prisma column default for every object-permission row, so it means + "no private grants configured here" and defers to the other scope, same as the agents check. + """ + key_skills: Final = _skills_of(user_api_key_dict.object_permission) + team_skills: Final = _skills_of(user_api_key_dict.team_object_permission) + match (bool(key_skills), bool(team_skills)): + case (True, True): + return key_skills & team_skills + case (True, False): + return key_skills + case _: + return team_skills + + +@dataclass(frozen=True, slots=True) +class SkillVisibility: + granted: frozenset[str] + sees_private: bool + + def allows(self, skill: _SkillRecord) -> bool: + return skill.enabled or self.sees_private or skill.name in self.granted + + def where(self) -> "LiteLLM_ClaudeCodePluginTableWhereInput": + if self.sees_private: + return {} + if not self.granted: + return {"enabled": True} + return {"OR": [{"enabled": True}, {"name": {"in": sorted(self.granted)}}]} + + +PUBLIC_ONLY: Final = SkillVisibility(granted=frozenset(), sees_private=False) + + +def skill_visibility(user_api_key_dict: UserAPIKeyAuth | None) -> SkillVisibility: + if user_api_key_dict is None: + return PUBLIC_ONLY + if is_proxy_admin(user_api_key_dict): + return SkillVisibility(granted=frozenset(), sees_private=True) + return SkillVisibility(granted=granted_skills(user_api_key_dict), sees_private=False) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 05c5aad9303..817df082d8c 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -282,6 +282,7 @@ model LiteLLM_ObjectPermissionTable { mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user search_tools String[] @default([]) // search_tool_name values this key/team/user may call mcp_tool_search_enabled Boolean? + skills String[] @default([]) // Claude Code plugin names granted to this key/team beyond the public (enabled) set teams LiteLLM_TeamTable[] projects LiteLLM_ProjectTable[] verification_tokens LiteLLM_VerificationToken[] diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py index 6b1f9c68e47..7736939c696 100644 --- a/litellm/repositories/object_permission_repository.py +++ b/litellm/repositories/object_permission_repository.py @@ -40,6 +40,7 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): blocked_tools: list[str] | None = None, mcp_toolsets: list[str] | None = None, search_tools: list[str] | None = None, + skills: list[str] | None = None, ) -> LiteLLM_ObjectPermissionTable: """Create a new object permission record.""" data: Final[dict[str, Any]] = {} @@ -63,6 +64,8 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): data["mcp_toolsets"] = mcp_toolsets if search_tools is not None: data["search_tools"] = search_tools + if skills is not None: + data["skills"] = skills return await self.create(data) @@ -79,6 +82,7 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): blocked_tools: list[str] | None = None, mcp_toolsets: list[str] | None = None, search_tools: list[str] | None = None, + skills: list[str] | None = None, ) -> LiteLLM_ObjectPermissionTable | None: """Update an object permission record.""" data: Final[dict[str, Any]] = {} @@ -102,6 +106,8 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): data["mcp_toolsets"] = mcp_toolsets if search_tools is not None: data["search_tools"] = search_tools + if skills is not None: + data["skills"] = skills return await self.update(object_permission_id, data, id_field="object_permission_id") diff --git a/litellm/types/object_permission.py b/litellm/types/object_permission.py index 1b391a3a1ef..68661aed891 100644 --- a/litellm/types/object_permission.py +++ b/litellm/types/object_permission.py @@ -8,7 +8,7 @@ can adopt the type without violating the SDK-must-not-import-from-proxy layering rule. """ -from typing_extensions import TypedDict +from typing_extensions import ReadOnly, TypedDict class ObjectPermissionDict(TypedDict, total=False): @@ -23,3 +23,4 @@ class ObjectPermissionDict(TypedDict, total=False): models: list[str] | None search_tools: list[str] | None mcp_tool_search_enabled: bool | None + skills: ReadOnly[list[str] | None] diff --git a/schema.prisma b/schema.prisma index 05c5aad9303..817df082d8c 100644 --- a/schema.prisma +++ b/schema.prisma @@ -282,6 +282,7 @@ model LiteLLM_ObjectPermissionTable { mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user search_tools String[] @default([]) // search_tool_name values this key/team/user may call mcp_tool_search_enabled Boolean? + skills String[] @default([]) // Claude Code plugin names granted to this key/team beyond the public (enabled) set teams LiteLLM_TeamTable[] projects LiteLLM_ProjectTable[] verification_tokens LiteLLM_VerificationToken[] diff --git a/tests/pass_through_unit_tests/test_claude_code_marketplace.py b/tests/pass_through_unit_tests/test_claude_code_marketplace.py index 2ca81f1d5d3..bedb8830559 100644 --- a/tests/pass_through_unit_tests/test_claude_code_marketplace.py +++ b/tests/pass_through_unit_tests/test_claude_code_marketplace.py @@ -216,7 +216,7 @@ async def test_get_marketplace(mock_prisma_client): ) # Now get the marketplace - response = await get_marketplace() + response = await get_marketplace(request=MagicMock()) # Response is a JSONResponse, get the body body = json.loads(response.body.decode()) diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py index aa2da262bae..a7c2bd7ba20 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py @@ -11,17 +11,20 @@ from fastapi import HTTPException from unittest.mock import AsyncMock, MagicMock import litellm -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, ProxyException, UserAPIKeyAuth from litellm.proxy.proxy_server import LitellmUserRoles from litellm.types.proxy.claude_code_endpoints import ( RegisterPluginRequest, UpdatePluginRequest, ) +from litellm.proxy.anthropic_endpoints.claude_code_endpoints import claude_code_marketplace from litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_marketplace import ( delete_plugin, disable_plugin, enable_plugin, get_marketplace, + get_plugin, + list_plugins, register_plugin, update_plugin, ) @@ -38,11 +41,15 @@ def _make_mock_prisma(): async def _find_unique(where): return store.get(where.get("name")) + def _matches(record, where) -> bool: + if "OR" in where: + return any(_matches(record, clause) for clause in where["OR"]) + if "enabled" in where and record.enabled != where["enabled"]: + return False + return "name" not in where or record.name in where["name"]["in"] + async def _find_many(where=None): - records = list(store.values()) - if where and "enabled" in where: - return [r for r in records if r.enabled == where["enabled"]] - return records + return [r for r in store.values() if _matches(r, where or {})] async def _create(data): record = MagicMock() @@ -52,6 +59,8 @@ def _make_mock_prisma(): record.description = data.get("description") record.manifest_json = data.get("manifest_json", "{}") record.enabled = data.get("enabled", True) + record.created_at = data.get("created_at") + record.updated_at = data.get("updated_at") store[data["name"]] = record return record @@ -271,13 +280,89 @@ async def test_get_marketplace_skips_plugin_with_null_manifest(): table = litellm.proxy.proxy_server.prisma_client.db.litellm_claudecodeplugintable await table.create(data={"name": "null-manifest-plugin", "manifest_json": None, "enabled": True}) - response = await get_marketplace() + response = await get_marketplace(request=MagicMock()) assert response.status_code == 200 body = json.loads(response.body) assert [plugin["name"] for plugin in body["plugins"]] == ["good-plugin"] +def _granted_user(skills: list[str]) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-granted", + user_id="granted-user", + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="perm-1", skills=skills), + ) + + +async def _register_public_and_private_plugins() -> None: + for name, enabled in (("public-skill", True), ("private-skill", False)): + await register_plugin( + request=RegisterPluginRequest(name=name, source=_GIT_SUBDIR_SOURCE, version="1.0.0"), + user_api_key_dict=_USER, + ) + if not enabled: + await disable_plugin(plugin_name=name, user_api_key_dict=_USER) + + +async def _listed_names(user: UserAPIKeyAuth) -> set[str]: + response = await list_plugins(user_api_key_dict=user) + return {plugin.name for plugin in response.plugins} + + +@pytest.mark.asyncio +async def test_list_plugins_shows_disabled_plugin_only_to_granted_key_or_admin(): + await _register_public_and_private_plugins() + + assert await _listed_names(_NON_ADMIN_USER) == {"public-skill"} + assert await _listed_names(_granted_user(["other-skill"])) == {"public-skill"} + assert await _listed_names(_granted_user(["private-skill"])) == {"public-skill", "private-skill"} + assert await _listed_names(_USER) == {"public-skill", "private-skill"} + + +@pytest.mark.asyncio +async def test_get_plugin_returns_403_for_disabled_plugin_the_key_is_not_granted(): + await _register_public_and_private_plugins() + + with pytest.raises(HTTPException) as exc_info: + await get_plugin(plugin_name="private-skill", user_api_key_dict=_NON_ADMIN_USER) + assert exc_info.value.status_code == 403 + + assert (await get_plugin(plugin_name="public-skill", user_api_key_dict=_NON_ADMIN_USER))["name"] == "public-skill" + granted = await get_plugin(plugin_name="private-skill", user_api_key_dict=_granted_user(["private-skill"])) + assert granted["name"] == "private-skill" + assert granted["enabled"] is False + + +async def _marketplace_names(key: str | None) -> list[str]: + response = await get_marketplace(request=MagicMock(), key=key) + assert response.status_code == 200 + return sorted(plugin["name"] for plugin in json.loads(response.body)["plugins"]) + + +@pytest.mark.asyncio +async def test_get_marketplace_key_query_param_adds_granted_disabled_plugins(monkeypatch): + await _register_public_and_private_plugins() + keys = {"sk-granted": _granted_user(["private-skill"]), "sk-plain": _NON_ADMIN_USER} + + async def _fake_auth(request, api_key: str) -> UserAPIKeyAuth: + token = api_key.removeprefix("Bearer ") + if token not in keys: + raise ProxyException(message="invalid key", type="auth_error", param="key", code=401) + return keys[token] + + monkeypatch.setattr(claude_code_marketplace, "user_api_key_auth", _fake_auth) + + assert await _marketplace_names(None) == ["public-skill"] + assert await _marketplace_names("sk-plain") == ["public-skill"] + assert await _marketplace_names("sk-granted") == ["private-skill", "public-skill"] + + with pytest.raises(ProxyException) as exc_info: + await get_marketplace(request=MagicMock(), key="sk-bogus") + assert exc_info.value.code == "401" + + @pytest.mark.asyncio async def test_register_plugin_git_subdir_missing_url(): """git-subdir without url field raises HTTP 400.""" diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_skill_access.py b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_skill_access.py new file mode 100644 index 00000000000..dc4016d6db0 --- /dev/null +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_skill_access.py @@ -0,0 +1,79 @@ +"""Unit tests for claude_code_skill_access.py: key/team grant resolution.""" + +from unittest.mock import MagicMock + +import pytest + +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_skill_access import ( + granted_skills, + skill_visibility, +) + + +def _perm(skills: list[str] | None) -> LiteLLM_ObjectPermissionTable: + return LiteLLM_ObjectPermissionTable(object_permission_id="perm", skills=skills) + + +def _key(key_skills: list[str] | None, team_skills: list[str] | None) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-user", + user_role=LitellmUserRoles.INTERNAL_USER, + object_permission=_perm(key_skills) if key_skills is not None else None, + team_object_permission=_perm(team_skills) if team_skills is not None else None, + ) + + +def _plugin(name: str, enabled: bool) -> MagicMock: + record = MagicMock() + record.name = name + record.enabled = enabled + return record + + +@pytest.mark.parametrize( + ("key_skills", "team_skills", "expected"), + [ + (None, None, frozenset()), + ([], [], frozenset()), + (["a", "b"], None, frozenset({"a", "b"})), + (None, ["a", "b"], frozenset({"a", "b"})), + (["a", "b"], ["b", "c"], frozenset({"b"})), + (["a"], ["c"], frozenset()), + (["a", "b"], [], frozenset({"a", "b"})), + ([], ["a", "b"], frozenset({"a", "b"})), + ], +) +def test_granted_skills_intersects_key_with_team(key_skills, team_skills, expected): + assert granted_skills(_key(key_skills, team_skills)) == expected + + +def test_visibility_enabled_plugin_is_public_for_everyone(): + public = _plugin("public-skill", enabled=True) + + assert skill_visibility(None).allows(public) + assert skill_visibility(_key(None, None)).allows(public) + assert skill_visibility(_key([], ["other"])).allows(public) + + +def test_visibility_disabled_plugin_needs_grant_or_admin(): + private = _plugin("private-skill", enabled=False) + admin = UserAPIKeyAuth(api_key="sk-1234", user_role=LitellmUserRoles.PROXY_ADMIN) + + assert not skill_visibility(None).allows(private) + assert not skill_visibility(_key(None, None)).allows(private) + assert not skill_visibility(_key(["other-skill"], None)).allows(private) + assert not skill_visibility(_key(["private-skill"], ["other-skill"])).allows(private) + assert skill_visibility(_key(["private-skill"], None)).allows(private) + assert skill_visibility(_key(None, ["private-skill"])).allows(private) + assert skill_visibility(admin).allows(private) + + +def test_where_clause_bounds_the_plugin_query_to_what_the_caller_may_see(): + admin = UserAPIKeyAuth(api_key="sk-1234", user_role=LitellmUserRoles.PROXY_ADMIN) + + assert skill_visibility(None).where() == {"enabled": True} + assert skill_visibility(_key(None, None)).where() == {"enabled": True} + assert skill_visibility(_key(["b", "a"], None)).where() == {"OR": [{"enabled": True}, {"name": {"in": ["a", "b"]}}]} + assert skill_visibility(_key(["a", "b"], ["b"])).where() == {"OR": [{"enabled": True}, {"name": {"in": ["b"]}}]} + assert skill_visibility(admin).where() == {} diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 0fd19f518d9..4f2ff1582dd 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -3826,3 +3826,17 @@ def test_team_disable_logging_stays_proxy_admin_only(): def test_neighbouring_team_routes_stay_closed(route): """The grant is the callback paths and nothing else on the team namespace.""" assert "Only proxy admin" in _gate(route, LitellmUserRoles.INTERNAL_USER.value) + + +@pytest.mark.parametrize( + "route", + [ + "/claude-code/marketplace.json", + "/claude-code/plugins", + "/claude-code/plugins/my-skill", + ], +) +def test_claude_code_marketplace_routes_open_to_internal_users(route): + """Per-skill visibility is enforced inside the handler, so the route gate must let non-admins through.""" + assert RouteChecks.is_llm_api_route(route) is True + assert _gate(route, LitellmUserRoles.INTERNAL_USER.value) == "allowed" diff --git a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py index 1225cb80224..bd59a82cbd2 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py @@ -763,6 +763,7 @@ _EXPECTED_CUSTOMER = { "blocked_tools": [], "search_tools": [], "mcp_tool_search_enabled": None, + "skills": None, }, } diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index d7ebb1f60bf..078315c2bf8 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -122,6 +122,29 @@ async def test_set_object_permission_persists_mcp_tool_search_enabled(): assert created_data["mcp_tool_search_enabled"] is True +@pytest.mark.asyncio +async def test_set_object_permission_persists_skills(): + mock_prisma_client = MagicMock() + mock_created_permission = MagicMock() + mock_created_permission.object_permission_id = "perm_id" + mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( + return_value=mock_created_permission + ) + + data_json = { + "object_permission": LiteLLM_ObjectPermissionBase(skills=["private-skill"]).model_dump(), + } + + await _set_object_permission(data_json=data_json, prisma_client=mock_prisma_client) + + created_data = ( + mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[ + "data" + ] + ) + assert created_data["skills"] == ["private-skill"] + + # ---- Tests for _extract_requested_mcp_server_ids ---- diff --git a/ui/litellm-dashboard/src/components/Teams.test.tsx b/ui/litellm-dashboard/src/components/Teams.test.tsx index c7df35197e6..70c30596c75 100644 --- a/ui/litellm-dashboard/src/components/Teams.test.tsx +++ b/ui/litellm-dashboard/src/components/Teams.test.tsx @@ -33,6 +33,14 @@ vi.mock("./mcp_server_management/MCPServerSelector", () => ({ ), })); +vi.mock("./skills/SkillSelector", () => ({ + default: ({ onChange }: { onChange: (selected: string[]) => void }) => ( + + ), +})); + const can = vi.fn(); vi.mock("@/app/(dashboard)/hooks/useCan", () => ({ default: (...args: unknown[]) => can(...args), @@ -1359,6 +1367,27 @@ describe("Teams - the exact bytes the create call sends", () => { }); }); + it("puts the selected skills into object_permission.skills and drops the form key", async () => { + await openCreateModal(); + await openSection("Skill Settings", /Allowed Skills/); + fireEvent.click(screen.getByTestId("select-private-skill")); + + const payload = await submit(); + + expect(payload.object_permission).toStrictEqual({ skills: ["private-skill"] }); + expect(payload).not.toHaveProperty("object_permission_skills"); + }); + + it("sends no object_permission when Skill Settings is opened but nothing is selected", async () => { + await openCreateModal(); + await openSection("Skill Settings", /Allowed Skills/); + + const payload = await submit(); + + expect(payload).not.toHaveProperty("object_permission"); + expect(payload).not.toHaveProperty("object_permission_skills"); + }); + it("includes selected MCP toolsets in the create object permission", async () => { await openCreateModal(); await openSection("MCP Settings", /Allowed MCP Servers/); diff --git a/ui/litellm-dashboard/src/components/Teams.tsx b/ui/litellm-dashboard/src/components/Teams.tsx index c4060163c78..d2d18139893 100644 --- a/ui/litellm-dashboard/src/components/Teams.tsx +++ b/ui/litellm-dashboard/src/components/Teams.tsx @@ -50,6 +50,7 @@ import { Organization, getDefaultTeamSettings, getGuardrailsList, getPoliciesLis import NumericalInput from "./shared/numerical_input"; import VectorStoreSelector from "./vector_store_management/VectorStoreSelector"; import SearchToolSelector from "./search_tools/SearchToolSelector"; +import SkillSelector from "./skills/SkillSelector"; import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog"; interface TeamProps { @@ -99,6 +100,7 @@ const teamCreateFieldsSchema = z.object({ mcp_tool_permissions: z.record(z.string(), z.array(z.string())).optional(), allowed_agents_and_groups: z.object({ agents: z.array(z.string()), accessGroups: z.array(z.string()) }).optional(), object_permission_search_tools: z.array(z.string()).optional(), + object_permission_skills: z.array(z.string()).optional(), }); type TeamCreateFormValues = z.infer; @@ -128,6 +130,7 @@ const EMPTY_TEAM_CREATE_VALUES: TeamCreateFormValues = { mcp_tool_permissions: {}, allowed_agents_and_groups: undefined, object_permission_search_tools: undefined, + object_permission_skills: undefined, }; const ADDITIONAL_SETTINGS_FIELDS = [ @@ -147,6 +150,7 @@ const ADDITIONAL_SETTINGS_FIELDS = [ const MCP_SETTINGS_FIELDS = ["allowed_mcp_servers_and_groups", "mcp_tool_permissions"] as const; const AGENT_SETTINGS_FIELDS = ["allowed_agents_and_groups"] as const; const SEARCH_TOOL_SETTINGS_FIELDS = ["object_permission_search_tools"] as const; +const SKILL_SETTINGS_FIELDS = ["object_permission_skills"] as const; const isParsableJson = (value: string | undefined): boolean => { if (!value) { @@ -214,6 +218,7 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser const [mcpSettingsOpen, setMcpSettingsOpen] = useState(false); const [agentSettingsOpen, setAgentSettingsOpen] = useState(false); const [searchToolSettingsOpen, setSearchToolSettingsOpen] = useState(false); + const [skillSettingsOpen, setSkillSettingsOpen] = useState(false); const adminOrgs = useMemo( () => getAdminOrganizations(userRole, userID, organizations), @@ -505,6 +510,14 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser delete formValues.object_permission_search_tools; } + if (Array.isArray(formValues.object_permission_skills) && formValues.object_permission_skills.length > 0) { + if (!formValues.object_permission) { + formValues.object_permission = {}; + } + formValues.object_permission.skills = formValues.object_permission_skills; + } + delete formValues.object_permission_skills; + // Add model_aliases if any are defined if (Object.keys(modelAliases).length > 0) { formValues.model_aliases = modelAliases; @@ -540,6 +553,7 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser ...(mcpSettingsOpen ? [] : MCP_SETTINGS_FIELDS), ...(agentSettingsOpen ? [] : AGENT_SETTINGS_FIELDS), ...(searchToolSettingsOpen ? [] : SEARCH_TOOL_SETTINGS_FIELDS), + ...(skillSettingsOpen ? [] : SKILL_SETTINGS_FIELDS), ]); return Object.fromEntries(Object.entries(values).filter(([key]) => !unmounted.has(key))); }; @@ -1163,6 +1177,38 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser + + + Skill Settings + + + + + {({ value, onChange }) => ( + + )} + + + + Logging Settings diff --git a/ui/litellm-dashboard/src/components/object_permissions_view.tsx b/ui/litellm-dashboard/src/components/object_permissions_view.tsx index e90645a9e9f..ae46729222a 100644 --- a/ui/litellm-dashboard/src/components/object_permissions_view.tsx +++ b/ui/litellm-dashboard/src/components/object_permissions_view.tsx @@ -30,6 +30,7 @@ export function ObjectPermissionsView({ const agents = objectPermission?.agents || []; const agentAccessGroups = objectPermission?.agent_access_groups || []; const searchTools = objectPermission?.search_tools || []; + const skills = objectPermission?.skills || []; const content = (
@@ -58,6 +59,16 @@ export function ObjectPermissionsView({

{searchTools.join(", ")}

)}
+
+

Skills

+ {skills.length === 0 ? ( +

+ No private skills granted. Only enabled (public) Claude Code plugins are visible. +

+ ) : ( +

{skills.join(", ")}

+ )} +
); diff --git a/ui/litellm-dashboard/src/components/organisms/createKeyPayload.test.ts b/ui/litellm-dashboard/src/components/organisms/createKeyPayload.test.ts index 8b7283370a2..fef94cc3c2b 100644 --- a/ui/litellm-dashboard/src/components/organisms/createKeyPayload.test.ts +++ b/ui/litellm-dashboard/src/components/organisms/createKeyPayload.test.ts @@ -301,6 +301,16 @@ describe("object_permission", () => { ).toStrictEqual(aliasOnly({ object_permission: { agents: ["a-1"], agent_access_groups: ["ag-1"] } })); }); + it("moves the selected skills under object_permission and off the top level", () => { + expect(payloadOf(build({ key_alias: "my-key", allowed_skills: ["private-skill"] }))).toStrictEqual( + aliasOnly({ object_permission: { skills: ["private-skill"] } }), + ); + }); + + it("sends no object_permission for an empty skill selection", () => { + expect(payloadOf(build({ key_alias: "my-key", allowed_skills: [] }))).toStrictEqual(aliasOnly()); + }); + it("merges every source into a single object_permission", () => { const everySource = { key_alias: "my-key", @@ -308,6 +318,7 @@ describe("object_permission", () => { allowed_mcp_servers_and_groups: { servers: ["s-1"], accessGroups: ["g-1"], toolsets: ["t-1"] }, mcp_tool_permissions: { "s-1": ["read"] }, allowed_agents_and_groups: { agents: ["a-1"], accessGroups: ["ag-1"] }, + allowed_skills: ["private-skill"], }; expect(payloadOf(build(everySource))).toStrictEqual( aliasOnly({ @@ -319,6 +330,7 @@ describe("object_permission", () => { mcp_tool_permissions: { "s-1": ["read"] }, agents: ["a-1"], agent_access_groups: ["ag-1"], + skills: ["private-skill"], }, }), ); diff --git a/ui/litellm-dashboard/src/components/organisms/createKeyPayload.ts b/ui/litellm-dashboard/src/components/organisms/createKeyPayload.ts index a528dbaeb2c..37a973d5def 100644 --- a/ui/litellm-dashboard/src/components/organisms/createKeyPayload.ts +++ b/ui/litellm-dashboard/src/components/organisms/createKeyPayload.ts @@ -112,6 +112,7 @@ interface PermissionSources { readonly toolPermissions: unknown | undefined; readonly extraMcpAccessGroups: unknown[] | undefined; readonly agents: AgentSelection | undefined; + readonly skills: unknown[] | undefined; } const readPermissionSources = (values: Record): PermissionSources => ({ @@ -120,6 +121,7 @@ const readPermissionSources = (values: Record): PermissionSourc toolPermissions: readToolPermissions(values.mcp_tool_permissions), extraMcpAccessGroups: nonEmptyList(values.allowed_mcp_access_groups), agents: readAgentSelection(values.allowed_agents_and_groups), + skills: nonEmptyList(values.allowed_skills), }); const buildObjectPermission = ({ @@ -128,6 +130,7 @@ const buildObjectPermission = ({ toolPermissions, extraMcpAccessGroups, agents, + skills, }: PermissionSources): Record | undefined => { const permission: Record = { ...(vectorStores && { vector_stores: vectorStores }), @@ -138,6 +141,7 @@ const buildObjectPermission = ({ ...(extraMcpAccessGroups && { mcp_access_groups: extraMcpAccessGroups }), ...(agents?.agents && { agents: agents.agents }), ...(agents?.accessGroups && { agent_access_groups: agents.accessGroups }), + ...(skills && { skills }), }; return Object.keys(permission).length > 0 ? permission : undefined; }; @@ -148,6 +152,7 @@ const consumedSourceKeys = ( ): ReadonlySet => new Set([ "mcp_tool_permissions", + "allowed_skills", ...(values.disable_global_guardrails ? [] : ["disable_global_guardrails"]), ...(vectorStores ? ["allowed_vector_store_ids"] : []), ...(mcp ? ["allowed_mcp_servers_and_groups"] : []), diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.integration.test.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.integration.test.tsx index 3471ef00eb0..2bf39bf4cd7 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.integration.test.tsx @@ -84,6 +84,13 @@ vi.mock("../networking", async (importOriginal) => { getPossibleUserRoles: vi.fn().mockResolvedValue({}), userFilterUICall: vi.fn().mockResolvedValue([]), getAgentsList: vi.fn().mockResolvedValue({ agents: [] }), + getClaudeCodePluginsList: vi.fn().mockResolvedValue({ + plugins: [ + { name: "public-skill", enabled: true }, + { name: "private-skill", enabled: false }, + ], + count: 2, + }), getPassThroughEndpointsCall: vi.fn().mockResolvedValue({ endpoints: [] }), vectorStoreListCall: vi.fn().mockResolvedValue({ data: [] }), listMCPTools: vi.fn().mockResolvedValue(emptyMcpTools), @@ -109,6 +116,7 @@ const OPENAPI_SCHEMA = { const SECTIONS = { mcp: /MCP Settings/i, agent: /Agent Settings/i, + skill: /Skill Settings/i, logging: /Logging Settings/i, router: /Router Settings/i, aliases: /Model Aliases/i, @@ -164,6 +172,7 @@ const ROUTER_SETTINGS_DEFAULT = { const SECTION_PAYLOAD_ADDITIONS: Record> = { mcp: { allowed_mcp_servers_and_groups: { servers: [], accessGroups: [] } }, agent: { allowed_agents_and_groups: undefined }, + skill: {}, logging: {}, router: { router_settings: ROUTER_SETTINGS_DEFAULT }, aliases: {}, @@ -329,6 +338,21 @@ describe("CreateKey", () => { expect(Object.keys(serialised).sort()).toStrictEqual([...wireKeys].sort()); }); + it("moves a picked private skill under object_permission.skills and off the top level", async () => { + await openModal(); + await nameTheKey(); + await openSection(/Optional Settings/i); + await openSection(SECTIONS.skill); + await userEvent.click(await screen.findByRole("combobox", { name: "Select skills (optional)" })); + await userEvent.click(await screen.findByRole("option", { name: "private-skill (private)" })); + await userEvent.keyboard("{Escape}"); + await submit(); + + const payload = await createdPayload(); + expect(payload.object_permission).toStrictEqual({ skills: ["private-skill"] }); + expect(payload).not.toHaveProperty("allowed_skills"); + }); + it("omits a budget typed into a section the user closed again, rather than sending it as null", async () => { await openModal(); await nameTheKey(); diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx index 16ea99ea94c..831cb0cf6a6 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx @@ -27,6 +27,7 @@ import React, { useEffect, useMemo, useRef, useState } from "react"; import { type Control, useForm, useWatch, type UseFormSetValue } from "react-hook-form"; import { rolesWithWriteAccess } from "../../utils/roles"; import AgentSelector from "../agent_management/AgentSelector"; +import SkillSelector from "../skills/SkillSelector"; import AccessGroupSelector from "../common_components/AccessGroupSelector"; import BudgetDurationDropdown from "../common_components/budget_duration_dropdown"; import SchemaFormFields from "../common_components/check_openapi_schema"; @@ -1557,6 +1558,36 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp + + + Skill Settings + + + + + Allowed Skills{" "} + + + + + } + name="allowed_skills" + help="Select private skills this key can access in the Claude Code marketplace" + > + {(control) => ( + + )} + + + + {premiumUser ? ( diff --git a/ui/litellm-dashboard/src/components/skills/SkillSelector.test.tsx b/ui/litellm-dashboard/src/components/skills/SkillSelector.test.tsx new file mode 100644 index 00000000000..59aff7c25de --- /dev/null +++ b/ui/litellm-dashboard/src/components/skills/SkillSelector.test.tsx @@ -0,0 +1,54 @@ +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { renderWithProviders, screen } from "../../../tests/test-utils"; +import { getClaudeCodePluginsList } from "../networking"; +import SkillSelector from "./SkillSelector"; + +vi.mock("../networking", () => ({ + getClaudeCodePluginsList: vi.fn(), +})); + +describe("SkillSelector", () => { + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(getClaudeCodePluginsList).mockResolvedValue({ + plugins: [ + { name: "public-skill", enabled: true }, + { name: "private-skill", enabled: false }, + ], + count: 2, + }); + }); + + it("should list marketplace plugins and mark disabled ones as private", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("combobox")); + + expect(await screen.findByRole("option", { name: "public-skill" })).toBeInTheDocument(); + expect(screen.getByRole("option", { name: "private-skill (private)" })).toBeInTheDocument(); + expect(getClaudeCodePluginsList).toHaveBeenCalledWith("token"); + }); + + it("should report the selected skill names", async () => { + const user = userEvent.setup(); + const onChange = vi.fn(); + renderWithProviders(); + + await user.click(screen.getByRole("combobox")); + await user.click(await screen.findByRole("option", { name: "private-skill (private)" })); + + expect(onChange).toHaveBeenCalledWith(["private-skill"]); + }); + + it("should clear all selected skills", async () => { + const user = userEvent.setup(); + const onChange = vi.fn(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: "Clear all skills" })); + + expect(onChange).toHaveBeenCalledWith([]); + }); +}); diff --git a/ui/litellm-dashboard/src/components/skills/SkillSelector.tsx b/ui/litellm-dashboard/src/components/skills/SkillSelector.tsx new file mode 100644 index 00000000000..b21c8b7e872 --- /dev/null +++ b/ui/litellm-dashboard/src/components/skills/SkillSelector.tsx @@ -0,0 +1,109 @@ +import React, { useEffect, useState } from "react"; +import { + Combobox, + ComboboxChip, + ComboboxChips, + ComboboxChipsInput, + ComboboxClear, + ComboboxContent, + ComboboxEmpty, + ComboboxItem, + ComboboxList, + ComboboxValue, + useComboboxAnchor, +} from "@/components/ui/combobox"; +import { cn } from "@/lib/cva.config"; +import { getClaudeCodePluginsList } from "../networking"; + +export interface SkillSelectorProps { + onChange: (selected: string[]) => void; + value?: string[]; + className?: string; + accessToken: string; + placeholder?: string; + disabled?: boolean; +} + +interface SkillOption { + readonly name: string; + readonly enabled: boolean; +} + +const readSkillOptions = (data: unknown): SkillOption[] => { + const plugins = (data as { plugins?: unknown[] } | undefined)?.plugins; + if (!Array.isArray(plugins)) return []; + return plugins.flatMap((plugin) => { + const record = plugin as { name?: unknown; enabled?: unknown }; + return typeof record.name === "string" && record.name.length > 0 + ? [{ name: record.name, enabled: record.enabled !== false }] + : []; + }); +}; + +const SkillSelector: React.FC = ({ + onChange, + value, + className, + accessToken, + placeholder = "Select skills (optional)", + disabled = false, +}) => { + const anchor = useComboboxAnchor(); + const [options, setOptions] = useState([]); + const [loading, setLoading] = useState(false); + + useEffect(() => { + const load = async () => { + if (!accessToken) return; + setLoading(true); + try { + setOptions(readSkillOptions(await getClaudeCodePluginsList(accessToken))); + } catch (e) { + console.error("Failed to load skills:", e); + } finally { + setLoading(false); + } + }; + load(); + }, [accessToken]); + + return ( + option.name)} + value={value ?? []} + onValueChange={(selected: string[]) => onChange(selected)} + disabled={disabled} + > + } className={cn("w-full", className)} aria-busy={loading}> + + {(selected: string[]) => + selected.map((skill) => ( + + {skill} + + )) + } + + + {value && value.length > 0 && } + + + {loading ? "Loading skills…" : "No skills found"} + + {(skill: string) => { + const isPrivate = options.some((option) => option.name === skill && !option.enabled); + return ( + + {skill} + {isPrivate && private} + + ); + }} + + + + ); +}; + +export default SkillSelector; diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx index b286c8c1303..a25ca28651a 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx @@ -52,6 +52,7 @@ vi.mock("@/components/networking", () => ({ listMCPTools: vi.fn().mockResolvedValue({ tools: [] }), vectorStoreListCall: vi.fn().mockResolvedValue({ data: [] }), getAgentsList: vi.fn().mockResolvedValue({ agents: [] }), + getClaudeCodePluginsList: vi.fn().mockResolvedValue({ plugins: [], count: 0 }), })); const can = vi.fn(); @@ -1990,6 +1991,7 @@ describe("TeamInfoView - the exact bytes the update call sends", () => { agents: [], agent_access_groups: [], vector_stores: ["vs-1"], + skills: [], }; it("leaves every team member key out of the request body for an untouched save with both sections closed", async () => { @@ -2052,6 +2054,47 @@ describe("TeamInfoView - the exact bytes the update call sends", () => { expect(objectPermission.agent_access_groups).toStrictEqual([]); }); + it("resends the stored skills when the selector is left untouched", async () => { + const user = userEvent.setup({ delay: null }); + vi.mocked(networking.teamInfoCall).mockResolvedValue( + createMockTeamData({ models: ["gpt-4"], object_permission: { skills: ["private-skill"] } }), + ); + vi.mocked(networking.teamUpdateCall).mockResolvedValue({ data: {}, team_id: "123" } as any); + + renderWithProviders(); + await waitFor(() => expect(screen.queryAllByText("Test Team").length).toBeGreaterThan(0)); + await user.click(screen.getByRole("tab", { name: "Settings" })); + await user.click(await screen.findByRole("button", { name: /edit settings/i })); + await screen.findByLabelText("Team Name"); + + const payload = await save(user); + + const objectPermission = wireBody(payload).object_permission as Record; + expect(objectPermission.skills).toStrictEqual(["private-skill"]); + }); + + it("sends an empty skills array after the last skill chip is removed", async () => { + const user = userEvent.setup({ delay: null }); + vi.mocked(networking.teamInfoCall).mockResolvedValue( + createMockTeamData({ models: ["gpt-4"], object_permission: { skills: ["private-skill"] } }), + ); + vi.mocked(networking.teamUpdateCall).mockResolvedValue({ data: {}, team_id: "123" } as any); + + renderWithProviders(); + await waitFor(() => expect(screen.queryAllByText("Test Team").length).toBeGreaterThan(0)); + await user.click(screen.getByRole("tab", { name: "Settings" })); + await user.click(await screen.findByRole("button", { name: /edit settings/i })); + await screen.findByLabelText("Team Name"); + + await user.click(within(screen.getByLabelText("private-skill")).getByRole("button")); + expect(screen.queryByLabelText("private-skill")).not.toBeInTheDocument(); + + const payload = await save(user); + + const objectPermission = wireBody(payload).object_permission as Record; + expect(objectPermission.skills).toStrictEqual([]); + }); + it("sends an empty vector_stores array after the last vector store chip is removed", async () => { const user = userEvent.setup({ delay: null }); await openEditor(user); diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index 4947cec1441..988e2e082aa 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -89,6 +89,7 @@ import ObjectPermissionsView from "../object_permissions_view"; import NumericalInput from "../shared/numerical_input"; import VectorStoreSelector from "../vector_store_management/VectorStoreSelector"; import SearchToolSelector from "../search_tools/SearchToolSelector"; +import SkillSelector from "../skills/SkillSelector"; import EditLoggingSettings from "./EditLoggingSettings"; import RouterSettingsAccordion, { RouterSettingsAccordionRef } from "../common_components/RouterSettingsAccordion"; import MemberModal from "./EditMembership"; @@ -371,6 +372,7 @@ const teamUpdateFieldsSchema = z.object({ mcp_tool_permissions: z.record(z.string(), z.array(z.string())).optional(), agents_and_groups: z.object({ agents: z.array(z.string()), accessGroups: z.array(z.string()) }).optional(), object_permission_search_tools: z.array(z.string()).optional(), + object_permission_skills: z.array(z.string()).optional(), organization_id: z.string().nullish(), logging_settings: z.array(z.unknown()).optional(), secret_manager_settings: z.string().optional(), @@ -419,6 +421,7 @@ const EMPTY_TEAM_UPDATE_VALUES: TeamUpdateFormValues = { mcp_tool_permissions: {}, agents_and_groups: { agents: [], accessGroups: [] }, object_permission_search_tools: [], + object_permission_skills: [], organization_id: null, logging_settings: [], secret_manager_settings: "", @@ -485,6 +488,7 @@ const toTeamFormValues = (info: TeamInfoRecord, effectiveGuardrails: string[]): accessGroups: info.object_permission?.agent_access_groups || [], }, object_permission_search_tools: info.object_permission?.search_tools || [], + object_permission_skills: info.object_permission?.skills || [], organization_id: info.organization_id, logging_settings: info.metadata?.logging || [], secret_manager_settings: info.metadata?.secret_manager_settings @@ -1047,6 +1051,10 @@ const TeamInfoView: React.FC = ({ updateData.object_permission.search_tools = values.object_permission_search_tools; } + if (Array.isArray(values.object_permission_skills)) { + updateData.object_permission.skills = values.object_permission_skills; + } + // Pass access_group_ids to the update request if (values.access_group_ids !== undefined) { updateData.access_group_ids = values.access_group_ids; @@ -1848,6 +1856,24 @@ const TeamInfoView: React.FC = ({ + + {({ value, onChange }) => ( + + )} + + {({ id, value, onChange }) => ( ( <> @@ -53,6 +55,36 @@ export const KeyTypeSelect = ({ ); +const SKILLS_HINT = + "Enabled skills are visible to every key. Grant disabled (private) Claude Code plugins to this key here."; + +export const KeyAgentAndSkillFields = ({ + control, + accessToken, +}: { + control: Control; + accessToken: string; +}) => ( + <> + + {({ value, onChange }) => ( + + )} + + + + {({ value, onChange }) => ( + + )} + + +); + export const KeyBudgetNumberField = ({ control, name, diff --git a/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx b/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx index 8a42ecff50c..522af5a85ad 100644 --- a/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx @@ -366,6 +366,45 @@ describe("KeyInfoView handleKeyUpdate mcp_toolsets", () => { }); }); +describe("KeyInfoView handleKeyUpdate skills", () => { + it("should forward the skills the edit form supplies into object_permission and drop the form key", async () => { + renderView(true); + + fireEvent.click(screen.getByText("Settings")); + fireEvent.click(screen.getByText("Edit Settings")); + (globalThis as any).__TEST_FORM_VALUES = { + token: "tok_123", + skills: ["private-skill"], + }; + + fireEvent.click(screen.getByText("Mock Submit")); + + await waitFor(() => expect(keyUpdateCallMock).toHaveBeenCalled()); + + const [, sentPayload] = keyUpdateCallMock.mock.calls[0]; + expect(sentPayload.object_permission.skills).toEqual(["private-skill"]); + expect(sentPayload).not.toHaveProperty("skills"); + }); + + it("should send an explicit empty skills list when the form clears every skill", async () => { + renderView(true); + + fireEvent.click(screen.getByText("Settings")); + fireEvent.click(screen.getByText("Edit Settings")); + (globalThis as any).__TEST_FORM_VALUES = { + token: "tok_123", + skills: [], + }; + + fireEvent.click(screen.getByText("Mock Submit")); + + await waitFor(() => expect(keyUpdateCallMock).toHaveBeenCalled()); + + const [, sentPayload] = keyUpdateCallMock.mock.calls[0]; + expect(sentPayload.object_permission.skills).toEqual([]); + }); +}); + describe("KeyInfoView handleKeyUpdate budget_duration", () => { it("should send a canonical budget_duration through unchanged", async () => { renderView(true); diff --git a/ui/litellm-dashboard/src/components/templates/keyEditFormValues.ts b/ui/litellm-dashboard/src/components/templates/keyEditFormValues.ts index f24fb6e5a86..fec8e749143 100644 --- a/ui/litellm-dashboard/src/components/templates/keyEditFormValues.ts +++ b/ui/litellm-dashboard/src/components/templates/keyEditFormValues.ts @@ -46,6 +46,7 @@ export interface KeyEditFormValues { mcp_servers_and_groups?: McpServersAndGroups; mcp_tool_permissions?: Record; agents_and_groups?: AgentsAndGroups; + skills?: string[]; organization_id?: string | null; team_id?: string | null; logging_settings?: unknown[]; @@ -102,6 +103,7 @@ export const toKeyEditFormValues = (keyData: KeyResponse): KeyEditFormValues => agents: keyData.object_permission?.agents || [], accessGroups: keyData.object_permission?.agent_access_groups || [], }, + skills: keyData.object_permission?.skills || [], organization_id: keyData.organization_id, team_id: keyData.team_id, logging_settings: extractLoggingSettings(keyData.metadata), @@ -148,6 +150,7 @@ export const keyEditFormSchema = z.object({ mcp_servers_and_groups: z.custom(), mcp_tool_permissions: z.custom | undefined>(), agents_and_groups: z.custom(), + skills: z.custom(), organization_id: z.custom(), team_id: z.custom(), logging_settings: z.custom(), @@ -196,6 +199,7 @@ export const toSubmittedValues = ( mcp_servers_and_groups: values.mcp_servers_and_groups, mcp_tool_permissions: values.mcp_tool_permissions, agents_and_groups: values.agents_and_groups, + skills: values.skills, organization_id: values.organization_id, team_id: values.team_id, logging_settings: values.logging_settings, diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx index 35e9268e1c2..b2c2a381b42 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx @@ -59,6 +59,7 @@ vi.mock("../networking", async () => { agents: [], }), getAgentAccessGroups: vi.fn().mockResolvedValue([]), + getClaudeCodePluginsList: vi.fn().mockResolvedValue({ plugins: [], count: 0 }), }; }); @@ -135,6 +136,14 @@ vi.mock("../agent_management/AgentSelector", () => ({ ), })); +vi.mock("../skills/SkillSelector", () => ({ + default: ({ onChange }: { onChange: (selected: string[]) => void }) => ( + + ), +})); + vi.mock("../common_components/AccessGroupSelector", () => ({ default: ({ value = [], onChange }: { value?: string[]; onChange?: (v: string[]) => void }) => ( { mcp_servers_and_groups: { servers: [], accessGroups: [], toolsets: [] }, mcp_tool_permissions: {}, agents_and_groups: { agents: [], accessGroups: [] }, + skills: [], organization_id: null, team_id: null, logging_settings: [], @@ -2241,6 +2251,36 @@ describe("KeyEditView", () => { expect(onSubmitMock.mock.calls[0][0].agents_and_groups.agents).toEqual(["agent-1"]); }); + it("carries a picked skill into the payload", async () => { + const onSubmitMock = vi.fn().mockResolvedValue(undefined); + renderForPayload(onSubmitMock); + await screen.findByRole("button", { name: /save changes/i }); + + await userEvent.click(screen.getByRole("button", { name: "pick skill" })); + await userEvent.click(screen.getByRole("button", { name: /save changes/i })); + + await waitFor(() => { + expect(onSubmitMock).toHaveBeenCalled(); + }); + expect(onSubmitMock.mock.calls[0][0].skills).toEqual(["private-skill"]); + }); + + it("preloads the stored skills into the payload when the selector is left untouched", async () => { + const onSubmitMock = vi.fn().mockResolvedValue(undefined); + renderForPayload(onSubmitMock, { + ...MOCK_KEY_DATA, + object_permission: { ...MOCK_KEY_DATA.object_permission, skills: ["stored-skill"] }, + } as KeyResponse); + await screen.findByRole("button", { name: /save changes/i }); + + await userEvent.click(screen.getByRole("button", { name: /save changes/i })); + + await waitFor(() => { + expect(onSubmitMock).toHaveBeenCalled(); + }); + expect(onSubmitMock.mock.calls[0][0].skills).toEqual(["stored-skill"]); + }); + it("carries an added logging integration into the payload", async () => { const onSubmitMock = vi.fn().mockResolvedValue(undefined); renderForPayload(onSubmitMock); diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx index b1bbbeed6f9..e464dfe5008 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx @@ -15,7 +15,6 @@ import { FormField } from "@/components/shared/form/FormField"; import React, { useEffect, useRef, useState } from "react"; import { hasCapability } from "../../utils/capabilities"; import { isProxyAdminRole, rolesWithWriteAccess } from "../../utils/roles"; -import AgentSelector from "../agent_management/AgentSelector"; import AccessGroupSelector from "../common_components/AccessGroupSelector"; import BudgetDurationDropdown from "../common_components/budget_duration_dropdown"; import { mapInternalToDisplayNames } from "../callback_info_helpers"; @@ -32,9 +31,8 @@ import { modelSentinelOptions, parseAllowedRoutes, } from "./keyEditFieldNormalizers"; -import { KeyBudgetNumberField, KeyTypeSelect, labelWithHint } from "./KeyEditViewControls"; +import { KeyAgentAndSkillFields, KeyBudgetNumberField, KeyTypeSelect, labelWithHint } from "./KeyEditViewControls"; import { - AgentsAndGroups, KeyEditFormValues, keyEditFormSchema, McpServersAndGroups, @@ -762,16 +760,7 @@ export function KeyEditView({ /> - - {({ value, onChange }) => ( - - )} - + * - claude plugin install @ * + * Without `key` the catalog holds the enabled (public) plugins. With `?key=sk-...` + * the key is authenticated and the catalog also holds the disabled plugins granted + * to it through `object_permission.skills` on the key or its team. + * * Returns: * Marketplace catalog with list of available plugins and their git sources. * * Example: * ```bash * claude plugin marketplace add http://localhost:4000/claude-code/marketplace.json + * claude plugin marketplace add "http://localhost:4000/claude-code/marketplace.json?key=sk-..." * claude plugin install my-plugin@litellm * ``` */ @@ -29333,6 +29338,8 @@ export interface components { models?: string[] | null; /** Search Tools */ search_tools?: string[] | null; + /** Skills */ + skills?: string[] | null; /** Vector Stores */ vector_stores?: string[] | null; }; @@ -29386,6 +29393,8 @@ export interface components { * @default [] */ search_tools: string[] | null; + /** Skills */ + skills?: string[] | null; /** * Vector Stores * @default [] @@ -43206,7 +43215,9 @@ export interface operations { }; get_marketplace_claude_code_marketplace_json_get: { parameters: { - query?: never; + query?: { + key?: string | null; + }; header?: never; path?: never; cookie?: never; @@ -43222,6 +43233,15 @@ export interface operations { "application/json": unknown; }; }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; }; }; list_plugins_claude_code_plugins_get: { From 1987e6e290a8344515f595538e85c3ae00e69a96 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 10 Sep 2026 14:54:46 -0700 Subject: [PATCH 59/60] test(proxy): pass the request to get_marketplace in the archive marketplace test --- .../proxy/anthropic_endpoints/test_claude_code_marketplace.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py index a7c2bd7ba20..a585666743f 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py @@ -481,7 +481,7 @@ async def test_archive_source_registers_and_is_served_verbatim_in_marketplace(): assert response.action == "created" assert response.plugin.source == _ARCHIVE_SOURCE - marketplace = json.loads((await get_marketplace()).body) + marketplace = json.loads((await get_marketplace(request=MagicMock())).body) assert marketplace["plugins"] == [{"name": "s3-skill", "source": _ARCHIVE_SOURCE, "version": "1.0.0"}] From d98522b6f662fec6f600ab94817855cb7ff7edce Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 22:14:37 +0000 Subject: [PATCH 60/60] feat(proxy): share database connections across workers with an in-container pgbouncer (#39683) * feat(proxy): share database connections across workers with an in-container pgbouncer Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): parse pgbouncer options iteratively to satisfy the recursion gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): refuse pgbouncer with token db auth and retry failed pooler restarts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): build pgbouncer 1.25.2 from a pinned source archive and verify pooler replacements The public Wolfi repository only carries pgbouncer 1.24.1-r3, which the image scan rejects (CVE-2026-6664, CVE-2026-6665, CVE-2026-6666, CVE-2025-12819). All three images now compile the checksummed 1.25.2 release in a builder stage. The supervisor now waits for a replacement pooler to listen before treating it as recovered, ends and retries one that never does, and takes the same lock for stop() and spawn so no replacement can be started after shutdown began. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): refuse to start pgbouncer on a loopback port another process already owns Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): count pgbouncer ready only once its own unix socket answers, not any listener on the port Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): refuse pgbouncer older than 1.19, whose unix socket cannot vouch for the tcp port Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(helm,terraform): expose the in-container pgbouncer pool for the componentized gateway Add database.connectionPool to helm/litellm and gateway_connection_pool_* to terraform/litellm/aws so the componentized gateway can receive the LITELLM_PGBOUNCER_* env the classic image already honours. Both reject the pool under IAM or Entra token auth at render/plan time: the pooler holds one static database password for the life of the pod or task. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(gateway): launch the componentized gateway image through a pgbouncer-aware supervisor (#40592) The componentized gateway image started uvicorn directly, so the in-container PgBouncer never ran for it: every worker opened its own Prisma pool to the database. It also passed no keep-alive timeout, so behind a load balancer with a 60s idle timeout uvicorn's 5s default closed idle connections first and the balancer returned 502s on scale-out gateway.launch assembles DATABASE_URL, starts PgBouncer once per pod when LITELLM_PGBOUNCER_ENABLED is set, hands the workers the loopback URL and then runs uvicorn on gateway.main:app with KEEPALIVE_TIMEOUT as --timeout-keep-alive. The image builds PgBouncer 1.25.2 from a checksummed tarball, copies the compiled Rust extension into the /app source tree it imports from (it was only in site-packages, which PYTHONPATH=/app shadows) and asserts the native bridge loads. The app user is added to stats_users so operators can read the PgBouncer console with the application credentials The supervisor returns the pooled URL instead of writing into the mapping it was handed, a database user whose name PgBouncer would split into several stats_users entries is refused before the config is written, and the launcher tests drive main() with an injected serve callable Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(terraform): describe the gateway.launch pooler entrypoint in the aws module README Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): run pgbouncer exit hooks only in the parent and copy the CA into the runtime dir Gunicorn workers inherit the parent's atexit table, so a recycled worker (max_requests) stopped the shared pooler and removed its runtime dir, then hung in the inherited Popen lock. The hooks now no-op unless os.getpid() is the process that started PgBouncer A verified TLS upstream named the operator's CA bundle directly, which is often a 0600 root-owned file that nobody (the user PgBouncer drops to) cannot read, so every server connection failed with "failed to load CA". The bundle is copied into the runtime dir next to the ini and chowned with it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(terraform): run the gateway through gateway.launch and add the gcp connection-pool variables Cloud Run and ECS overrode the image command with uvicorn gateway.main:app, which skips the supervisor that starts the in-container PgBouncer, so LITELLM_PGBOUNCER_ENABLED was inert on both stacks. Both now exec python -m gateway.launch (under ddtrace-run when USE_DDTRACE is set), and the gcp module gains gateway_connection_pool_enabled / gateway_pool_max_db_connections / gateway_pool_max_client_conn wired to the gateway service only The test_launch password_env fixture now restores DATABASE_URL even when it was unset: monkeypatch.delenv records nothing for an absent var, so main() left postgresql://...@db.internal in the xdist worker's environ and the key-rotation e2e test in the same proxy-infra shard stopped skipping and tried to reach db.internal Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(pgbouncer): keep channel_binding and gssencmode off the loopback URL Prisma would demand TLS channel binding from a pooler that only speaks plain TCP on 127.0.0.1. Also pass the request the marketplace test started needing after #40518 landed on top of #40496 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/workflows/test-unit.yml | 1 + Dockerfile | 19 +- docker/Dockerfile.database | 19 +- docker/Dockerfile.non_root | 19 +- gateway/Dockerfile | 28 +- gateway/launch.py | 71 ++ helm/litellm-helm/templates/deployment.yaml | 8 + .../tests/connection_pool_tests.yaml | 61 ++ helm/litellm-helm/values.yaml | 14 + helm/litellm/templates/_helpers.tpl | 17 + .../litellm/templates/gateway/deployment.yaml | 3 + helm/litellm/tests/connection_pool_tests.yaml | 117 ++++ helm/litellm/values.yaml | 20 + litellm/proxy/db/pgbouncer.py | 525 +++++++++++++++ litellm/proxy/proxy_cli.py | 16 + terraform/litellm/aws/README.md | 33 + terraform/litellm/aws/ecs.tf | 14 +- .../aws/tests/connection_pool.tftest.hcl | 123 ++++ terraform/litellm/aws/variables.tf | 39 ++ terraform/litellm/gcp/README.md | 34 + terraform/litellm/gcp/cloudrun.tf | 10 +- .../gcp/tests/connection_pool.tftest.hcl | 135 ++++ terraform/litellm/gcp/variables.tf | 39 ++ tests/test_gateway/test_launch.py | 155 +++++ tests/test_litellm/proxy/db/test_pgbouncer.py | 624 ++++++++++++++++++ .../test_litellm/test_component_entrypoint.py | 72 +- 26 files changed, 2195 insertions(+), 21 deletions(-) create mode 100644 gateway/launch.py create mode 100644 helm/litellm-helm/tests/connection_pool_tests.yaml create mode 100644 helm/litellm/tests/connection_pool_tests.yaml create mode 100644 litellm/proxy/db/pgbouncer.py create mode 100644 terraform/litellm/aws/tests/connection_pool.tftest.hcl create mode 100644 terraform/litellm/gcp/tests/connection_pool.tftest.hcl create mode 100644 tests/test_gateway/test_launch.py create mode 100644 tests/test_litellm/proxy/db/test_pgbouncer.py diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index cc606339a20..f55c87c2ae5 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -200,6 +200,7 @@ jobs: tests/test_litellm/proxy/types_utils tests/test_litellm/proxy/logging_endpoints tests/test_litellm/proxy/test_*.py + tests/test_gateway workers: 4 reruns: 2 timeout-minutes: 20 diff --git a/Dockerfile b/Dockerfile index 0a92aa9a68c..759dac76795 100644 --- a/Dockerfile +++ b/Dockerfile @@ -8,9 +8,25 @@ ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7 ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a # Pinned by digest like the other base images; bump explicitly on Node upgrades. ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43 +# Checksum from https://www.pgbouncer.org/downloads/ (the Wolfi repo only carries 1.24.x) +ARG PGBOUNCER_VERSION=1.25.2 +ARG PGBOUNCER_SHA256=924ad35113fd0a71c8e2dbe85b5d03445532e2b7b37a9f8a48983beea238b332 FROM $UV_IMAGE AS uvbin +FROM $LITELLM_BUILD_IMAGE AS pgbouncer-builder +ARG PGBOUNCER_VERSION +ARG PGBOUNCER_SHA256 +USER root +RUN apk add --no-cache build-base pkgconf libevent-dev openssl-dev curl +WORKDIR /build +RUN curl -fsSL -o pgbouncer.tar.gz "https://www.pgbouncer.org/downloads/files/${PGBOUNCER_VERSION}/pgbouncer-${PGBOUNCER_VERSION}.tar.gz" && \ + echo "${PGBOUNCER_SHA256} pgbouncer.tar.gz" | sha256sum -c - && \ + tar xzf pgbouncer.tar.gz --strip-components=1 && \ + ./configure --prefix=/usr/local --with-openssl=/usr && \ + make -j"$(nproc)" pgbouncer && \ + install -m 0755 pgbouncer /usr/local/bin/pgbouncer + # Admin UI builder. Pinned to the build platform so the architecture-independent # Next.js static export compiles once natively even in a multi-arch build, # instead of once per target arch under QEMU. @@ -110,7 +126,8 @@ USER root RUN echo "https://packages.wolfi.dev/os" >> /etc/apk/repositories # node (without npm) is required by the prisma CLI at runtime -RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile +RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile libevent +COPY --from=pgbouncer-builder /usr/local/bin/pgbouncer /usr/local/bin/pgbouncer WORKDIR /app ENV PATH="/app/.venv/bin:${PATH}" \ diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database index e9ad2849bb2..b0bf935c616 100644 --- a/docker/Dockerfile.database +++ b/docker/Dockerfile.database @@ -8,9 +8,25 @@ ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7 ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a # Pinned by digest like the other base images; bump explicitly on Node upgrades. ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43 +# Checksum from https://www.pgbouncer.org/downloads/ (the Wolfi repo only carries 1.24.x) +ARG PGBOUNCER_VERSION=1.25.2 +ARG PGBOUNCER_SHA256=924ad35113fd0a71c8e2dbe85b5d03445532e2b7b37a9f8a48983beea238b332 FROM $UV_IMAGE AS uvbin +FROM $LITELLM_BUILD_IMAGE AS pgbouncer-builder +ARG PGBOUNCER_VERSION +ARG PGBOUNCER_SHA256 +USER root +RUN apk add --no-cache build-base pkgconf libevent-dev openssl-dev curl +WORKDIR /build +RUN curl -fsSL -o pgbouncer.tar.gz "https://www.pgbouncer.org/downloads/files/${PGBOUNCER_VERSION}/pgbouncer-${PGBOUNCER_VERSION}.tar.gz" && \ + echo "${PGBOUNCER_SHA256} pgbouncer.tar.gz" | sha256sum -c - && \ + tar xzf pgbouncer.tar.gz --strip-components=1 && \ + ./configure --prefix=/usr/local --with-openssl=/usr && \ + make -j"$(nproc)" pgbouncer && \ + install -m 0755 pgbouncer /usr/local/bin/pgbouncer + # Admin UI builder. Pinned to the build platform so the architecture-independent # Next.js static export compiles once natively even in a multi-arch build, # instead of once per target arch under QEMU. @@ -101,7 +117,8 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime USER root # node (without npm) is required by the prisma CLI at runtime -RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile +RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile libevent +COPY --from=pgbouncer-builder /usr/local/bin/pgbouncer /usr/local/bin/pgbouncer WORKDIR /app ENV PATH="/app/.venv/bin:${PATH}" \ diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root index edf20e8bbff..5d729046678 100644 --- a/docker/Dockerfile.non_root +++ b/docker/Dockerfile.non_root @@ -7,9 +7,25 @@ ARG PROXY_EXTRAS_SOURCE=published ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a # Pinned by digest like the other base images; bump explicitly on Node upgrades. ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43 +# Checksum from https://www.pgbouncer.org/downloads/ (the Wolfi repo only carries 1.24.x) +ARG PGBOUNCER_VERSION=1.25.2 +ARG PGBOUNCER_SHA256=924ad35113fd0a71c8e2dbe85b5d03445532e2b7b37a9f8a48983beea238b332 FROM $UV_IMAGE AS uvbin +FROM $LITELLM_BUILD_IMAGE AS pgbouncer-builder +ARG PGBOUNCER_VERSION +ARG PGBOUNCER_SHA256 +USER root +RUN apk add --no-cache build-base pkgconf libevent-dev openssl-dev curl +WORKDIR /build +RUN curl -fsSL -o pgbouncer.tar.gz "https://www.pgbouncer.org/downloads/files/${PGBOUNCER_VERSION}/pgbouncer-${PGBOUNCER_VERSION}.tar.gz" && \ + echo "${PGBOUNCER_SHA256} pgbouncer.tar.gz" | sha256sum -c - && \ + tar xzf pgbouncer.tar.gz --strip-components=1 && \ + ./configure --prefix=/usr/local --with-openssl=/usr && \ + make -j"$(nproc)" pgbouncer && \ + install -m 0755 pgbouncer /usr/local/bin/pgbouncer + # Admin UI builder. Pinned to the build platform so the architecture-independent # Next.js static export compiles once natively even in a multi-arch build, # instead of once per target arch under QEMU. @@ -128,8 +144,9 @@ RUN for i in 1 2 3; do \ apk upgrade --no-cache && break || sleep 5; \ done && \ for i in 1 2 3; do \ - apk add --no-cache python-3.13 bash openssl tzdata libsndfile nodejs && break || sleep 5; \ + apk add --no-cache python-3.13 bash openssl tzdata libsndfile nodejs libevent && break || sleep 5; \ done +COPY --from=pgbouncer-builder /usr/local/bin/pgbouncer /usr/local/bin/pgbouncer # Copy only what runtime needs. The application is installed inside the venv; # the rest of the builder's /app is source and build metadata that must not diff --git a/gateway/Dockerfile b/gateway/Dockerfile index 308d70a6b26..33d3791dbba 100644 --- a/gateway/Dockerfile +++ b/gateway/Dockerfile @@ -1,9 +1,25 @@ ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a +# Checksum from https://www.pgbouncer.org/downloads/ (the Wolfi repo only carries 1.24.x) +ARG PGBOUNCER_VERSION=1.25.2 +ARG PGBOUNCER_SHA256=924ad35113fd0a71c8e2dbe85b5d03445532e2b7b37a9f8a48983beea238b332 FROM $UV_IMAGE AS uvbin +FROM $LITELLM_BUILD_IMAGE AS pgbouncer-builder +ARG PGBOUNCER_VERSION +ARG PGBOUNCER_SHA256 +USER root +RUN apk add --no-cache build-base pkgconf libevent-dev openssl-dev curl +WORKDIR /build +RUN curl -fsSL -o pgbouncer.tar.gz "https://www.pgbouncer.org/downloads/files/${PGBOUNCER_VERSION}/pgbouncer-${PGBOUNCER_VERSION}.tar.gz" && \ + echo "${PGBOUNCER_SHA256} pgbouncer.tar.gz" | sha256sum -c - && \ + tar xzf pgbouncer.tar.gz --strip-components=1 && \ + ./configure --prefix=/usr/local --with-openssl=/usr && \ + make -j"$(nproc)" pgbouncer && \ + install -m 0755 pgbouncer /usr/local/bin/pgbouncer + # ---------- Builder ---------- FROM $LITELLM_BUILD_IMAGE AS builder @@ -61,6 +77,10 @@ RUN --mount=type=cache,target=/root/.cache/uv \ --extra bedrock-realtime \ --python python3.13 +# PYTHONPATH=/app makes the source tree shadow the installed package, so the +# compiled Rust extension must live next to the source or it is never imported. +RUN cp "$(python -c 'import sysconfig; print(sysconfig.get_paths()["purelib"])')"/litellm/rust_bridge/_native*.so litellm/rust_bridge/ + RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \ npm_config_cache=/root/.npm \ prisma generate --schema=./schema.prisma @@ -73,7 +93,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime USER root RUN for i in 1 2 3; do \ - apk add --no-cache bash openssl tzdata python-3.13 libsndfile libatomic && break; \ + apk add --no-cache bash openssl tzdata python-3.13 libsndfile libatomic libevent && break; \ [ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \ sleep 5; \ done @@ -90,15 +110,17 @@ ENV HOME=/home/nonroot \ COPY --from=builder --chown=nonroot:nonroot /app /app COPY --from=builder /opt/prisma /opt/prisma +COPY --from=pgbouncer-builder /usr/local/bin/pgbouncer /usr/local/bin/pgbouncer RUN find /app/.venv -type f -path "*/tornado/test/*" -delete && \ find /app/.venv -type d -path "*/tornado/test" -delete && \ chmod -R a+rX /opt/prisma && \ - python -c "from prisma.client import BINARY_PATHS; paths = list(BINARY_PATHS.query_engine.values()); assert paths and all(p.startswith('/opt/prisma/') for p in paths), paths" + python -c "from prisma.client import BINARY_PATHS; paths = list(BINARY_PATHS.query_engine.values()); assert paths and all(p.startswith('/opt/prisma/') for p in paths), paths" && \ + python -c "import litellm; from litellm.rust_bridge.loader import native_bridge_available; assert litellm.__file__ == '/app/litellm/__init__.py', litellm.__file__; assert native_bridge_available()" USER nonroot EXPOSE 4000/tcp -ENTRYPOINT ["sh", "-c", "exec /app/docker/component_entrypoint.sh uvicorn gateway.main:app --workers \"${NUM_WORKERS:-1}\" \"$@\"", "--"] +ENTRYPOINT ["sh", "-c", "exec /app/docker/component_entrypoint.sh python -m gateway.launch --workers \"${NUM_WORKERS:-1}\" \"$@\"", "--"] CMD ["--host", "0.0.0.0", "--port", "4000"] diff --git a/gateway/launch.py b/gateway/launch.py new file mode 100644 index 00000000000..6b28fafdcf6 --- /dev/null +++ b/gateway/launch.py @@ -0,0 +1,71 @@ +"""Gateway supervisor: assemble DATABASE_URL, start the in-container PgBouncer, then run uvicorn. + +``gateway/main.py`` assembles ``DATABASE_URL`` inside every uvicorn worker, which +is fine for a plain Postgres URL but not for the pooler: PgBouncer must be +started exactly once per pod, before the workers fork, and the workers must be +handed the loopback URL it listens on. A pre-existing ``DATABASE_URL`` wins in +``DatabaseURLSettings.apply_to_env`` under password auth, so setting it here is +enough for every worker to pick the pooled URL up unchanged. + +Run with: + python -m gateway.launch --workers 4 --host 0.0.0.0 --port 4000 +""" + +import os +import sys +from collections.abc import Callable, Mapping, Sequence +from typing import Final + +from uvicorn.main import main as uvicorn_main + +from litellm.proxy.db.db_url_settings import DatabaseURLSettings +from litellm.proxy.db.pgbouncer import PgBouncerError, PgBouncerSettings, start_in_container_pgbouncer + +GATEWAY_APP: Final = "gateway.main:app" +KEEPALIVE_FLAG: Final = "--timeout-keep-alive" + + +def uvicorn_argv(argv: Sequence[str], environ: Mapping[str, str]) -> tuple[str, ...]: + """Honor ``KEEPALIVE_TIMEOUT`` like ``proxy_cli.py`` does, unless the flag was passed explicitly.""" + keepalive: Final = environ.get("KEEPALIVE_TIMEOUT") + if keepalive is None or any(arg == KEEPALIVE_FLAG or arg.startswith(f"{KEEPALIVE_FLAG}=") for arg in argv): + return (GATEWAY_APP, *argv) + return (GATEWAY_APP, *argv, KEEPALIVE_FLAG, keepalive) + + +def pool_database_url( + settings: DatabaseURLSettings, + pgbouncer: PgBouncerSettings, + environ: Mapping[str, str], +) -> str | PgBouncerError | None: + """Start the in-container PgBouncer and return its loopback URL, or None when ``pgbouncer.enabled`` is off. + + The upstream URL is whatever ``apply_to_env`` assembled from the discrete + ``DATABASE_*`` vars (or an operator-pinned ``DATABASE_URL``). Token auth is + rejected by the pooler itself, since it holds one password for its lifetime. + """ + if not pgbouncer.enabled: + return None + upstream_url: Final = environ.get("DATABASE_URL") + if upstream_url is None: + return PgBouncerError("LITELLM_PGBOUNCER_ENABLED is set but no DATABASE_URL could be assembled") + return start_in_container_pgbouncer(pgbouncer, upstream_url, token_auth_enabled=settings.token_auth() is not None) + + +def _serve(argv: Sequence[str]) -> None: + uvicorn_main(tuple(argv), prog_name="uvicorn") + + +def main(argv: Sequence[str], serve: Callable[[Sequence[str]], None] = _serve) -> None: + settings: Final = DatabaseURLSettings.from_env() + settings.apply_to_env() + pooled_url: Final = pool_database_url(settings, PgBouncerSettings(), os.environ) + if isinstance(pooled_url, PgBouncerError): + sys.exit(f"LiteLLM gateway: in-container pgbouncer could not start: {pooled_url.reason}") + if pooled_url is not None: + os.environ["DATABASE_URL"] = pooled_url + serve(uvicorn_argv(argv, os.environ)) + + +if __name__ == "__main__": + main(sys.argv[1:]) diff --git a/helm/litellm-helm/templates/deployment.yaml b/helm/litellm-helm/templates/deployment.yaml index f7c918a6827..834071eb9b2 100644 --- a/helm/litellm-helm/templates/deployment.yaml +++ b/helm/litellm-helm/templates/deployment.yaml @@ -117,6 +117,14 @@ spec: - name: DATABASE_URL_READ_REPLICA value: {{ .Values.db.readReplicaUrl | quote }} {{- end }} + {{- if .Values.db.connectionPool.enabled }} + - name: LITELLM_PGBOUNCER_ENABLED + value: "true" + - name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS + value: {{ .Values.db.connectionPool.maxDbConnections | quote }} + - name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN + value: {{ .Values.db.connectionPool.maxClientConn | quote }} + {{- end }} - name: PROXY_MASTER_KEY valueFrom: secretKeyRef: diff --git a/helm/litellm-helm/tests/connection_pool_tests.yaml b/helm/litellm-helm/tests/connection_pool_tests.yaml new file mode 100644 index 00000000000..af23512dafc --- /dev/null +++ b/helm/litellm-helm/tests/connection_pool_tests.yaml @@ -0,0 +1,61 @@ +suite: test in-container connection pool +templates: + - deployment.yaml + - configmap-litellm.yaml +tests: + - it: should not emit pgbouncer env vars by default + template: deployment.yaml + asserts: + - notContains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_ENABLED + value: "true" + - notContains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS + value: "20" + + - it: should enable the pool with the default sizing when connectionPool.enabled is set + template: deployment.yaml + set: + db.connectionPool.enabled: true + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_ENABLED + value: "true" + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS + value: "20" + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN + value: "1000" + + - it: should pass custom sizing through as strings next to the worker count + template: deployment.yaml + set: + numWorkers: 4 + db.connectionPool.enabled: true + db.connectionPool.maxDbConnections: 8 + db.connectionPool.maxClientConn: 400 + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS + value: "8" + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN + value: "400" + - contains: + path: spec.template.spec.containers[0].args + content: "4" diff --git a/helm/litellm-helm/values.yaml b/helm/litellm-helm/values.yaml index 907f85de8fc..db596e2d68e 100644 --- a/helm/litellm-helm/values.yaml +++ b/helm/litellm-helm/values.yaml @@ -355,6 +355,20 @@ db: # only (e.g. when IAM_TOKEN_DB_AUTH supplies the token at runtime). readReplicaUrl: "" + # In-container connection pool (PgBouncer, transaction mode) shared by every + # worker in the pod. Without it each --num_workers worker opens its own + # connection_limit connections to Postgres, so a pod's footprint against the + # database's connection ceiling is workers x connection_limit and grows with + # every replica. With it, the pod holds at most maxDbConnections upstream + # connections no matter how many workers run; the workers connect to the pool + # over loopback, with no extra network hop. Migrations still go straight to + # Postgres. Starting profile for numWorkers: 4 is maxDbConnections: 20, so + # a database with a 5000-connection ceiling fits roughly 200 replicas. + connectionPool: + enabled: false + maxDbConnections: 20 + maxClientConn: 1000 + # Use the Stackgres Helm chart to deploy an instance of a Stackgres cluster. # The Stackgres Operator must already be installed within the target # Kubernetes cluster. diff --git a/helm/litellm/templates/_helpers.tpl b/helm/litellm/templates/_helpers.tpl index c459512c7b9..4ad3cfe0484 100644 --- a/helm/litellm/templates/_helpers.tpl +++ b/helm/litellm/templates/_helpers.tpl @@ -360,6 +360,23 @@ harmless no-op for the Job and authoritative for the app pods. {{- end }} {{- end -}} +{{/* +In-container PgBouncer env for the gateway container. Fails at render time under IAM or Entra auth: the pooler holds one static password for the life of the pod. +*/}} +{{- define "litellm.connectionPoolEnv" -}} +{{- if or .Values.database.writer.useIAMAuth .Values.database.writer.useAzureEntraAuth }} +{{- fail "database.connectionPool.enabled cannot be combined with database.writer.useIAMAuth or database.writer.useAzureEntraAuth: the in-container pgbouncer holds a static database password and cannot follow a rotating token. Disable the pool or use a static database password" }} +{{- end }} +{{- with .Values.database.connectionPool -}} +- name: LITELLM_PGBOUNCER_ENABLED + value: "true" +- name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS + value: {{ required "database.connectionPool.maxDbConnections is required when the pool is enabled" .maxDbConnections | quote }} +- name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN + value: {{ required "database.connectionPool.maxClientConn is required when the pool is enabled" .maxClientConn | quote }} +{{- end }} +{{- end -}} + {{/* PodDisruptionBudget shared by gateway, backend, and ui. diff --git a/helm/litellm/templates/gateway/deployment.yaml b/helm/litellm/templates/gateway/deployment.yaml index 9cb6b07e77b..1e9041f0a33 100644 --- a/helm/litellm/templates/gateway/deployment.yaml +++ b/helm/litellm/templates/gateway/deployment.yaml @@ -61,6 +61,9 @@ spec: - name: NUM_WORKERS value: {{ .Values.gateway.numWorkers | quote }} {{- end }} + {{- if .Values.database.connectionPool.enabled }} + {{- include "litellm.connectionPoolEnv" $ | nindent 12 }} + {{- end }} {{- if .Values.billingMetrics.enabled }} {{- include "litellm.billingMetricsEnv" . | nindent 12 }} {{- end }} diff --git a/helm/litellm/tests/connection_pool_tests.yaml b/helm/litellm/tests/connection_pool_tests.yaml new file mode 100644 index 00000000000..31be7575f3c --- /dev/null +++ b/helm/litellm/tests/connection_pool_tests.yaml @@ -0,0 +1,117 @@ +suite: test in-container connection pool env vars +templates: + - gateway/deployment.yaml + - gateway/configmap.yaml + - backend/deployment.yaml + - backend/configmap.yaml +values: + - ./values/required.yaml +tests: + - it: renders no pool env by default + template: gateway/deployment.yaml + asserts: + - notContains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_ENABLED + any: true + - notContains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS + any: true + - notContains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN + any: true + + - it: enabled pool renders the three pgbouncer vars with the configured sizes + template: gateway/deployment.yaml + set: + gateway.numWorkers: 4 + database.connectionPool.enabled: true + database.connectionPool.maxDbConnections: 8 + database.connectionPool.maxClientConn: 250 + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_ENABLED + value: "true" + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS + value: "8" + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN + value: "250" + - contains: + path: spec.template.spec.containers[0].env + content: + name: NUM_WORKERS + value: "4" + + - it: enabled pool uses the chart default sizes + template: gateway/deployment.yaml + set: + database.connectionPool.enabled: true + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS + value: "20" + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN + value: "1000" + + - it: backend never gets the pool env + template: backend/deployment.yaml + set: + database.connectionPool.enabled: true + asserts: + - notContains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_ENABLED + any: true + + - it: pool with IAM auth fails at render time + template: gateway/deployment.yaml + set: + database.connectionPool.enabled: true + database.writer.useIAMAuth: true + asserts: + - failedTemplate: + errorMessage: "database.connectionPool.enabled cannot be combined with database.writer.useIAMAuth or database.writer.useAzureEntraAuth: the in-container pgbouncer holds a static database password and cannot follow a rotating token. Disable the pool or use a static database password" + + - it: pool with Entra auth fails at render time + template: gateway/deployment.yaml + set: + database.connectionPool.enabled: true + database.writer.useAzureEntraAuth: true + asserts: + - failedTemplate: + errorMessage: "database.connectionPool.enabled cannot be combined with database.writer.useIAMAuth or database.writer.useAzureEntraAuth: the in-container pgbouncer holds a static database password and cannot follow a rotating token. Disable the pool or use a static database password" + + - it: IAM auth without the pool still renders + template: gateway/deployment.yaml + set: + database.writer.useIAMAuth: true + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: IAM_TOKEN_DB_AUTH + value: "true" + - notContains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_PGBOUNCER_ENABLED + any: true diff --git a/helm/litellm/values.yaml b/helm/litellm/values.yaml index be1e6d43987..45ab5229f38 100644 --- a/helm/litellm/values.yaml +++ b/helm/litellm/values.yaml @@ -225,6 +225,26 @@ database: usernameKey: username passwordKey: password + # In-container connection pool (PgBouncer, transaction mode) shared by every + # gateway worker in the pod. Without it each of the `gateway.numWorkers` + # workers opens its own Prisma pool straight to Postgres, so a pod's + # footprint against the database's connection ceiling is + # numWorkers x connection_limit and grows with every replica. With it, the + # pod holds at most maxDbConnections upstream connections no matter how many + # workers run; the workers connect to the pool over loopback, with no extra + # network hop. The chart emits LITELLM_PGBOUNCER_ENABLED / + # LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS / LITELLM_PGBOUNCER_MAX_CLIENT_CONN on + # the gateway container only: the backend runs a single worker and the + # migrations Job must keep a direct connection. The pool holds a static + # password, so it cannot be combined with `database.writer.useIAMAuth` or + # `useAzureEntraAuth` (rendering fails). Starting profile for + # `gateway.numWorkers: 4` is maxDbConnections: 20, so a database with a + # 5000-connection ceiling fits roughly 200 gateway replicas. + connectionPool: + enabled: false + maxDbConnections: 20 + maxClientConn: 1000 + # Optional Redis. Leave host empty to disable. # # This is the proxy's coordination store: cross-pod tpm/rpm rate limits, spend diff --git a/litellm/proxy/db/pgbouncer.py b/litellm/proxy/db/pgbouncer.py new file mode 100644 index 00000000000..7792867600c --- /dev/null +++ b/litellm/proxy/db/pgbouncer.py @@ -0,0 +1,525 @@ +"""In-container PgBouncer shared by every proxy worker. + +Each uvicorn worker owns a Prisma query engine with its own pool of +``connection_limit`` server connections, so the connections a pod holds open +against Postgres scale as ``workers * connection_limit`` and a database with a +fixed connection ceiling runs out of room as pods and workers are added. + +When ``LITELLM_PGBOUNCER_ENABLED`` is set, the supervisor process starts one +PgBouncer next to the workers (no extra network hop: it listens on loopback +inside the pod) in transaction pooling mode, points ``DATABASE_URL`` at it +with ``pgbouncer=true`` so Prisma stops using server-side prepared statements, +and keeps it running for the life of the proxy. Every worker's pool then +becomes cheap client connections to PgBouncer while the upstream connection +count is capped at ``LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS`` per pod, no matter +how many workers run. + +Migrations and the schema diff run in the supervisor before the pooler is +started, so they always go straight to Postgres. ``DATABASE_URL_READ_REPLICA`` +is left untouched. + +The pooler holds the database password from startup, so it cannot be combined +with ``IAM_TOKEN_DB_AUTH`` or ``AZURE_POSTGRESQL_AUTH``: those rotate the +password inside every worker on their own schedule, and PgBouncer would keep +authenticating upstream with the expired token. +""" + +from __future__ import annotations + +import atexit +import os +import re +import shlex +import shutil +import socket +import subprocess +import tempfile +import threading +import time +import urllib.parse +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from pathlib import Path +from types import MappingProxyType +from typing import Final + +from pydantic import Field +from pydantic_settings import BaseSettings, SettingsConfigDict + +from litellm._logging import verbose_proxy_logger +from litellm.proxy.db.token_auth import AZURE_POSTGRESQL_AUTH_ENV_VAR, IAM_TOKEN_DB_AUTH_ENV_VAR + +PGBOUNCER_ENV_PREFIX: Final = "LITELLM_PGBOUNCER_" +PGBOUNCER_LISTEN_ADDR: Final = "127.0.0.1" +PGBOUNCER_INI_NAME: Final = "pgbouncer.ini" +PGBOUNCER_USERLIST_NAME: Final = "userlist.txt" +PGBOUNCER_CA_NAME: Final = "server-ca.pem" +PGBOUNCER_RESTART_DELAY_SECONDS: Final = 1.0 +PGBOUNCER_READY_TIMEOUT_SECONDS: Final = 15.0 +PGBOUNCER_STOP_GRACE_SECONDS: Final = 10.0 +PGBOUNCER_UNPRIVILEGED_USER: Final = "nobody" +PGBOUNCER_MIN_VERSION: Final = (1, 19) +PGBOUNCER_VERSION_PATTERN: Final = re.compile(r"PgBouncer (\d+)\.(\d+)") +PGBOUNCER_LIST_DELIMITER_PATTERN: Final = re.compile(r"[,\s]") +PGBOUNCER_TOKEN_AUTH_CONFLICT: Final = ( + f"the in-container pgbouncer cannot be combined with {IAM_TOKEN_DB_AUTH_ENV_VAR} or " + f"{AZURE_POSTGRESQL_AUTH_ENV_VAR}: each worker rotates the database password on its own schedule and the pooler " + "would keep using the expired token upstream. Disable the pooler or use a static database password" +) + +# Prisma's client-side TLS params describe the hop to Postgres, which becomes +# PgBouncer's server side. They move into ``server_tls_*`` and must not stay on +# the loopback URL: the listener speaks plain TCP and Prisma would refuse it +# under ``sslmode=require`` or ``channel_binding=require``. +PRISMA_TLS_PARAM_KEYS: Final[frozenset[str]] = frozenset( + {"sslmode", "sslcert", "sslaccept", "sslidentity", "sslpassword", "channel_binding", "gssencmode"} +) +POOLED_URL_DROPPED_KEYS: Final[frozenset[str]] = PRISMA_TLS_PARAM_KEYS | frozenset(("options", "pgbouncer")) +PGBOUNCER_SSLMODES: Final[frozenset[str]] = frozenset( + {"disable", "allow", "prefer", "require", "verify-ca", "verify-full"} +) + + +class PgBouncerSettings(BaseSettings): + """``LITELLM_PGBOUNCER_*`` env vars, read once in the supervisor.""" + + model_config = SettingsConfigDict( + env_prefix=PGBOUNCER_ENV_PREFIX, case_sensitive=False, extra="ignore", frozen=True + ) + + enabled: bool = False + port: int = Field(default=6432, ge=1, le=65535) + max_db_connections: int = Field(default=20, ge=1) + max_client_conn: int = Field(default=1000, ge=1) + binary: str = "pgbouncer" + + +@dataclass(frozen=True, slots=True) +class PgBouncerPlan: + ini: str + userlist: str + pooled_url: str + ca_source: str | None = None + + +@dataclass(frozen=True, slots=True) +class PgBouncerError: + reason: str + + +def _single_quoted(value: str) -> str: + """Quote for SQL and for PgBouncer's ``[databases]`` connection string: both double a literal ``'``.""" + return "'" + value.replace("'", "''") + "'" + + +def _userlist_quote(value: str) -> str: + return '"' + value.replace('"', '""') + '"' + + +def _option_settings(tokens: Sequence[str]) -> tuple[str, ...] | None: + """The ``name=value`` settings in a libpq ``options`` string, or None if it holds anything else. + + Accepts ``-c name=value``, ``-cname=value`` and ``--name=value``; a + detached ``-c`` is folded into the token that follows it first. + """ + folded: Final = tuple( + f"-c{tokens[index + 1]}" if token == "-c" and index + 1 < len(tokens) else token + for index, token in enumerate(tokens) + if index == 0 or tokens[index - 1] != "-c" + ) + settings: Final = tuple(token[2:] for token in folded if token.startswith(("-c", "--")) and "=" in token[2:]) + return settings if len(settings) == len(folded) else None + + +def _connect_query(options: str) -> str | PgBouncerError: + """Turn Prisma's ``options=-c name=value ...`` startup param into ``SET`` statements. + + PgBouncer rejects any ``-c`` setting in ``options`` that is not one of the + handful it tracks (``statement_timeout`` and ``lock_timeout`` are not), so + the settings are applied to each new server connection instead. Every + client shares them, which is what the single ``DATABASE_URL`` gave anyway. + """ + settings: Final = _option_settings(tuple(shlex.split(options))) + if settings is None: + return PgBouncerError(f"cannot translate the DATABASE_URL options {options!r} into PgBouncer settings") + return "; ".join( + f"SET {name.strip()} TO {_single_quoted(value.strip())}" + for name, value in (setting.split("=", 1) for setting in settings) + ) + + +def _server_tls_settings(sslmode: str, sslcert: str, sslaccept: str, ca_path: Path) -> tuple[str, ...] | PgBouncerError: + """``server_tls_*`` lines naming ``ca_path``, the runtime-dir copy of the bundle: the original (or the + 0600 root pinned by ``pin_bundle_root``) is often unreadable for the user PgBouncer drops to.""" + if sslmode not in PGBOUNCER_SSLMODES: + return PgBouncerError(f"unsupported sslmode {sslmode!r} on DATABASE_URL") + verify: Final = sslmode in ("verify-ca", "verify-full") or (sslmode == "require" and sslaccept == "strict") + if verify and not sslcert: + return PgBouncerError( + "DATABASE_URL asks for a verified TLS connection but names no CA bundle; " + "add sslcert= (or sslrootcert=) so the in-container PgBouncer can verify Postgres" + ) + mode: Final = "verify-full" if verify else sslmode + return (f"server_tls_sslmode = {mode}", *((f"server_tls_ca_file = {ca_path}",) if sslcert else ())) + + +def plan_pgbouncer( + upstream_url: str, + settings: PgBouncerSettings, + runtime_dir: Path, + run_as_user: str | None, +) -> PgBouncerPlan | PgBouncerError: + """Render the PgBouncer config for ``upstream_url`` and the loopback URL Prisma uses instead. + + Params describing Prisma's own pool (``connection_limit``, ``pool_timeout``, + ...) stay on the pooled URL; the TLS params and ``options`` describe the hop + to Postgres and move into the PgBouncer config. ``run_as_user`` is the + unprivileged user PgBouncer drops to when the proxy runs as root, which + PgBouncer itself refuses to do. + """ + parsed: Final = urllib.parse.urlsplit(upstream_url) + params: Final[Mapping[str, str]] = MappingProxyType( + dict(urllib.parse.parse_qsl(parsed.query, keep_blank_values=True)) + ) + dbname: Final = urllib.parse.unquote(parsed.path.lstrip("/")) + username: Final = urllib.parse.unquote(parsed.username or "") + password: Final = None if parsed.password is None else urllib.parse.unquote(parsed.password) + if not parsed.hostname or not username or password is None or not dbname: + return PgBouncerError( + "DATABASE_URL must carry a host, user, password and database name for the in-container PgBouncer" + ) + if PGBOUNCER_LIST_DELIMITER_PATTERN.search(username): + return PgBouncerError( + f"the database user {username!r} cannot be named in PgBouncer's stats_users list: " + "PgBouncer splits list settings on commas and whitespace and has no quoting for them" + ) + if "sslidentity" in params: + return PgBouncerError("client certificates (sslidentity) are not supported with the in-container PgBouncer") + tls: Final = _server_tls_settings( + params.get("sslmode", "prefer"), + params.get("sslcert", ""), + params.get("sslaccept", ""), + runtime_dir / PGBOUNCER_CA_NAME, + ) + if isinstance(tls, PgBouncerError): + return tls + connect_query: Final = _connect_query(params["options"]) if params.get("options") else "" + if isinstance(connect_query, PgBouncerError): + return connect_query + upstream: Final = " ".join( + ( + f"host={_single_quoted(parsed.hostname)}", + f"port={parsed.port or 5432}", + f"dbname={_single_quoted(dbname)}", + f"user={_single_quoted(username)}", + f"password={_single_quoted(password)}", + *((f"connect_query={_single_quoted(connect_query)}",) if connect_query else ()), + ) + ) + ini: Final = "\n".join( + ( + "[databases]", + f"{dbname} = {upstream}", + "", + "[pgbouncer]", + f"listen_addr = {PGBOUNCER_LISTEN_ADDR}", + f"listen_port = {settings.port}", + f"unix_socket_dir = {runtime_dir}", + f"auth_file = {runtime_dir / PGBOUNCER_USERLIST_NAME}", + "auth_type = scram-sha-256", + f"stats_users = {username}", + "pool_mode = transaction", + f"max_client_conn = {settings.max_client_conn}", + f"default_pool_size = {settings.max_db_connections}", + f"max_db_connections = {settings.max_db_connections}", + "ignore_startup_parameters = extra_float_digits", + *tls, + *((f"user = {run_as_user}",) if run_as_user else ()), + "", + ) + ) + userlist: Final = f"{_userlist_quote(username)} {_userlist_quote(password)}\n" + pooled_query: Final = urllib.parse.urlencode( + (*((key, value) for key, value in params.items() if key not in POOLED_URL_DROPPED_KEYS), ("pgbouncer", "true")) + ) + credentials: Final = f"{urllib.parse.quote(username, safe='')}:{urllib.parse.quote(password, safe='')}" + pooled_url: Final = urllib.parse.urlunsplit( + parsed._replace(netloc=f"{credentials}@{PGBOUNCER_LISTEN_ADDR}:{settings.port}", query=pooled_query) + ) + return PgBouncerPlan(ini=ini, userlist=userlist, pooled_url=pooled_url, ca_source=params.get("sslcert") or None) + + +def write_pgbouncer_files(plan: PgBouncerPlan, runtime_dir: Path, run_as_user: str | None) -> Path | PgBouncerError: + """Write the ini, userlist (both hold the password, so mode 0600) and CA copy, and return the ini path. + + ``run_as_user`` is the user PgBouncer drops to when started as root; it has + to own the files it re-reads on reload and the socket directory. + """ + ini_path: Final = runtime_dir / PGBOUNCER_INI_NAME + userlist_path: Final = runtime_dir / PGBOUNCER_USERLIST_NAME + ca_path: Final = runtime_dir / PGBOUNCER_CA_NAME + if plan.ca_source is not None: + try: + shutil.copyfile(plan.ca_source, ca_path) + except OSError as error: + return PgBouncerError(f"cannot read the CA bundle {plan.ca_source!r} named by sslcert: {error}") + for path, content in ((userlist_path, plan.userlist), (ini_path, plan.ini)): + path.touch(mode=0o600) + path.write_text(content, encoding="utf-8") + if run_as_user is not None: + runtime_dir.chmod(0o700) + for path in (runtime_dir, ini_path, userlist_path, *((ca_path,) if plan.ca_source is not None else ())): + shutil.chown(path, user=run_as_user) + return ini_path + + +def _port_open(port: int) -> bool: + try: + with socket.create_connection((PGBOUNCER_LISTEN_ADDR, port), timeout=0.5): + return True + except OSError: + return False + + +def _unix_socket_open(path: Path) -> bool: + with socket.socket(socket.AF_UNIX) as probe: + probe.settimeout(0.5) + try: + probe.connect(str(path)) + except OSError: + return False + return True + + +def unix_socket_path(runtime_dir: Path, port: int) -> Path: + return runtime_dir / f".s.PGSQL.{port}" + + +def pgbouncer_version(binary: str) -> tuple[int, int] | PgBouncerError: + """``(major, minor)`` from `` --version``. + + Readiness relies on PgBouncer exiting when it cannot bind its TCP port, + which it does from 1.19 on. Older releases log a warning and serve the unix + socket alone, so their socket would vouch for a port held by someone else. + """ + try: + output: Final = subprocess.run( + (binary, "--version"), capture_output=True, text=True, check=False, timeout=10 + ).stdout + except (OSError, subprocess.TimeoutExpired) as run_error: + return PgBouncerError(f"could not run {binary!r} --version: {run_error}") + found: Final = PGBOUNCER_VERSION_PATTERN.search(output) + if found is None: + return PgBouncerError(f"{binary!r} --version did not report a PgBouncer version: {output.strip()!r}") + return int(found[1]), int(found[2]) + + +def _end(process: subprocess.Popen[bytes]) -> None: + if process.poll() is not None: + return + process.terminate() + try: + process.wait(timeout=PGBOUNCER_STOP_GRACE_SECONDS) + except subprocess.TimeoutExpired: + process.kill() + process.wait() + + +class PgBouncerProcess: + """Runs ``argv`` as a foreground child and restarts it whenever it exits on its own. + + Prisma reconnects by itself after a failed query, so a PgBouncer crash + costs the requests in flight plus one failed query per idle pooled + connection the crash severed, and nothing else once the replacement is + listening again. A replacement that cannot be spawned, finds its port + taken, exits again or never starts listening is retried every + ``restart_delay_seconds`` until ``stop`` is called. + + A connect probe of ``port`` cannot tell the child from another process + that grabbed the port after the availability check, so readiness also + needs ``socket_path``: the unix socket PgBouncer creates in the private + runtime directory, which it only does once every TCP listener is bound + (PgBouncer 1.19 or newer, see ``pgbouncer_version``). + """ + + def __init__( + self, + argv: Sequence[str], + port: int, + socket_path: Path, + restart_delay_seconds: float = PGBOUNCER_RESTART_DELAY_SECONDS, + ready_timeout_seconds: float = PGBOUNCER_READY_TIMEOUT_SECONDS, + ) -> None: + self.argv: Final = tuple(argv) + self.port: Final = port + self.socket_path: Final = socket_path + self.restart_delay_seconds: Final = restart_delay_seconds + self.ready_timeout_seconds: Final = ready_timeout_seconds + self._stopping: Final = threading.Event() + self._lock: Final = threading.Lock() + self._process: subprocess.Popen[bytes] | None = None + + @property + def pid(self) -> int | None: + with self._lock: + return None if self._process is None else self._process.pid + + def _spawn(self) -> subprocess.Popen[bytes] | PgBouncerError | None: + """Start a child, or None once ``stop`` ran; both take the lock so no child can slip in after a stop. + + The port has to be free first: a listener that is already there would + pass the readiness check while the child fails to bind. + """ + with self._lock: + if self._stopping.is_set(): + return None + if _port_open(self.port): + return PgBouncerError(f"{PGBOUNCER_LISTEN_ADDR}:{self.port} is already in use by another process") + try: + process: Final = subprocess.Popen(self.argv) + except OSError as spawn_error: + return PgBouncerError(f"could not start {self.argv[0]!r}: {spawn_error}") + self._process = process + return process + + def _wait_ready(self, process: subprocess.Popen[bytes]) -> PgBouncerError | None: + deadline: Final = time.monotonic() + self.ready_timeout_seconds + while time.monotonic() < deadline: + if process.poll() is not None: + return PgBouncerError(f"pgbouncer exited with status {process.returncode} during startup") + if _port_open(self.port) and _unix_socket_open(self.socket_path): + return None + time.sleep(0.1) + if _port_open(self.port): + return PgBouncerError( + f"{PGBOUNCER_LISTEN_ADDR}:{self.port} is served by another process, not the pgbouncer that was started" + ) + return PgBouncerError( + f"pgbouncer did not start listening on {PGBOUNCER_LISTEN_ADDR}:{self.port} " + f"within {self.ready_timeout_seconds:.0f}s" + ) + + def start(self) -> PgBouncerError | None: + """Spawn PgBouncer, wait until it listens on port and unix socket, then supervise it from a daemon thread.""" + process: Final = self._spawn() + if process is None: + return PgBouncerError("pgbouncer was stopped before it started") + if isinstance(process, PgBouncerError): + return process + not_ready: Final = self._wait_ready(process) + if not_ready is not None: + self.stop() + return not_ready + self._watch(process) + return None + + def _watch(self, process: subprocess.Popen[bytes]) -> None: + threading.Thread( + target=self._supervise, args=(process,), daemon=True, name="litellm-pgbouncer-supervisor" + ).start() + + def _supervise(self, process: subprocess.Popen[bytes]) -> None: + status: Final = process.wait() + if self._stopping.is_set(): + return + verbose_proxy_logger.error( + "In-container pgbouncer (pid %s) exited with status %s; restarting in %.1fs.", + process.pid, + status, + self.restart_delay_seconds, + ) + self._restart_after_delay() + + def _restart_after_delay(self) -> None: + time.sleep(self.restart_delay_seconds) + process: Final = self._spawn() + if process is None: + return + if isinstance(process, PgBouncerError): + self._retry_restart(process.reason) + return + not_ready: Final = self._wait_ready(process) + if not_ready is None: + self._watch(process) + return + _end(process) + self._retry_restart(not_ready.reason) + + def _retry_restart(self, reason: str) -> None: + if self._stopping.is_set(): + return + verbose_proxy_logger.error( + "In-container pgbouncer could not be restarted (%s); retrying in %.1fs.", reason, self.restart_delay_seconds + ) + threading.Thread(target=self._restart_after_delay, daemon=True, name="litellm-pgbouncer-supervisor").start() + + def stop(self) -> None: + with self._lock: + self._stopping.set() + process: Final = self._process + if process is not None: + _end(process) + + +def _only_in_this_process(action: Callable[[], None]) -> Callable[[], None]: + """An exit hook that does nothing in a forked child, which inherits the parent's ``atexit`` table.""" + owner_pid: Final = os.getpid() + + def run() -> None: + if os.getpid() == owner_pid: + action() + + return run + + +def start_in_container_pgbouncer( + settings: PgBouncerSettings, + upstream_url: str, + token_auth_enabled: bool = False, + register_exit_hook: Callable[[Callable[[], None]], object] = atexit.register, +) -> str | PgBouncerError: + """Start the pooler for ``upstream_url`` and return the loopback URL the workers must use. + + The pooler lives as long as this process: it is stopped from the exit hooks + once the worker manager has returned, and only by the process that started + it (gunicorn forks its workers, so they carry the hooks too). PgBouncer + refuses to run as root, so a root proxy (the default image) has it drop to + ``nobody``. + """ + if token_auth_enabled: + return PgBouncerError(PGBOUNCER_TOKEN_AUTH_CONFLICT) + version: Final = pgbouncer_version(settings.binary) + if isinstance(version, PgBouncerError): + return version + if version < PGBOUNCER_MIN_VERSION: + return PgBouncerError( + f"PgBouncer {version[0]}.{version[1]} keeps running after failing to bind its TCP port, so the proxy " + f"cannot tell it apart from another listener; {PGBOUNCER_MIN_VERSION[0]}.{PGBOUNCER_MIN_VERSION[1]} " + "or newer is required" + ) + runtime_dir: Final = Path(tempfile.mkdtemp(prefix="litellm-pgbouncer-")) + register_exit_hook(_only_in_this_process(lambda: shutil.rmtree(runtime_dir, ignore_errors=True))) + run_as_user: Final = PGBOUNCER_UNPRIVILEGED_USER if os.geteuid() == 0 else None + plan: Final = plan_pgbouncer(upstream_url, settings, runtime_dir, run_as_user) + if isinstance(plan, PgBouncerError): + return plan + ini_path: Final = write_pgbouncer_files(plan, runtime_dir, run_as_user) + if isinstance(ini_path, PgBouncerError): + return ini_path + pooler: Final = PgBouncerProcess( + argv=(settings.binary, str(ini_path)), + port=settings.port, + socket_path=unix_socket_path(runtime_dir, settings.port), + ) + failed: Final = pooler.start() + if failed is not None: + return failed + register_exit_hook(_only_in_this_process(pooler.stop)) + verbose_proxy_logger.info( + "In-container pgbouncer (pid %s) listening on %s:%s; capping this pod at %s upstream database connections.", + pooler.pid, + PGBOUNCER_LISTEN_ADDR, + settings.port, + settings.max_db_connections, + ) + return plan.pooled_url diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index e245367b1b4..16d76ff0415 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -18,6 +18,7 @@ from pydantic import BaseModel, ConfigDict import litellm from litellm.constants import DEFAULT_NUM_WORKERS_LITELLM_PROXY +from litellm.proxy.db.pgbouncer import PgBouncerError, PgBouncerSettings, start_in_container_pgbouncer from litellm.proxy.db.query_engine_reaper import start_query_engine_reaper if TYPE_CHECKING: @@ -1377,6 +1378,21 @@ def run_server( print( f"Unable to connect to DB. DATABASE_URL found in environment, but prisma package not found." # noqa: F541 ) + pgbouncer_settings: Final = PgBouncerSettings() + upstream_database_url: Final = os.getenv("DATABASE_URL") + if pgbouncer_settings.enabled and upstream_database_url is not None: + pooled_database_url: Final = start_in_container_pgbouncer( + pgbouncer_settings, upstream_database_url, token_auth_enabled=wants_rds_iam or wants_azure_entra + ) + if isinstance(pooled_database_url, PgBouncerError): + print( + f"\033[1;31mLiteLLM Proxy: LITELLM_PGBOUNCER_ENABLED is set but the in-container pgbouncer " + f"could not start: {pooled_database_url.reason}\033[0m", + file=sys.stderr, + flush=True, + ) + sys.exit(1) + os.environ["DATABASE_URL"] = pooled_database_url if port == 4000 and ProxyInitializationHelpers._is_port_in_use(port): port = random.randint(1024, 49152) if prometheus_metrics_port == port: diff --git a/terraform/litellm/aws/README.md b/terraform/litellm/aws/README.md index 6d3e0269a15..986428151e0 100644 --- a/terraform/litellm/aws/README.md +++ b/terraform/litellm/aws/README.md @@ -258,6 +258,39 @@ gateway_metrics_port = 4001 gateway_metrics_scrape_cidrs = ["10.0.0.0/16"] ``` +### In-container connection pool + +Each of the `gateway_num_workers` uvicorn workers opens its own Prisma pool +straight to Postgres, so one task holds `workers x connection_limit` +connections and the fleet's footprint against the database ceiling grows with +every task. `gateway_connection_pool_enabled` runs a PgBouncer (transaction +mode, loopback) inside the gateway container that all workers share, capping +the task at `gateway_pool_max_db_connections` upstream connections however +many workers it runs. `gateway_pool_max_client_conn` bounds the worker-side +connections the pooler accepts. The module sets +`LITELLM_PGBOUNCER_ENABLED`, `LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS` and +`LITELLM_PGBOUNCER_MAX_CLIENT_CONN` on the gateway container only; the backend +and the migration task keep the direct connection. + +```hcl +create_database = false +database_url = "postgresql://litellm:@db.internal:5432/litellm" +gateway_num_workers = 4 +gateway_connection_pool_enabled = true +gateway_pool_max_db_connections = 20 +gateway_pool_max_client_conn = 1000 +``` + +The pool needs a static database password, so it is only valid with an +existing database via `database_url`. The module-created Aurora authenticates +with rotating IAM tokens (see [Aurora + IAM auth](#aurora--iam-auth)), which +the pooler cannot follow, and `terraform plan` rejects that combination. + +The componentized `gateway_image` starts through `python -m gateway.launch`, +which reads these variables, starts the pooler once per task and hands the +workers its loopback URL; the classic `litellm` image honours them the same +way. + ### Scaling the gateway on requests and tokens By default the gateway service target-tracks CPU (`gateway_cpu_target`) and diff --git a/terraform/litellm/aws/ecs.tf b/terraform/litellm/aws/ecs.tf index aa2c3d558e2..149842c14ff 100644 --- a/terraform/litellm/aws/ecs.tf +++ b/terraform/litellm/aws/ecs.tf @@ -213,6 +213,12 @@ locals { # otherwise we keep the image's ENTRYPOINT and only override `command`. gateway_uvicorn_args = "--host 0.0.0.0 --port 4000 --workers ${var.gateway_num_workers}" + gateway_pool_env = var.gateway_connection_pool_enabled ? [ + { name = "LITELLM_PGBOUNCER_ENABLED", value = "true" }, + { name = "LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS", value = tostring(var.gateway_pool_max_db_connections) }, + { name = "LITELLM_PGBOUNCER_MAX_CLIENT_CONN", value = tostring(var.gateway_pool_max_client_conn) }, + ] : [] + metrics_enabled = var.gateway_metrics_port != null metrics_multiproc_dir = "/tmp/litellm_prometheus_multiproc" metrics_volume = "prometheus-multiproc" @@ -253,7 +259,7 @@ locals { backend_uvicorn_args = "--host 0.0.0.0 --port 4001" - gateway_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args};; *) exec uvicorn gateway.main:app ${local.gateway_uvicorn_args};; esac" + gateway_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run python -m gateway.launch ${local.gateway_uvicorn_args};; *) exec python -m gateway.launch ${local.gateway_uvicorn_args};; esac" backend_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn backend.main:app ${local.backend_uvicorn_args};; *) exec uvicorn backend.main:app ${local.backend_uvicorn_args};; esac" gateway_proxy_overrides = local.proxy_config_enabled ? { @@ -298,6 +304,11 @@ resource "aws_ecs_task_definition" "gateway" { ) error_message = "billing_metrics_client_cert_pem and billing_metrics_client_key_pem are both required when billing_metrics_endpoint is set." } + + precondition { + condition = !var.gateway_connection_pool_enabled || local.byo_database + error_message = "gateway_connection_pool_enabled requires an existing database via database_url with create_database = false: the module-created Aurora authenticates with IAM tokens, which the in-container pgbouncer cannot follow because it holds a static database password." + } } family = "${local.name}-gateway" @@ -323,6 +334,7 @@ resource "aws_ecs_task_definition" "gateway" { local.gateway_extra_env_list, local.proxy_config_env, local.metrics_env, + local.gateway_pool_env, ) secrets = concat(local.shared_secrets, local.gateway_extra_secrets_list) mountPoints = local.metrics_mount_points diff --git a/terraform/litellm/aws/tests/connection_pool.tftest.hcl b/terraform/litellm/aws/tests/connection_pool.tftest.hcl new file mode 100644 index 00000000000..38aa7051134 --- /dev/null +++ b/terraform/litellm/aws/tests/connection_pool.tftest.hcl @@ -0,0 +1,123 @@ +# Plan-only coverage for the in-container PgBouncer knobs on the gateway task. +# `mock_provider` keeps this offline: no AWS credentials, no API calls, no +# resources. Run from terraform/litellm/aws with `terraform test`. + +mock_provider "aws" { + mock_data "aws_iam_policy_document" { + defaults = { + json = "{\"Version\":\"2012-10-17\",\"Statement\":[]}" + } + } +} +mock_provider "random" {} + +variables { + region = "us-east-1" + tenant = "acme" + env = "test" + azs = ["us-east-1a", "us-east-1b"] + allow_plaintext_alb = true +} + +run "pool_off_by_default" { + command = plan + + assert { + condition = length(local.gateway_pool_env) == 0 + error_message = "The gateway must get no LITELLM_PGBOUNCER_* env unless gateway_connection_pool_enabled is set." + } +} + +run "pool_enabled_renders_the_three_vars_with_configured_sizes" { + command = plan + + variables { + create_database = false + database_url = "postgresql://litellm:pw@db.internal:5432/litellm" + gateway_num_workers = 4 + gateway_connection_pool_enabled = true + gateway_pool_max_db_connections = 8 + gateway_pool_max_client_conn = 250 + } + + assert { + condition = alltrue([ + length(local.gateway_pool_env) == 3, + local.gateway_pool_env[0].name == "LITELLM_PGBOUNCER_ENABLED" && local.gateway_pool_env[0].value == "true", + local.gateway_pool_env[1].name == "LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS" && local.gateway_pool_env[1].value == "8", + local.gateway_pool_env[2].name == "LITELLM_PGBOUNCER_MAX_CLIENT_CONN" && local.gateway_pool_env[2].value == "250", + ]) + error_message = "The pool env must carry the enabled flag and the configured sizes as strings." + } +} + +run "pool_enabled_uses_the_module_default_sizes" { + command = plan + + variables { + create_database = false + database_url = "postgresql://litellm:pw@db.internal:5432/litellm" + gateway_connection_pool_enabled = true + } + + assert { + condition = alltrue([ + length(local.gateway_pool_env) == 3, + local.gateway_pool_env[1].name == "LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS" && local.gateway_pool_env[1].value == "20", + local.gateway_pool_env[2].name == "LITELLM_PGBOUNCER_MAX_CLIENT_CONN" && local.gateway_pool_env[2].value == "1000", + ]) + error_message = "The pool env must fall back to the module defaults of 20 upstream and 1000 client connections." + } +} + +run "gateway_starts_through_the_pool_aware_launcher" { + command = plan + + variables { + gateway_num_workers = 4 + } + + assert { + condition = alltrue([ + strcontains(local.gateway_launch_cmd, "exec python -m gateway.launch --host 0.0.0.0 --port 4000 --workers 4"), + strcontains(local.gateway_launch_cmd, "exec ddtrace-run python -m gateway.launch --host 0.0.0.0 --port 4000 --workers 4"), + !strcontains(local.gateway_launch_cmd, "uvicorn gateway.main:app"), + local.gateway_proxy_overrides.command[0] == local.gateway_launch_cmd, + ]) + error_message = "The gateway must start through gateway.launch (with and without ddtrace) so the pooler starts once before uvicorn forks the workers." + } +} + +run "pool_with_module_created_iam_aurora_fails_at_plan" { + command = plan + + variables { + gateway_connection_pool_enabled = true + } + + expect_failures = [ + aws_ecs_task_definition.gateway, + ] +} + +run "pool_without_any_database_fails_at_plan" { + command = plan + + variables { + create_database = false + gateway_connection_pool_enabled = true + } + + expect_failures = [ + aws_ecs_task_definition.gateway, + ] +} + +run "module_created_iam_aurora_without_the_pool_still_plans" { + command = plan + + assert { + condition = contains(local.managed_db_env, { name = "IAM_TOKEN_DB_AUTH", value = "true" }) + error_message = "Without the pool the module-created Aurora must keep IAM token auth." + } +} diff --git a/terraform/litellm/aws/variables.tf b/terraform/litellm/aws/variables.tf index 10a9b392673..ec8363dbbaf 100644 --- a/terraform/litellm/aws/variables.tf +++ b/terraform/litellm/aws/variables.tf @@ -200,6 +200,45 @@ variable "gateway_num_workers" { } } +variable "gateway_connection_pool_enabled" { + description = <<-EOT + Run an in-container PgBouncer (transaction mode, loopback) in each gateway + task, shared by every uvicorn worker. Without it each of the + `gateway_num_workers` workers opens its own Prisma pool straight to + Postgres, so a task's footprint against the database connection ceiling is + workers x connection_limit and grows with every task. Sets + LITELLM_PGBOUNCER_ENABLED / LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS / + LITELLM_PGBOUNCER_MAX_CLIENT_CONN on the gateway container only. Requires + an existing database via `database_url`: the module-created Aurora + authenticates with IAM tokens, which the pooler cannot follow because it + holds one static password for the life of the task. + EOT + type = bool + default = false +} + +variable "gateway_pool_max_db_connections" { + description = "Upstream Postgres connections one gateway task may hold when gateway_connection_pool_enabled is set, regardless of gateway_num_workers. 20 suits 4 workers; a 5000-connection database then fits roughly 200 tasks." + type = number + default = 20 + + validation { + condition = var.gateway_pool_max_db_connections >= 1 + error_message = "gateway_pool_max_db_connections must be >= 1." + } +} + +variable "gateway_pool_max_client_conn" { + description = "Client connections the in-container PgBouncer accepts from the gateway workers when gateway_connection_pool_enabled is set." + type = number + default = 1000 + + validation { + condition = var.gateway_pool_max_client_conn >= 1 + error_message = "gateway_pool_max_client_conn must be >= 1." + } +} + variable "backend_cpu" { description = "Fargate CPU units for the backend task (1024 = 1 vCPU)." type = number diff --git a/terraform/litellm/gcp/README.md b/terraform/litellm/gcp/README.md index 66095ec0108..9c71d3e15b7 100644 --- a/terraform/litellm/gcp/README.md +++ b/terraform/litellm/gcp/README.md @@ -287,6 +287,40 @@ the gateway on GKE with the Helm chart's `targetTokensPerSecond` (see "Dependencies only" below) rather than wiring the counter into Cloud Monitoring, which the autoscaler would ignore +### In-container connection pool + +Each of the `gateway_num_workers` uvicorn workers opens its own Prisma pool +straight to Cloud SQL, so one instance holds `workers x connection_limit` +connections and the fleet's footprint against the database ceiling grows with +every instance Cloud Run adds. `gateway_connection_pool_enabled` runs a +PgBouncer (transaction mode, loopback) inside the gateway container that all +workers share, capping the instance at `gateway_pool_max_db_connections` +upstream connections however many workers it runs. +`gateway_pool_max_client_conn` bounds the worker-side connections the pooler +accepts. The module sets `LITELLM_PGBOUNCER_ENABLED`, +`LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS` and `LITELLM_PGBOUNCER_MAX_CLIENT_CONN` +on the gateway service only; the backend service and the migrations job keep +the direct connection + +```hcl +gateway_num_workers = 4 +gateway_connection_pool_enabled = true +gateway_pool_max_db_connections = 20 +gateway_pool_max_client_conn = 1000 +``` + +The pooler holds one static database password for the life of the instance. +This stack authenticates to Cloud SQL with the Secret Manager password (see +[Database authentication](#database-authentication)), so nothing else is +needed; a Cloud SQL Auth Proxy sidecar with IAM auth would not work with the +pool + +The gateway container starts through `python -m gateway.launch` (the +componentized image's own entrypoint) rather than `uvicorn` directly. The +launcher reads these variables, starts the pooler once per instance before +uvicorn forks the workers and hands them its loopback `DATABASE_URL`. It also +honours `KEEPALIVE_TIMEOUT` from `gateway_extra_env` the way the image does + ## Tenant deployment Every resource the stack creates is named `${tenant}-litellm-${env}` (or diff --git a/terraform/litellm/gcp/cloudrun.tf b/terraform/litellm/gcp/cloudrun.tf index 405893b132b..78d4ffa2152 100644 --- a/terraform/litellm/gcp/cloudrun.tf +++ b/terraform/litellm/gcp/cloudrun.tf @@ -138,10 +138,16 @@ locals { "export DATABASE_URL_READ_REPLICA=\"postgresql://$${DATABASE_USER}:$${DATABASE_PASSWORD}@$${DATABASE_HOST_READ_REPLICA}:$${DATABASE_PORT_READ_REPLICA}/$${DATABASE_NAME}\"", ] + gateway_pool_env = var.gateway_connection_pool_enabled ? [ + { name = "LITELLM_PGBOUNCER_ENABLED", value = "true" }, + { name = "LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS", value = tostring(var.gateway_pool_max_db_connections) }, + { name = "LITELLM_PGBOUNCER_MAX_CLIENT_CONN", value = tostring(var.gateway_pool_max_client_conn) }, + ] : [] + gateway_uvicorn_args = "--host 0.0.0.0 --port 4000 --workers ${var.gateway_num_workers}" backend_uvicorn_args = "--host 0.0.0.0 --port 4001" - gateway_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args};; *) exec uvicorn gateway.main:app ${local.gateway_uvicorn_args};; esac" + gateway_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run python -m gateway.launch ${local.gateway_uvicorn_args};; *) exec python -m gateway.launch ${local.gateway_uvicorn_args};; esac" backend_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn backend.main:app ${local.backend_uvicorn_args};; *) exec uvicorn backend.main:app ${local.backend_uvicorn_args};; esac" gateway_args = join(" && ", concat( @@ -229,7 +235,7 @@ resource "google_cloud_run_v2_service" "gateway" { } dynamic "env" { - for_each = concat(local.shared_env_kv, local.gateway_otel_env_kv, local.billing_metrics_env_kv, local.gateway_extra_env_kv, local.proxy_config_env, local.metrics_env_kv) + for_each = concat(local.shared_env_kv, local.gateway_otel_env_kv, local.billing_metrics_env_kv, local.gateway_extra_env_kv, local.proxy_config_env, local.metrics_env_kv, local.gateway_pool_env) content { name = env.value.name value = env.value.value diff --git a/terraform/litellm/gcp/tests/connection_pool.tftest.hcl b/terraform/litellm/gcp/tests/connection_pool.tftest.hcl new file mode 100644 index 00000000000..439fe6b1d71 --- /dev/null +++ b/terraform/litellm/gcp/tests/connection_pool.tftest.hcl @@ -0,0 +1,135 @@ +# Plan-only coverage for the in-container PgBouncer knobs on the gateway +# service. `mock_provider` keeps this offline: no GCP credentials, no API +# calls, no resources. Run from terraform/litellm/gcp with `terraform test`. + +mock_provider "google" { + mock_resource "google_redis_instance" { + defaults = { + host = "10.0.0.4" + port = 6379 + server_ca_certs = [{ + cert = "-----BEGIN CERTIFICATE-----\nmock\n-----END CERTIFICATE-----" + }] + } + } +} + +mock_provider "google-beta" {} +mock_provider "random" {} + +variables { + project_id = "test-project" + tenant = "tenant" + env = "test" + allow_plaintext_lb = true + image_registry = "us-central1-docker.pkg.dev/test-project/litellm" +} + +run "pool_off_by_default" { + command = plan + + assert { + condition = length(local.gateway_pool_env) == 0 + error_message = "The gateway must get no LITELLM_PGBOUNCER_* env unless gateway_connection_pool_enabled is set." + } + + assert { + condition = !anytrue([ + for e in google_cloud_run_v2_service.gateway[0].template[0].containers[0].env : startswith(e.name, "LITELLM_PGBOUNCER_") + ]) + error_message = "The gateway service must carry no LITELLM_PGBOUNCER_* env by default." + } +} + +run "pool_enabled_renders_the_three_vars_with_configured_sizes" { + command = plan + + variables { + gateway_num_workers = 4 + gateway_connection_pool_enabled = true + gateway_pool_max_db_connections = 8 + gateway_pool_max_client_conn = 250 + } + + assert { + condition = alltrue([ + length(local.gateway_pool_env) == 3, + local.gateway_pool_env[0].name == "LITELLM_PGBOUNCER_ENABLED" && local.gateway_pool_env[0].value == "true", + local.gateway_pool_env[1].name == "LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS" && local.gateway_pool_env[1].value == "8", + local.gateway_pool_env[2].name == "LITELLM_PGBOUNCER_MAX_CLIENT_CONN" && local.gateway_pool_env[2].value == "250", + ]) + error_message = "The pool env must carry the enabled flag and the configured sizes as strings." + } + + assert { + condition = alltrue([ + contains([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[0].env : e.name], "LITELLM_PGBOUNCER_ENABLED"), + contains([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[0].env : e.name], "LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS"), + contains([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[0].env : e.name], "LITELLM_PGBOUNCER_MAX_CLIENT_CONN"), + ]) + error_message = "The gateway service must receive all three LITELLM_PGBOUNCER_* env vars." + } + + assert { + condition = !anytrue(concat( + [for e in google_cloud_run_v2_service.backend[0].template[0].containers[0].env : startswith(e.name, "LITELLM_PGBOUNCER_")], + [for e in google_cloud_run_v2_job.migrations[0].template[0].template[0].containers[0].env : startswith(e.name, "LITELLM_PGBOUNCER_")], + )) + error_message = "The backend service and the migrations job must keep their direct database connection." + } +} + +run "pool_enabled_uses_the_module_default_sizes" { + command = plan + + variables { + gateway_connection_pool_enabled = true + } + + assert { + condition = alltrue([ + length(local.gateway_pool_env) == 3, + local.gateway_pool_env[1].name == "LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS" && local.gateway_pool_env[1].value == "20", + local.gateway_pool_env[2].name == "LITELLM_PGBOUNCER_MAX_CLIENT_CONN" && local.gateway_pool_env[2].value == "1000", + ]) + error_message = "The pool env must fall back to the module defaults of 20 upstream and 1000 client connections." + } +} + +run "gateway_starts_through_the_pool_aware_launcher" { + command = plan + + variables { + gateway_num_workers = 4 + } + + assert { + condition = alltrue([ + strcontains(local.gateway_launch_cmd, "exec python -m gateway.launch --host 0.0.0.0 --port 4000 --workers 4"), + strcontains(local.gateway_launch_cmd, "exec ddtrace-run python -m gateway.launch --host 0.0.0.0 --port 4000 --workers 4"), + !strcontains(local.gateway_launch_cmd, "uvicorn gateway.main:app"), + endswith(google_cloud_run_v2_service.gateway[0].template[0].containers[0].args[0], local.gateway_launch_cmd), + ]) + error_message = "The gateway must start through gateway.launch (with and without ddtrace) so the pooler starts once before uvicorn forks the workers." + } + + assert { + condition = strcontains(local.backend_launch_cmd, "uvicorn backend.main:app") + error_message = "The backend has no workers to share a pooler and keeps starting uvicorn directly." + } +} + +run "pool_sizes_below_one_fail_at_plan" { + command = plan + + variables { + gateway_connection_pool_enabled = true + gateway_pool_max_db_connections = 0 + gateway_pool_max_client_conn = 0 + } + + expect_failures = [ + var.gateway_pool_max_db_connections, + var.gateway_pool_max_client_conn, + ] +} diff --git a/terraform/litellm/gcp/variables.tf b/terraform/litellm/gcp/variables.tf index a753b4584a6..f298a0431c0 100644 --- a/terraform/litellm/gcp/variables.tf +++ b/terraform/litellm/gcp/variables.tf @@ -206,6 +206,45 @@ variable "gateway_num_workers" { } } +variable "gateway_connection_pool_enabled" { + description = <<-EOT + Run an in-container PgBouncer (transaction mode, loopback) in each gateway + instance, shared by every uvicorn worker. Without it each of the + `gateway_num_workers` workers opens its own Prisma pool straight to + Cloud SQL, so an instance's footprint against the database connection + ceiling is workers x connection_limit and grows with every instance. Sets + LITELLM_PGBOUNCER_ENABLED / LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS / + LITELLM_PGBOUNCER_MAX_CLIENT_CONN on the gateway container only. The + module's Cloud SQL authenticates with the static password in Secret + Manager, which is what the pooler needs. Mirrors the AWS stack's + gateway_connection_pool_enabled. + EOT + type = bool + default = false +} + +variable "gateway_pool_max_db_connections" { + description = "Upstream Cloud SQL connections one gateway instance may hold when gateway_connection_pool_enabled is set, regardless of gateway_num_workers. 20 suits 4 workers." + type = number + default = 20 + + validation { + condition = var.gateway_pool_max_db_connections >= 1 + error_message = "gateway_pool_max_db_connections must be >= 1." + } +} + +variable "gateway_pool_max_client_conn" { + description = "Client connections the in-container PgBouncer accepts from the gateway workers when gateway_connection_pool_enabled is set." + type = number + default = 1000 + + validation { + condition = var.gateway_pool_max_client_conn >= 1 + error_message = "gateway_pool_max_client_conn must be >= 1." + } +} + # Cloud Run autoscales out of the box (request-rate driven). The min/max # bounds mirror the HPA replica bounds in helm/litellm/values.yaml so each # stack scales over the same range. Cloud Run has no direct CPU-utilization diff --git a/tests/test_gateway/test_launch.py b/tests/test_gateway/test_launch.py new file mode 100644 index 00000000000..964a7021416 --- /dev/null +++ b/tests/test_gateway/test_launch.py @@ -0,0 +1,155 @@ +import os +import socket +import sys +import textwrap +import urllib.parse +from pathlib import Path +from typing import Final, cast + +import pytest +from uvicorn.importer import import_from_string +from uvicorn.main import main as uvicorn_main + +import gateway.main +from gateway.launch import GATEWAY_APP, main, pool_database_url, uvicorn_argv +from litellm.proxy.db.db_url_settings import DatabaseURLSettings +from litellm.proxy.db.pgbouncer import PgBouncerError, PgBouncerSettings + +DB_ENV: Final = { + "DATABASE_HOST": "db.internal", + "DATABASE_PORT": "5432", + "DATABASE_USER": "litellm_pool", + "DATABASE_NAME": "litellm", + "DATABASE_PASSWORD": "p@ss", +} + + +def _free_port() -> int: + with socket.socket() as probe: + probe.bind(("127.0.0.1", 0)) + return cast(tuple[str, int], probe.getsockname())[1] + + +def _fake_pooler(tmp_path: Path) -> Path: + script: Final = tmp_path / "fake-pgbouncer" + script.write_text( + textwrap.dedent( + f"""\ + #!{sys.executable} + import configparser, select, socket, sys + if sys.argv[1:] == ["--version"]: + print("PgBouncer 1.25.2") + sys.exit(0) + ini = configparser.ConfigParser() + ini.read(sys.argv[1]) + port = ini.getint("pgbouncer", "listen_port") + tcp = socket.socket() + tcp.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + tcp.bind(("127.0.0.1", port)) + tcp.listen() + unix = socket.socket(socket.AF_UNIX) + unix.bind(ini.get("pgbouncer", "unix_socket_dir") + f"/.s.PGSQL.{{port}}") + unix.listen() + while True: + for ready in select.select([tcp, unix], [], [])[0]: + ready.accept()[0].close() + """ + ) + ) + script.chmod(0o700) + return script + + +def _query(url: str) -> dict[str, str]: + return dict(urllib.parse.parse_qsl(urllib.parse.urlsplit(url).query)) + + +@pytest.fixture +def password_env(monkeypatch: pytest.MonkeyPatch) -> dict[str, str]: + for var in ("DATABASE_URL", "IAM_TOKEN_DB_AUTH", "AZURE_POSTGRESQL_AUTH", "DATABASE_HOST_READ_REPLICA"): + monkeypatch.setenv(var, "") + monkeypatch.delenv(var) + for var, value in DB_ENV.items(): + monkeypatch.setenv(var, value) + return dict(DB_ENV) + + +def _uvicorn_params(argv: tuple[str, ...]) -> dict[str, object]: + return uvicorn_main.make_context("uvicorn", list(argv)).params + + +class TestUvicornArgv: + def test_keepalive_env_reaches_uvicorn(self): + params: Final = _uvicorn_params(uvicorn_argv(("--workers", "4"), {"KEEPALIVE_TIMEOUT": "75"})) + assert params["app"] == GATEWAY_APP + assert params["workers"] == 4 + assert params["timeout_keep_alive"] == 75 + + def test_unset_env_keeps_the_uvicorn_default(self): + assert _uvicorn_params(uvicorn_argv(("--workers", "4"), {}))["timeout_keep_alive"] == 5 + + def test_an_explicit_flag_wins_over_the_env(self): + argv: Final = uvicorn_argv(("--timeout-keep-alive", "30"), {"KEEPALIVE_TIMEOUT": "75"}) + assert _uvicorn_params(argv)["timeout_keep_alive"] == 30 + + def test_the_app_uvicorn_is_told_to_serve_is_the_trimmed_gateway(self): + assert import_from_string(cast(str, _uvicorn_params(uvicorn_argv((), {}))["app"])) is gateway.main.app + + +class TestPoolDatabaseUrl: + def test_a_disabled_pooler_yields_no_url_to_install(self, password_env: dict[str, str]): + settings: Final = DatabaseURLSettings.from_env() + settings.apply_to_env() + environ: Final = {"DATABASE_URL": "postgresql://litellm_pool:p%40ss@db.internal:5432/litellm"} + assert pool_database_url(settings, PgBouncerSettings(enabled=False), environ) is None + + def test_a_missing_upstream_url_is_reported(self, password_env: dict[str, str]): + environ: Final[dict[str, str]] = {} + outcome: Final = pool_database_url(DatabaseURLSettings.from_env(), PgBouncerSettings(enabled=True), environ) + assert isinstance(outcome, PgBouncerError) + assert "DATABASE_URL" in outcome.reason + + def test_token_auth_is_refused(self, password_env: dict[str, str], monkeypatch: pytest.MonkeyPatch, tmp_path: Path): + monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true") + environ: Final = {"DATABASE_URL": "postgresql://litellm:token@db.internal:5432/litellm"} + outcome: Final = pool_database_url( + DatabaseURLSettings.from_env(), + PgBouncerSettings(enabled=True, port=_free_port(), binary=str(_fake_pooler(tmp_path))), + environ, + ) + assert isinstance(outcome, PgBouncerError) + assert "IAM_TOKEN_DB_AUTH" in outcome.reason + + +class TestMain: + def test_workers_inherit_the_loopback_url_the_supervisor_installed( + self, password_env: dict[str, str], monkeypatch: pytest.MonkeyPatch, tmp_path: Path + ): + port: Final = _free_port() + monkeypatch.setenv("LITELLM_PGBOUNCER_ENABLED", "true") + monkeypatch.setenv("LITELLM_PGBOUNCER_PORT", str(port)) + monkeypatch.setenv("LITELLM_PGBOUNCER_BINARY", str(_fake_pooler(tmp_path))) + monkeypatch.setenv("KEEPALIVE_TIMEOUT", "75") + served: Final[list[tuple[str, ...]]] = [] + main(("--workers", "4"), serve=lambda argv: served.append(tuple(argv))) + + pooled: Final = os.environ["DATABASE_URL"] + assert urllib.parse.urlsplit(pooled).netloc == f"litellm_pool:p%40ss@127.0.0.1:{port}" + assert _query(pooled)["pgbouncer"] == "true" + assert _uvicorn_params(served[0])["timeout_keep_alive"] == 75 + + DatabaseURLSettings.from_env().apply_to_env() + worker_url: Final = os.environ["DATABASE_URL"] + assert urllib.parse.urlsplit(worker_url).netloc == f"litellm_pool:p%40ss@127.0.0.1:{port}" + assert _query(worker_url)["pgbouncer"] == "true" + + def test_a_pooler_that_cannot_start_stops_the_gateway_before_uvicorn( + self, password_env: dict[str, str], monkeypatch: pytest.MonkeyPatch, tmp_path: Path + ): + monkeypatch.setenv("LITELLM_PGBOUNCER_ENABLED", "true") + monkeypatch.setenv("LITELLM_PGBOUNCER_BINARY", str(tmp_path / "missing-pgbouncer")) + served: Final[list[tuple[str, ...]]] = [] + with pytest.raises(SystemExit) as stopped: + main(("--workers", "4"), serve=lambda argv: served.append(tuple(argv))) + assert "missing-pgbouncer" in str(stopped.value) + assert served == [] diff --git a/tests/test_litellm/proxy/db/test_pgbouncer.py b/tests/test_litellm/proxy/db/test_pgbouncer.py new file mode 100644 index 00000000000..da1d4c91f09 --- /dev/null +++ b/tests/test_litellm/proxy/db/test_pgbouncer.py @@ -0,0 +1,624 @@ +import configparser +import logging +import os +import signal +import socket +import stat +import sys +import tempfile +import textwrap +import time +import urllib.parse +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Final, cast + +import pytest + +from litellm._logging import verbose_proxy_logger +from litellm.proxy.db.pgbouncer import ( + PgBouncerError, + PgBouncerPlan, + PgBouncerProcess, + PgBouncerSettings, + pgbouncer_version, + plan_pgbouncer, + start_in_container_pgbouncer, + unix_socket_path, + write_pgbouncer_files, +) + +UPSTREAM: Final = ( + "postgresql://app:p%40ss%27w@db.internal:5433/litellm" + "?schema=public&connection_limit=10&pool_timeout=20" + "&sslmode=require&sslaccept=strict&sslcert=/certs/ca.pem" + "&options=-c%20statement_timeout%3D7000%20-c%20lock_timeout%3D3000" +) +SETTINGS: Final = PgBouncerSettings(enabled=True, port=6543, max_db_connections=8, max_client_conn=400) + + +def _plan(url: str = UPSTREAM, run_as_user: str | None = None) -> PgBouncerPlan: + plan: Final = plan_pgbouncer(url, SETTINGS, Path("/run/pgb"), run_as_user) + assert isinstance(plan, PgBouncerPlan), plan + return plan + + +def _ini(plan: PgBouncerPlan) -> configparser.ConfigParser: + parser: Final = configparser.ConfigParser(interpolation=None) + parser.read_string(plan.ini) + return parser + + +def _query(url: str) -> dict[str, str]: + return dict(urllib.parse.parse_qsl(urllib.parse.urlsplit(url).query, keep_blank_values=True)) + + +class TestPlanPgBouncer: + def test_upstream_credentials_and_timeouts_move_into_the_pgbouncer_config(self): + ini: Final = _ini(_plan()) + assert ini["databases"]["litellm"] == ( + "host='db.internal' port=5433 dbname='litellm' user='app' password='p@ss''w' " + "connect_query='SET statement_timeout TO ''7000''; SET lock_timeout TO ''3000'''" + ) + assert _plan().userlist == '"app" "p@ss\'w"\n' + + def test_an_upstream_without_a_port_is_reached_on_the_postgres_default(self): + ini: Final = _ini(_plan("postgresql://app:pw@db/litellm")) + assert ini["databases"]["litellm"] == "host='db' port=5432 dbname='litellm' user='app' password='pw'" + + def test_pool_is_sized_from_settings_in_transaction_mode(self): + pgb: Final = _ini(_plan())["pgbouncer"] + assert pgb["pool_mode"] == "transaction" + assert pgb["max_db_connections"] == "8" + assert pgb["default_pool_size"] == "8" + assert pgb["max_client_conn"] == "400" + assert pgb["auth_type"] == "scram-sha-256" + assert pgb["listen_addr"] == "127.0.0.1" + assert pgb["listen_port"] == "6543" + assert pgb["auth_file"] == "/run/pgb/userlist.txt" + assert pgb["unix_socket_dir"] == "/run/pgb" + + def test_the_app_user_can_read_the_pgbouncer_console(self): + assert _ini(_plan())["pgbouncer"]["stats_users"] == "app" + + @pytest.mark.parametrize("user", ["app,admin", "app%20admin", "app%09admin"]) + def test_a_user_pgbouncer_would_split_into_several_console_users_is_refused(self, user: str): + outcome: Final = plan_pgbouncer(f"postgresql://{user}:pw@db/litellm", SETTINGS, Path("/run/pgb"), None) + assert isinstance(outcome, PgBouncerError) + assert "stats_users" in outcome.reason + + def test_pooled_url_points_prisma_at_loopback_without_prepared_statements(self): + pooled: Final = urllib.parse.urlsplit(_plan().pooled_url) + assert (pooled.hostname, pooled.port, pooled.path) == ("127.0.0.1", 6543, "/litellm") + assert (pooled.username, pooled.password) == ("app", "p%40ss%27w") + assert _query(_plan().pooled_url) == { + "schema": "public", + "connection_limit": "10", + "pool_timeout": "20", + "pgbouncer": "true", + } + + @pytest.mark.parametrize("hop_param", ["channel_binding=require", "gssencmode=require"]) + def test_transport_params_for_the_postgres_hop_stay_off_the_plain_tcp_loopback_url(self, hop_param: str): + pooled: Final = _plan(f"postgresql://app:pw@db/litellm?connection_limit=5&{hop_param}").pooled_url + assert _query(pooled) == {"connection_limit": "5", "pgbouncer": "true"} + + def test_verified_tls_becomes_server_side_verify_full_with_a_ca_copy_in_the_runtime_dir(self): + plan: Final = _plan() + pgb: Final = _ini(plan)["pgbouncer"] + assert pgb["server_tls_sslmode"] == "verify-full" + assert pgb["server_tls_ca_file"] == "/run/pgb/server-ca.pem" + assert plan.ca_source == "/certs/ca.pem" + + def test_unverified_require_stays_require_without_a_ca_file(self): + plan: Final = _plan("postgresql://app:pw@db/litellm?sslmode=require") + pgb: Final = _ini(plan)["pgbouncer"] + assert pgb["server_tls_sslmode"] == "require" + assert "server_tls_ca_file" not in pgb + assert plan.ca_source is None + + def test_no_tls_params_default_to_prefer(self): + assert _ini(_plan("postgresql://app:pw@db/litellm"))["pgbouncer"]["server_tls_sslmode"] == "prefer" + + def test_verification_without_a_ca_bundle_is_refused(self): + outcome: Final = plan_pgbouncer( + "postgresql://app:pw@db/litellm?sslmode=require&sslaccept=strict", SETTINGS, Path("/run/pgb"), None + ) + assert isinstance(outcome, PgBouncerError) + assert "sslcert" in outcome.reason + + def test_client_certificates_are_refused(self): + outcome: Final = plan_pgbouncer( + "postgresql://app:pw@db/litellm?sslidentity=/certs/client.p12", SETTINGS, Path("/run/pgb"), None + ) + assert isinstance(outcome, PgBouncerError) + assert "sslidentity" in outcome.reason + + @pytest.mark.parametrize( + "url", + [ + "postgresql://app@db/litellm", + "postgresql://app:pw@db", + "postgresql://:pw@db/litellm", + ], + ) + def test_urls_missing_forwardable_credentials_are_refused(self, url: str): + outcome: Final = plan_pgbouncer(url, SETTINGS, Path("/run/pgb"), None) + assert isinstance(outcome, PgBouncerError) + + def test_every_options_spelling_becomes_a_set_statement(self): + options: Final = urllib.parse.quote("-c a=1 -cb=2 --c=3") + ini: Final = _ini(_plan(f"postgresql://app:pw@db/litellm?options={options}")) + assert ini["databases"]["litellm"].endswith("connect_query='SET a TO ''1''; SET b TO ''2''; SET c TO ''3'''") + + def test_options_that_are_not_settings_are_refused(self): + outcome: Final = plan_pgbouncer( + "postgresql://app:pw@db/litellm?options=-c%20search_path", SETTINGS, Path("/run/pgb"), None + ) + assert isinstance(outcome, PgBouncerError) + assert "options" in outcome.reason + + def test_run_as_user_is_only_written_when_given(self): + assert _ini(_plan(run_as_user="nobody"))["pgbouncer"]["user"] == "nobody" + assert "user" not in _ini(_plan())["pgbouncer"] + + +class TestWritePgBouncerFiles: + def test_files_hold_the_plan_and_are_private_to_the_owner(self, tmp_path: Path): + plan: Final = _plan("postgresql://app:pw@db/litellm") + ini_path: Final = write_pgbouncer_files(plan, tmp_path, None) + assert isinstance(ini_path, Path), ini_path + assert ini_path == tmp_path / "pgbouncer.ini" + assert ini_path.read_text() == plan.ini + assert (tmp_path / "userlist.txt").read_text() == plan.userlist + for path in (ini_path, tmp_path / "userlist.txt"): + assert stat.S_IMODE(path.stat().st_mode) == 0o600 + assert not (tmp_path / "server-ca.pem").exists() + + def test_the_ca_bundle_is_copied_next_to_the_ini_pgbouncer_reads(self, tmp_path: Path): + bundle: Final = tmp_path / "rds-root.pem" + bundle.write_text("-----BEGIN CERTIFICATE-----\nMIIB\n-----END CERTIFICATE-----\n") + runtime_dir: Final = tmp_path / "run" + runtime_dir.mkdir() + plan: Final = plan_pgbouncer( + f"postgresql://app:pw@db/litellm?sslmode=verify-full&sslcert={bundle}", SETTINGS, runtime_dir, None + ) + assert isinstance(plan, PgBouncerPlan), plan + ini_path: Final = write_pgbouncer_files(plan, runtime_dir, None) + assert isinstance(ini_path, Path), ini_path + ca_file: Final = Path(_ini(plan)["pgbouncer"]["server_tls_ca_file"]) + assert ca_file.parent == runtime_dir + assert ca_file.read_text() == bundle.read_text() + + def test_an_unreadable_ca_bundle_is_reported(self, tmp_path: Path): + plan: Final = plan_pgbouncer( + f"postgresql://app:pw@db/litellm?sslmode=verify-full&sslcert={tmp_path / 'missing.pem'}", + SETTINGS, + tmp_path, + None, + ) + assert isinstance(plan, PgBouncerPlan), plan + outcome: Final = write_pgbouncer_files(plan, tmp_path, None) + assert isinstance(outcome, PgBouncerError) + assert "missing.pem" in outcome.reason + assert not (tmp_path / "pgbouncer.ini").exists() + + +def _bound_port(sock: socket.socket) -> int: + return cast(tuple[str, int], sock.getsockname())[1] + + +def _free_port() -> int: + with socket.socket() as probe: + probe.bind(("127.0.0.1", 0)) + return _bound_port(probe) + + +def _fake_pooler( + tmp_path: Path, + port: int, + exit_immediately: bool = False, + port_file: Path | None = None, + bind_delay_seconds: float = 0.0, + version_banner: str = "PgBouncer 1.25.2\nlibevent 2.1.13-stable", +) -> Path: + """An executable that listens like PgBouncer: on the TCP port first, then on ``.s.PGSQL.`` in the socket dir. + + Port and socket dir come from the ini it is given, else from ``port`` and + ``tmp_path``. With ``port_file`` each start reads the port from that file + instead. ``bind_delay_seconds`` holds the bind back, like a slow start. + ``--version`` prints ``version_banner``. + """ + script: Final = tmp_path / "fake-pgbouncer" + script.write_text( + textwrap.dedent( + f"""\ + #!{sys.executable} + import configparser, os, pathlib, select, socket, sys, time + if sys.argv[1:] == ["--version"]: + print({version_banner!r}) + sys.exit(0) + if {exit_immediately!r}: + sys.exit(3) + ini = configparser.ConfigParser() + ini.read(sys.argv[1:2]) + port = ini.getint("pgbouncer", "listen_port", fallback={port}) + if not {port_file is None!r}: + port = int(pathlib.Path({str(port_file)!r}).read_text()) + socket_dir = ini.get("pgbouncer", "unix_socket_dir", fallback={str(tmp_path)!r}) + socket_path = f"{{socket_dir}}/.s.PGSQL.{{port}}" + time.sleep({bind_delay_seconds!r}) + listener = socket.socket() + listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + listener.bind(("127.0.0.1", port)) + listener.listen() + if os.path.exists(socket_path): + os.unlink(socket_path) + unix_listener = socket.socket(socket.AF_UNIX) + unix_listener.bind(socket_path) + unix_listener.listen() + while True: + for ready in select.select([listener, unix_listener], [], [])[0]: + conn, _ = ready.accept() + conn.close() + """ + ) + ) + script.chmod(0o700) + return script + + +def _listening(port: int) -> bool: + try: + with socket.create_connection(("127.0.0.1", port), timeout=0.5): + return True + except OSError: + return False + + +def _wait_until(condition: Callable[[], bool], timeout_seconds: float = 5.0) -> bool: + deadline: Final = time.monotonic() + timeout_seconds + while time.monotonic() < deadline: + if condition(): + return True + time.sleep(0.05) + return False + + +class TestPgBouncerProcess: + def test_start_waits_for_the_listener_and_stop_ends_it(self, tmp_path: Path): + port: Final = _free_port() + pooler: Final = PgBouncerProcess( + argv=(str(_fake_pooler(tmp_path, port)),), port=port, socket_path=unix_socket_path(tmp_path, port) + ) + assert pooler.start() is None + assert _listening(port) + pid: Final = pooler.pid + assert pid is not None + pooler.stop() + assert _wait_until(lambda: not _listening(port)) + with pytest.raises(ProcessLookupError): + os.kill(pid, 0) + + def test_a_crashed_pooler_is_restarted_with_a_new_pid(self, tmp_path: Path): + port: Final = _free_port() + pooler: Final = PgBouncerProcess( + argv=(str(_fake_pooler(tmp_path, port)),), + port=port, + socket_path=unix_socket_path(tmp_path, port), + restart_delay_seconds=0.1, + ) + assert pooler.start() is None + first_pid: Final = pooler.pid + assert first_pid is not None + os.kill(first_pid, signal.SIGKILL) + assert _wait_until(lambda: pooler.pid not in (None, first_pid) and _listening(port)) + pooler.stop() + assert _wait_until(lambda: not _listening(port)) + + def test_a_failed_restart_is_retried_until_the_pooler_is_back( + self, tmp_path: Path, caplog: pytest.LogCaptureFixture + ): + port: Final = _free_port() + script: Final = _fake_pooler(tmp_path, port) + pooler: Final = PgBouncerProcess( + argv=(str(script),), port=port, socket_path=unix_socket_path(tmp_path, port), restart_delay_seconds=0.1 + ) + assert pooler.start() is None + first_pid: Final = pooler.pid + assert first_pid is not None + hidden: Final = script.rename(tmp_path / "hidden") + with caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name): + os.kill(first_pid, signal.SIGKILL) + assert _wait_until(lambda: any("could not be restarted" in record.message for record in caplog.records)) + assert not _listening(port) + hidden.rename(script) + assert _wait_until(lambda: pooler.pid not in (None, first_pid) and _listening(port)) + pooler.stop() + assert _wait_until(lambda: not _listening(port)) + + def test_a_replacement_that_never_listens_is_replaced_again(self, tmp_path: Path, caplog: pytest.LogCaptureFixture): + port: Final = _free_port() + port_file: Final = tmp_path / "port" + port_file.write_text(str(port)) + script: Final = _fake_pooler(tmp_path, port, port_file=port_file) + pooler: Final = PgBouncerProcess( + argv=(str(script),), + port=port, + socket_path=unix_socket_path(tmp_path, port), + restart_delay_seconds=0.1, + ready_timeout_seconds=0.3, + ) + assert pooler.start() is None + first_pid: Final = pooler.pid + assert first_pid is not None + wrong_port: Final = _free_port() + port_file.write_text(str(wrong_port)) + with caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name): + os.kill(first_pid, signal.SIGKILL) + assert _wait_until(lambda: _listening(wrong_port)) + assert _wait_until(lambda: any("did not start listening" in record.message for record in caplog.records)) + port_file.write_text(str(port)) + assert _wait_until(lambda: _listening(port)) + assert _wait_until(lambda: not _listening(wrong_port)) + pooler.stop() + assert _wait_until(lambda: not _listening(port)) + + def test_stopping_during_the_restart_delay_leaves_no_pooler_behind(self, tmp_path: Path): + port: Final = _free_port() + pooler: Final = PgBouncerProcess( + argv=(str(_fake_pooler(tmp_path, port)),), + port=port, + socket_path=unix_socket_path(tmp_path, port), + restart_delay_seconds=0.3, + ) + assert pooler.start() is None + first_pid: Final = pooler.pid + assert first_pid is not None + os.kill(first_pid, signal.SIGKILL) + assert _wait_until(lambda: not _listening(port)) + pooler.stop() + time.sleep(1.0) + assert not _listening(port) + assert pooler.pid == first_pid + + def test_a_stopped_pooler_is_not_restarted(self, tmp_path: Path, caplog: pytest.LogCaptureFixture): + port: Final = _free_port() + pooler: Final = PgBouncerProcess( + argv=(str(_fake_pooler(tmp_path, port)),), + port=port, + socket_path=unix_socket_path(tmp_path, port), + restart_delay_seconds=0.1, + ) + assert pooler.start() is None + with caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name): + pooler.stop() + time.sleep(0.5) + assert not _listening(port) + assert caplog.records == [] + + def test_a_pooler_that_exits_during_startup_is_reported(self, tmp_path: Path): + port: Final = _free_port() + pooler: Final = PgBouncerProcess( + argv=(str(_fake_pooler(tmp_path, port, exit_immediately=True)),), + port=port, + socket_path=unix_socket_path(tmp_path, port), + ) + outcome: Final = pooler.start() + assert isinstance(outcome, PgBouncerError) + assert "status 3" in outcome.reason + + def test_a_missing_binary_is_reported(self, tmp_path: Path): + outcome: Final = PgBouncerProcess( + argv=("/nonexistent/pgbouncer",), port=_free_port(), socket_path=tmp_path / "sock" + ).start() + assert isinstance(outcome, PgBouncerError) + assert "/nonexistent/pgbouncer" in outcome.reason + + def test_a_port_owned_by_someone_else_is_refused_before_spawning(self, tmp_path: Path): + with socket.socket() as squatter: + squatter.bind(("127.0.0.1", 0)) + squatter.listen() + port: Final = _bound_port(squatter) + pooler: Final = PgBouncerProcess( + argv=(str(_fake_pooler(tmp_path, port)),), port=port, socket_path=unix_socket_path(tmp_path, port) + ) + outcome: Final = pooler.start() + assert isinstance(outcome, PgBouncerError) + assert f"127.0.0.1:{port} is already in use" in outcome.reason + assert pooler.pid is None + + def test_a_replacement_waits_until_a_squatter_leaves_the_port( + self, tmp_path: Path, caplog: pytest.LogCaptureFixture + ): + port: Final = _free_port() + pooler: Final = PgBouncerProcess( + argv=(str(_fake_pooler(tmp_path, port)),), + port=port, + socket_path=unix_socket_path(tmp_path, port), + restart_delay_seconds=0.5, + ) + assert pooler.start() is None + first_pid: Final = pooler.pid + assert first_pid is not None + os.kill(first_pid, signal.SIGKILL) + assert _wait_until(lambda: not _listening(port)) + with socket.socket() as squatter, caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name): + squatter.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + squatter.bind(("127.0.0.1", port)) + squatter.listen() + assert _wait_until(lambda: any("already in use" in record.message for record in caplog.records)) + assert pooler.pid == first_pid + assert _wait_until(lambda: pooler.pid not in (None, first_pid) and _listening(port)) + pooler.stop() + assert _wait_until(lambda: not _listening(port)) + + def test_a_listener_that_grabs_the_port_after_the_spawn_is_not_taken_for_the_pooler(self, tmp_path: Path): + port: Final = _free_port() + pooler: Final = PgBouncerProcess( + argv=(str(_fake_pooler(tmp_path, port, bind_delay_seconds=0.5)),), + port=port, + socket_path=unix_socket_path(tmp_path, port), + ready_timeout_seconds=3.0, + ) + with socket.socket() as squatter, ThreadPoolExecutor(max_workers=1) as starter: + squatter.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + starting: Final = starter.submit(pooler.start) + assert _wait_until(lambda: pooler.pid is not None) + squatter.bind(("127.0.0.1", port)) + squatter.listen() + outcome: Final = starting.result() + assert isinstance(outcome, PgBouncerError) + assert "exited with status 1" in outcome.reason + + def test_a_port_served_by_a_stranger_while_the_pooler_is_still_starting_is_reported(self, tmp_path: Path): + port: Final = _free_port() + pooler: Final = PgBouncerProcess( + argv=(str(_fake_pooler(tmp_path, port, bind_delay_seconds=30.0)),), + port=port, + socket_path=unix_socket_path(tmp_path, port), + ready_timeout_seconds=0.5, + ) + with socket.socket() as squatter, ThreadPoolExecutor(max_workers=1) as starter: + squatter.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + starting: Final = starter.submit(pooler.start) + assert _wait_until(lambda: pooler.pid is not None) + squatter.bind(("127.0.0.1", port)) + squatter.listen() + outcome: Final = starting.result() + assert isinstance(outcome, PgBouncerError) + assert f"127.0.0.1:{port} is served by another process" in outcome.reason + pid: Final = pooler.pid + assert pid is not None + with pytest.raises(ProcessLookupError): + os.kill(pid, 0) + + def test_a_pooler_that_never_listens_times_out(self, tmp_path: Path): + port: Final = _free_port() + pooler: Final = PgBouncerProcess( + argv=(str(_fake_pooler(tmp_path, _free_port())),), + port=port, + socket_path=unix_socket_path(tmp_path, port), + ready_timeout_seconds=0.5, + ) + outcome: Final = pooler.start() + assert isinstance(outcome, PgBouncerError) + assert "did not start listening" in outcome.reason + pid: Final = pooler.pid + assert pid is not None + with pytest.raises(ProcessLookupError): + os.kill(pid, 0) + + +def _runtime_dir_listening_on(port: int) -> Path: + matches: Final = tuple( + ini.parent + for ini in Path(tempfile.gettempdir()).glob("litellm-pgbouncer-*/pgbouncer.ini") + if f"listen_port = {port}\n" in ini.read_text() + ) + assert len(matches) == 1, matches + return matches[0] + + +class TestStartInContainerPgBouncer: + def test_returns_the_loopback_url_once_the_pooler_listens(self, tmp_path: Path): + port: Final = _free_port() + settings: Final = PgBouncerSettings(enabled=True, port=port, binary=str(_fake_pooler(tmp_path, port))) + pooled: Final = start_in_container_pgbouncer(settings, "postgresql://app:pw@db/litellm?connection_limit=5") + assert pooled == f"postgresql://app:pw@127.0.0.1:{port}/litellm?connection_limit=5&pgbouncer=true" + assert _listening(port) + + @pytest.mark.filterwarnings("ignore:This process .* is multi-threaded:DeprecationWarning") + def test_a_forked_worker_exiting_leaves_the_pooler_and_its_files_to_the_parent(self, tmp_path: Path): + port: Final = _free_port() + settings: Final = PgBouncerSettings(enabled=True, port=port, binary=str(_fake_pooler(tmp_path, port))) + exit_hooks: Final[list[Callable[[], None]]] = [] + pooled: Final = start_in_container_pgbouncer( + settings, "postgresql://app:pw@db/litellm", register_exit_hook=exit_hooks.append + ) + assert isinstance(pooled, str), pooled + runtime_dir: Final = _runtime_dir_listening_on(port) + + worker: Final = os.fork() + if worker == 0: + try: + for hook in exit_hooks: + hook() + finally: + os._exit(0) + if not _wait_until(lambda: os.waitpid(worker, os.WNOHANG) != (0, 0)): + os.kill(worker, signal.SIGKILL) + pytest.fail("the forked worker did not exit: an exit hook blocked on state inherited from the parent") + assert _listening(port) + assert (runtime_dir / "pgbouncer.ini").exists() + + for hook in exit_hooks: + hook() + assert _wait_until(lambda: not _listening(port)) + assert not runtime_dir.exists() + + def test_a_bad_upstream_url_is_reported_without_starting_anything(self, tmp_path: Path): + port: Final = _free_port() + settings: Final = PgBouncerSettings(enabled=True, port=port, binary=str(_fake_pooler(tmp_path, port))) + outcome: Final = start_in_container_pgbouncer(settings, "postgresql://app@db/litellm") + assert isinstance(outcome, PgBouncerError) + assert not _listening(port) + + def test_token_auth_is_refused_without_starting_anything(self, tmp_path: Path): + port: Final = _free_port() + settings: Final = PgBouncerSettings(enabled=True, port=port, binary=str(_fake_pooler(tmp_path, port))) + outcome: Final = start_in_container_pgbouncer( + settings, "postgresql://app:pw@db/litellm", token_auth_enabled=True + ) + assert isinstance(outcome, PgBouncerError) + assert "IAM_TOKEN_DB_AUTH" in outcome.reason + assert "AZURE_POSTGRESQL_AUTH" in outcome.reason + assert not _listening(port) + + def test_a_pgbouncer_that_survives_a_failed_tcp_bind_is_refused_without_starting(self, tmp_path: Path): + port: Final = _free_port() + binary: Final = _fake_pooler(tmp_path, port, version_banner="PgBouncer 1.18.1\nlibevent 2.1.12-stable") + settings: Final = PgBouncerSettings(enabled=True, port=port, binary=str(binary)) + outcome: Final = start_in_container_pgbouncer(settings, "postgresql://app:pw@db/litellm") + assert isinstance(outcome, PgBouncerError) + assert "PgBouncer 1.18" in outcome.reason + assert "1.19" in outcome.reason + assert not _listening(port) + + def test_the_first_version_that_dies_on_a_failed_tcp_bind_is_accepted(self, tmp_path: Path): + port: Final = _free_port() + binary: Final = _fake_pooler(tmp_path, port, version_banner="PgBouncer 1.19.0") + settings: Final = PgBouncerSettings(enabled=True, port=port, binary=str(binary)) + assert start_in_container_pgbouncer(settings, "postgresql://app:pw@db/litellm") == ( + f"postgresql://app:pw@127.0.0.1:{port}/litellm?pgbouncer=true" + ) + assert _listening(port) + + +class TestPgBouncerVersion: + def test_reads_major_and_minor_from_the_banner(self, tmp_path: Path): + assert pgbouncer_version(str(_fake_pooler(tmp_path, _free_port()))) == (1, 25) + + def test_a_binary_that_cannot_run_is_reported(self, tmp_path: Path): + outcome: Final = pgbouncer_version(str(tmp_path / "missing-pgbouncer")) + assert isinstance(outcome, PgBouncerError) + assert "missing-pgbouncer" in outcome.reason + + def test_a_banner_without_a_version_is_reported(self, tmp_path: Path): + outcome: Final = pgbouncer_version(str(_fake_pooler(tmp_path, _free_port(), version_banner="something else"))) + assert isinstance(outcome, PgBouncerError) + assert "something else" in outcome.reason + + +class TestPgBouncerSettings: + def test_reads_the_litellm_pgbouncer_env_vars(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_PGBOUNCER_ENABLED", "true") + monkeypatch.setenv("LITELLM_PGBOUNCER_PORT", "7000") + monkeypatch.setenv("LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS", "12") + settings: Final = PgBouncerSettings() + assert (settings.enabled, settings.port, settings.max_db_connections) == (True, 7000, 12) + + def test_defaults_are_off(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("LITELLM_PGBOUNCER_ENABLED", raising=False) + assert PgBouncerSettings().enabled is False diff --git a/tests/test_litellm/test_component_entrypoint.py b/tests/test_litellm/test_component_entrypoint.py index 09837d2b233..0c2a533b8bc 100644 --- a/tests/test_litellm/test_component_entrypoint.py +++ b/tests/test_litellm/test_component_entrypoint.py @@ -38,12 +38,16 @@ _STUB_TEMPLATE = """#!/bin/sh _ENTRYPOINT_RE = re.compile(r"^ENTRYPOINT\s+(\[.*\])\s*$", re.MULTILINE) _CMD_RE = re.compile(r"^CMD\s+(\[.*\])\s*$", re.MULTILINE) _COPY_RE = re.compile(r"^COPY\s+(?!--from)(\S+)\s+(\S+)\s*$", re.MULTILINE) -_APP_TARGET_RE = re.compile(r"(?:gateway|backend)\.main:app") +_APP_TARGET_RE = re.compile(r"(?:gateway|backend)\.main:app|gateway\.launch") _TF_STRING_LOCAL_RE = re.compile(r'^\s*(\w+)\s*=\s*"((?:[^"\\]|\\.)*)"\s*$', re.MULTILINE) _TF_INTERPOLATION_RE = re.compile(r"\$\{(local|var)\.(\w+)\}") TERRAFORM_LAUNCH_SITES = {TERRAFORM_ECS: 2, TERRAFORM_CLOUDRUN: 2} TERRAFORM_VAR_STUBS = {"gateway_num_workers": "2"} +COMPONENT_LAUNCHERS = { + "gateway": ("python", "-m", "gateway.launch"), + "backend": ("uvicorn", "backend.main:app"), +} _MAX_INTERPOLATION_PASSES = 5 @@ -63,7 +67,7 @@ def _run_entrypoint( """Run `script` with stubbed executables on PATH and return the recorded lines.""" bin_dir = tmp_path / "bin" bin_dir.mkdir(parents=True) - _write_stubs(bin_dir, ("ddtrace-run", "uvicorn", "litellm")) + _write_stubs(bin_dir, ("ddtrace-run", "uvicorn", "python", "litellm")) record = tmp_path / "record.txt" env = { @@ -330,22 +334,63 @@ def test_entrypoint_script_has_no_carriage_returns() -> None: @pytest.mark.parametrize( - "dockerfile, app_target", + "dockerfile, launcher", [ - (GATEWAY_DOCKERFILE, "gateway.main:app"), - (BACKEND_DOCKERFILE, "backend.main:app"), + (GATEWAY_DOCKERFILE, "python -m gateway.launch"), + (BACKEND_DOCKERFILE, "uvicorn backend.main:app"), ], ) -def test_component_images_launch_uvicorn_through_the_entrypoint(dockerfile: Path, app_target: str) -> None: +def test_component_images_launch_uvicorn_through_the_entrypoint(dockerfile: Path, launcher: str) -> None: entrypoint = " ".join(_entrypoint_argv(dockerfile)) assert IMAGE_ENTRYPOINT_PATH in entrypoint, f"{dockerfile} bypasses the ddtrace-aware entrypoint" - assert app_target in entrypoint - assert entrypoint.index(IMAGE_ENTRYPOINT_PATH) < entrypoint.index("uvicorn"), ( + assert launcher in entrypoint + assert entrypoint.index(IMAGE_ENTRYPOINT_PATH) < entrypoint.index(launcher), ( f"{dockerfile} must invoke uvicorn through the entrypoint, not the other way around" ) +@pytest.mark.parametrize( + "use_ddtrace, num_workers, expected_exec, expected_args", + [ + (None, "4", "exec=python", "args=-m gateway.launch --workers 4 --host 0.0.0.0 --port 4000"), + (None, None, "exec=python", "args=-m gateway.launch --workers 1 --host 0.0.0.0 --port 4000"), + ("true", "4", "exec=ddtrace-run", "args=python -m gateway.launch --workers 4 --host 0.0.0.0 --port 4000"), + ], +) +def test_gateway_image_execs_the_supervisor_with_its_worker_count( + use_ddtrace: str | None, num_workers: str | None, expected_exec: str, expected_args: str, tmp_path: Path +) -> None: + """Run the gateway image's ENTRYPOINT + CMD and record what the container execs. + + The Dockerfile's `/app/...` script path is resolved to the checked-in script and `python` + is stubbed on PATH, so the assertion is on the argv `gateway.launch` receives, not on the + Dockerfile text. + """ + entrypoint = tuple( + part.replace(IMAGE_ENTRYPOINT_PATH, str(COMPONENT_ENTRYPOINT)) for part in _entrypoint_argv(GATEWAY_DOCKERFILE) + ) + bin_dir = tmp_path / "bin" + bin_dir.mkdir(parents=True) + _write_stubs(bin_dir, ("ddtrace-run", "python", "uvicorn")) + record = tmp_path / "record.txt" + overrides = {"USE_DDTRACE": use_ddtrace, "NUM_WORKERS": num_workers} + env = { + **{k: v for k, v in os.environ.items() if k not in ("DD_TRACE_OPENAI_ENABLED", *overrides)}, + **{k: v for k, v in overrides.items() if v is not None}, + "PATH": f"{bin_dir}{os.pathsep}{os.environ['PATH']}", + "RECORD": str(record), + "PYTHONPATH": PYTHONPATH_SENTINEL, + } + + result = subprocess.run( + [*entrypoint, *_cmd_argv(GATEWAY_DOCKERFILE)], env=env, capture_output=True, text=True, check=False + ) + + assert result.returncode == 0, f"stdout={result.stdout} stderr={result.stderr}" + assert tuple(record.read_text().splitlines())[:2] == (expected_exec, expected_args) + + @pytest.mark.parametrize("dockerfile", [GATEWAY_DOCKERFILE, BACKEND_DOCKERFILE]) def test_component_images_make_the_entrypoint_executable(dockerfile: Path) -> None: body = dockerfile.read_text() @@ -367,17 +412,18 @@ def test_terraform_launch_command_matches_the_script_contract( implementations under the same environment and asserts they agree on which binary is exec'd and on whether the openai integration is disabled. """ - app_target = f"{component}.main:app" + launcher = COMPONENT_LAUNCHERS[component] + app_target = " ".join(launcher[1:]) command = _resolve_tf_local(terraform_file, f"{component}_launch_cmd") bin_dir = tmp_path / "bin" bin_dir.mkdir(parents=True) - _write_stubs(bin_dir, ("ddtrace-run", "uvicorn")) + _write_stubs(bin_dir, ("ddtrace-run", "uvicorn", "python")) from_terraform = _run_shell_command(command, bin_dir, tmp_path / "terraform.txt", use_ddtrace) from_script = _run_entrypoint( COMPONENT_ENTRYPOINT, - ("uvicorn", app_target), + launcher, use_ddtrace=use_ddtrace, tmp_path=tmp_path / "script", ) @@ -387,12 +433,14 @@ def test_terraform_launch_command_matches_the_script_contract( ) assert from_terraform[2] == from_script[2], f"{terraform_file} disagrees with the script on the openai integration" assert app_target in from_terraform[1] + assert "gateway.main:app" not in from_terraform[1], f"{terraform_file} bypasses the gateway.launch supervisor" if use_ddtrace in TRUTHY_USE_DDTRACE: assert from_terraform[0] == "exec=ddtrace-run" + assert from_terraform[1].startswith(f"args={launcher[0]} ") assert from_terraform[2] == "DD_TRACE_OPENAI_ENABLED=False" else: - assert from_terraform[0] == "exec=uvicorn" + assert from_terraform[0] == f"exec={launcher[0]}" assert from_terraform[2] == "DD_TRACE_OPENAI_ENABLED="