diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 1301bfb0e60..85291b49880 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -40,6 +40,7 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac The proof must be completely e2e with no mocks, using, for example, actual LLM calls costing real $. `pytest` commands are not enough For bug fixes: show reproduction before the fix and passing behavior after Include the commit hash each proof was captured at, for both the before and the after runs + If the change applies to all three LLM endpoints (/v1/responses, /v1/chat/completions, /v1/messages), include proof for every single one of them, not just one For new features: show the feature working end-to-end For UI changes: include before/after screenshots --> diff --git a/.github/workflows/auto_update_price_and_context_window.yml b/.github/workflows/auto_update_price_and_context_window.yml index 1a638a4a331..d391c0bd6ce 100644 --- a/.github/workflows/auto_update_price_and_context_window.yml +++ b/.github/workflows/auto_update_price_and_context_window.yml @@ -24,9 +24,12 @@ jobs: - name: Update JSON Data run: | uv run --frozen --with 'aiohttp==3.13.3' python ".github/workflows/auto_update_price_and_context_window_file.py" + - name: Regenerate JSON Schema + run: | + uv run --frozen python ci_cd/generate_model_prices_schema.py - name: Create Pull Request run: | - git add model_prices_and_context_window.json + git add model_prices_and_context_window.json model_prices_and_context_window.schema.json git commit -m "Update model_prices_and_context_window.json file: $(date +'%Y-%m-%d')" gh pr create --title "Update model_prices_and_context_window.json file" \ --body "Automated update for model_prices_and_context_window.json" \ diff --git a/.github/workflows/test-model-map.yaml b/.github/workflows/test-model-map.yaml index b2170d9f6a4..cf4b0eb21a1 100644 --- a/.github/workflows/test-model-map.yaml +++ b/.github/workflows/test-model-map.yaml @@ -22,3 +22,12 @@ jobs: - name: Validate model_prices_and_context_window.json run: | jq empty model_prices_and_context_window.json + + - name: Set up uv + uses: ./.github/actions/setup-uv-with-retries + with: + version: "0.10.9" + + - name: Check model_prices_and_context_window.schema.json is in sync + run: | + uv run --frozen python ci_cd/generate_model_prices_schema.py --check diff --git a/.gitignore b/.gitignore index b812d45e349..13f2202305d 100644 --- a/.gitignore +++ b/.gitignore @@ -141,3 +141,4 @@ crash.*.log .coverage ui/litellm-dashboard/out/ +litellm.log diff --git a/CLAUDE.md b/CLAUDE.md index 1a4826d51e9..c3c8138d2ac 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -61,6 +61,8 @@ Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages When working on a PR, keep the PR description in sync with new commits being made +Replies/rebuttals to AI PR review bots must be 15-25 word human-readable replies + Monkeypatching attributes of a class to do testing is an anti-pattern. Prefer dependency-injecting things into classes. That way, at unit test time, you can pass a mocked dependency in Do not put names of customers or customer company names in code, PR descriptions, issue bodies, etc. This means never mention literally any company name. Especially if you're about to say a sentence mentioning that the reason the PR exists was a feature/model/bug fix/etc. requested by a company. That's the indication that you should replace that company name with "the customer". e.g. not "Model request from Acme (Pylon #1234)" but "Model request from a customer (Pylon #1234)". This is because the codebase is public. The only exception is for publicly known providers or vendors such as OpenAI, Anthropic, AWS Bedrock, etc. only IF we're adding support for that provider/vendor in general and NOT if that PR or whatnot was a request by one of them, and they're actually one of our customers. diff --git a/Makefile b/Makefile index 8b657dcb465..e9b2fb9d8f1 100644 --- a/Makefile +++ b/Makefile @@ -176,6 +176,8 @@ lint-ruff-FULL-dev: install-dev if [ -n "$$files" ]; then echo "$$files" | xargs $(UV_RUN) ruff check; \ else echo "No changed .py files to check."; fi +lint-basedpyright lint-basedpyright-budget-update: export NODE_OPTIONS := --max-old-space-size=12288 + lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 28602fc235f..65142091712 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,12 +1,12 @@ { "reportAny": { - "limit": 34906 + "limit": 31903 }, "reportArgumentType": { - "limit": 2701 + "limit": 2645 }, "reportAssignmentType": { - "limit": 330 + "limit": 329 }, "reportAttributeAccessIssue": { "limit": 516 @@ -18,13 +18,13 @@ "limit": 59 }, "reportDeprecated": { - "limit": 326 + "limit": 325 }, "reportDuplicateImport": { "limit": 42 }, "reportExplicitAny": { - "limit": 10230 + "limit": 10214 }, "reportFunctionMemberAccess": { "limit": 11 @@ -54,10 +54,10 @@ "limit": 0 }, "reportMissingParameterType": { - "limit": 5893 + "limit": 5869 }, "reportMissingTypeArgument": { - "limit": 15886 + "limit": 15861 }, "reportMissingTypeStubs": { "limit": 41 @@ -72,7 +72,7 @@ "limit": 0 }, "reportOptionalMemberAccess": { - "limit": 1085 + "limit": 1079 }, "reportOptionalOperand": { "limit": 0 @@ -84,13 +84,13 @@ "limit": 77 }, "reportPrivateUsage": { - "limit": 2438 + "limit": 2437 }, "reportRedeclaration": { "limit": 12 }, "reportReturnType": { - "limit": 225 + "limit": 219 }, "reportTypedDictNotRequiredAccess": { "limit": 27 @@ -99,31 +99,31 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 45870 + "limit": 45366 }, "reportUnknownLambdaType": { "limit": 113 }, "reportUnknownMemberType": { - "limit": 40525 + "limit": 40477 }, "reportUnknownParameterType": { - "limit": 20384 + "limit": 20338 }, "reportUnknownVariableType": { - "limit": 32099 + "limit": 32047 }, "reportUnnecessaryCast": { "limit": 177 }, "reportUnnecessaryComparison": { - "limit": 1023 + "limit": 1021 }, "reportUnnecessaryContains": { "limit": 7 }, "reportUnnecessaryIsInstance": { - "limit": 1206 + "limit": 1205 }, "reportUntypedBaseClass": { "limit": 165 @@ -135,10 +135,10 @@ "limit": 33 }, "reportUnusedFunction": { - "limit": 206 + "limit": 204 }, "reportUnusedImport": { - "limit": 1005 + "limit": 1003 }, "reportUnusedVariable": { "limit": 1297 diff --git a/ci_cd/generate_model_prices_schema.py b/ci_cd/generate_model_prices_schema.py new file mode 100644 index 00000000000..0f449f01ec9 --- /dev/null +++ b/ci_cd/generate_model_prices_schema.py @@ -0,0 +1,325 @@ +from __future__ import annotations + +import json +import sys +from pathlib import Path +from typing import Optional + +import jsonschema + +REPO_ROOT = Path(__file__).parent.parent +PRICES_PATH = REPO_ROOT / "model_prices_and_context_window.json" +SCHEMA_PATH = REPO_ROOT / "model_prices_and_context_window.schema.json" + +SPECIAL_ROOT_KEYS = frozenset({"sample_spec", "fallback_generalizations"}) + +JsonSchema = dict + +NONNEG_NUMBER: JsonSchema = {"type": "number", "minimum": 0} +NONNEG_INTEGER: JsonSchema = {"type": "integer", "minimum": 0} +BOOLEAN: JsonSchema = {"type": "boolean"} +STRING: JsonSchema = {"type": "string"} + +EXTRA_BOOLEAN_KEYS = frozenset( + { + "gemini_native_audio", + "gemini_audio_only_live", + "uses_embed_content", + "use_openai_responses_path", + "bedrock_converse_supports_strict_tools", + } +) + +OBJECT_KEYS: dict[str, JsonSchema] = { + "search_context_cost_per_query": { + "type": "object", + "description": "USD cost per web search query, keyed by search context size.", + "properties": { + "search_context_size_low": NONNEG_NUMBER, + "search_context_size_medium": NONNEG_NUMBER, + "search_context_size_high": NONNEG_NUMBER, + }, + "additionalProperties": False, + }, + "metadata": { + "type": "object", + "description": "Free-form notes about the entry (e.g. pricing derivation).", + }, + "provider_specific_entry": { + "type": "object", + "description": "Provider-internal routing hints (e.g. bedrock_invocation_schema).", + }, +} + +ARRAY_KEYS: dict[str, JsonSchema] = { + "supported_endpoints": { + "type": "array", + "description": "OpenAI-style API routes this model can be called through, e.g. /v1/chat/completions.", + "items": STRING, + }, + "supported_modalities": { + "type": "array", + "description": "Input modalities the model accepts.", + "items": {"type": "string", "enum": ["text", "image", "audio", "video"]}, + }, + "supported_output_modalities": { + "type": "array", + "description": "Output modalities the model can produce.", + "items": {"type": "string", "enum": ["text", "image", "audio", "video", "code"]}, + }, + "supported_regions": { + "type": "array", + "description": "Cloud regions the model is available in ('global' or region ids).", + "items": STRING, + }, + "tiered_pricing": { + "type": "array", + "description": "Context-length or result-count tiered rates; each tier's costs apply within its range.", + "items": { + "type": "object", + "properties": { + "range": { + "type": "array", + "description": "[min, max] prompt-token span this tier applies to.", + "items": NONNEG_NUMBER, + "minItems": 2, + "maxItems": 2, + }, + "max_results_range": { + "type": "array", + "description": "[min, max] result-count span this tier applies to (search models).", + "items": NONNEG_NUMBER, + "minItems": 2, + "maxItems": 2, + }, + "input_cost_per_token": NONNEG_NUMBER, + "output_cost_per_token": NONNEG_NUMBER, + "output_cost_per_reasoning_token": NONNEG_NUMBER, + "cache_read_input_token_cost": NONNEG_NUMBER, + "input_cost_per_query": NONNEG_NUMBER, + }, + "additionalProperties": False, + }, + }, +} + +INTEGER_KEYS: dict[str, JsonSchema] = { + "max_tokens": { + **NONNEG_INTEGER, + "description": "Legacy field: max output tokens if the provider specifies it, else max input tokens.", + }, + "max_input_tokens": { + **NONNEG_INTEGER, + "description": "Maximum prompt/context tokens the model accepts.", + }, + "max_output_tokens": { + **NONNEG_INTEGER, + "description": "Maximum tokens the model can generate in one response.", + }, + "output_vector_size": { + **NONNEG_INTEGER, + "description": "Embedding dimension for embedding models.", + }, + "prompt_cache_min_tokens": { + **NONNEG_INTEGER, + "description": "Smallest prefix the provider will actually cache; absent means the provider default applies.", + }, + "tpm": {**NONNEG_INTEGER, "description": "Provider default tokens-per-minute limit."}, + "rpm": {**NONNEG_INTEGER, "description": "Provider default requests-per-minute limit."}, +} + +NUMBER_KEYS: dict[str, JsonSchema] = { + "regional_processing_uplift_multiplier_eu": { + "type": "number", + "minimum": 1, + "description": "Multiplier applied to all token costs for EU data residency (e.g. 1.10 = +10%).", + }, + "regional_processing_uplift_multiplier_us": { + "type": "number", + "minimum": 1, + "description": "Multiplier applied to all token costs for US data residency (e.g. 1.10 = +10%).", + }, +} + +COST_DESCRIPTIONS: dict[str, str] = { + "input_cost_per_token": "USD per prompt token.", + "output_cost_per_token": "USD per generated token.", + "output_cost_per_reasoning_token": "USD per reasoning/thinking token, when billed separately.", + "cache_creation_input_token_cost": "USD per token written to the provider's prompt cache.", + "cache_read_input_token_cost": "USD per prompt token served from the provider's prompt cache.", + "input_cost_per_token_batches": "USD per prompt token via the provider's batch API.", + "output_cost_per_token_batches": "USD per generated token via the provider's batch API.", +} + + +def cost_description(key: str) -> Optional[str]: + if key in COST_DESCRIPTIONS: + return COST_DESCRIPTIONS[key] + if key.endswith("_flex"): + return "Flex service-tier rate for the same-named base field." + if key.endswith("_priority"): + return "Priority service-tier rate for the same-named base field." + if "_above_" in key: + return "Rate applied once the prompt exceeds the token threshold in the field name." + return None + + +def cost_schema(key: str) -> JsonSchema: + description = cost_description(key) + return {**NONNEG_NUMBER, "description": description} if description else dict(NONNEG_NUMBER) + + +def string_key_schemas(modes: tuple) -> dict[str, JsonSchema]: + return { + "litellm_provider": { + "type": "string", + "description": "LiteLLM provider slug; one of https://docs.litellm.ai/docs/providers.", + }, + "mode": { + "type": "string", + "description": "Primary API surface / task type of the model.", + "enum": list(modes), + }, + "source": { + "type": "string", + "description": "URL of the provider pricing/model page this entry was taken from.", + }, + "deprecation_date": { + "type": "string", + "description": "Date the provider deprecates the model, YYYY-MM-DD.", + "format": "date", + "pattern": "^\\d{4}-(0[1-9]|1[0-2])-(0[1-9]|[12]\\d|3[01])$", + }, + "web_search_billing_unit": { + "type": "string", + "description": "Whether web search is billed per query or per prompt.", + "enum": ["per_query", "per_prompt"], + }, + "bedrock_output_config_effort_ceiling": { + "type": "string", + "description": "Highest reasoning effort the Bedrock output_config accepts for this model.", + "enum": ["low", "medium", "high", "max", "xhigh"], + }, + "comment": STRING, + "audio_transcription_config": STRING, + } + + +def classify(key: str, modes: tuple) -> Optional[JsonSchema]: + curated = {**OBJECT_KEYS, **ARRAY_KEYS, **string_key_schemas(modes), **INTEGER_KEYS, **NUMBER_KEYS} + if key in curated: + return curated[key] + if key.startswith("supports_") or key in EXTRA_BOOLEAN_KEYS: + return BOOLEAN + if "cost" in key: + return cost_schema(key) + return None + + +def build_schema(prices: dict) -> JsonSchema: + entries = {name: entry for name, entry in prices.items() if name not in SPECIAL_ROOT_KEYS} + all_keys = tuple(sorted({key for entry in entries.values() for key in entry})) + modes = tuple(sorted({entry["mode"] for entry in entries.values() if "mode" in entry})) + unclassified = tuple(key for key in all_keys if classify(key, modes) is None) + if unclassified: + raise SystemExit( + f"Unclassified keys in {PRICES_PATH.name}: {', '.join(unclassified)}. " + f"Add them to the key tables in {Path(__file__).name} and rerun it." + ) + entry_properties = {key: classify(key, modes) for key in all_keys} + return { + "$schema": "https://json-schema.org/draft/2020-12/schema", + "title": "LiteLLM model_prices_and_context_window.json", + "description": ( + "Schema for LiteLLM's model price and context window registry " + "(https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json). " + "Every top-level key except 'sample_spec' and 'fallback_generalizations' is a model id, " + "optionally prefixed with its provider (e.g. 'azure/gpt-5.4'), mapping to a model entry. " + "All costs are USD per unit. New optional fields are added regularly, so consumers should " + "ignore unknown fields rather than reject them." + ), + "type": "object", + "properties": { + "sample_spec": { + "type": "object", + "description": ( + "Documentation placeholder illustrating the entry shape; not a real model and not " + "schema-conformant (several values are prose)." + ), + }, + "fallback_generalizations": { + "type": "object", + "description": "Regex rules that generalize unknown model ids to known families; not a model entry.", + "properties": { + "rules": { + "type": "array", + "items": { + "type": "object", + "properties": { + "name": STRING, + "pattern": STRING, + "description": STRING, + }, + "required": ["name", "pattern"], + "additionalProperties": True, + }, + } + }, + "additionalProperties": False, + }, + }, + "additionalProperties": {"$ref": "#/$defs/modelEntry"}, + "$defs": { + "modelEntry": { + "type": "object", + "description": ( + "Pricing, limits, and capability flags for one model. Fields other than litellm_provider " + "are optional; boolean capability flags are simply omitted when unknown or false." + ), + "required": ["litellm_provider"], + "properties": entry_properties, + "additionalProperties": True, + } + }, + } + + +def render(schema: JsonSchema) -> str: + return json.dumps(schema, indent=2) + "\n" + + +def validation_errors(prices: dict, schema: JsonSchema) -> tuple: + validator = jsonschema.Draft202012Validator( + schema, format_checker=jsonschema.Draft202012Validator.FORMAT_CHECKER + ) + return tuple( + f"{'.'.join(str(part) for part in error.absolute_path)}: {error.message}" + for error in validator.iter_errors(prices) + ) + + +def main() -> int: + check = "--check" in sys.argv[1:] + prices = json.loads(PRICES_PATH.read_text()) + rendered = render(build_schema(prices)) + errors = validation_errors(prices, json.loads(rendered)) + if errors: + print(f"{PRICES_PATH.name} does not validate against the generated schema:") + print("\n".join(errors[:20])) + return 1 + if not check: + SCHEMA_PATH.write_text(rendered) + print(f"wrote {SCHEMA_PATH}") + return 0 + if not SCHEMA_PATH.exists() or SCHEMA_PATH.read_text() != rendered: + print( + f"{SCHEMA_PATH.name} is out of sync with {PRICES_PATH.name}. " + f"Run `python {Path(__file__).relative_to(REPO_ROOT)}` and commit the result." + ) + return 1 + print(f"{SCHEMA_PATH.name} is in sync and {PRICES_PATH.name} validates against it") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/grafana_dashboard.json b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/grafana_dashboard.json new file mode 100644 index 00000000000..70608a2ffe8 --- /dev/null +++ b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/grafana_dashboard.json @@ -0,0 +1,523 @@ +{ + "annotations": { + "list": [] + }, + "editable": true, + "fiscalYearStartMonth": 0, + "graphTooltip": 0, + "links": [], + "panels": [ + { + "type": "stat", + "title": "Requests", + "gridPos": { + "h": 4, + "w": 6, + "x": 0, + "y": 0 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "short", + "decimals": 0, + "color": { + "mode": "fixed", + "fixedColor": "blue" + } + }, + "overrides": [] + }, + "options": { + "reduceOptions": { + "calcs": [ + "lastNotNull" + ], + "fields": "", + "values": false + }, + "colorMode": "background", + "graphMode": "none" + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "instant": true, + "expr": "sum(increase(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))" + } + ], + "id": 1 + }, + { + "type": "stat", + "title": "Spend", + "description": "LiteLLM's computed cost for the selected window, from gen_ai.usage.cost", + "gridPos": { + "h": 4, + "w": 6, + "x": 6, + "y": 0 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "currencyUSD", + "decimals": 4, + "color": { + "mode": "fixed", + "fixedColor": "green" + } + }, + "overrides": [] + }, + "options": { + "reduceOptions": { + "calcs": [ + "lastNotNull" + ], + "fields": "", + "values": false + }, + "colorMode": "background", + "graphMode": "none" + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "instant": true, + "expr": "sum(increase(gen_ai_usage_cost_USD_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))" + } + ], + "id": 2 + }, + { + "type": "stat", + "title": "Tokens", + "gridPos": { + "h": 4, + "w": 6, + "x": 12, + "y": 0 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "short", + "decimals": 0, + "color": { + "mode": "fixed", + "fixedColor": "purple" + } + }, + "overrides": [] + }, + "options": { + "reduceOptions": { + "calcs": [ + "lastNotNull" + ], + "fields": "", + "values": false + }, + "colorMode": "background", + "graphMode": "none" + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "instant": true, + "expr": "sum(increase(gen_ai_client_token_usage_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))" + } + ], + "id": 3 + }, + { + "type": "stat", + "title": "p95 request duration", + "gridPos": { + "h": 4, + "w": 6, + "x": 18, + "y": 0 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "s", + "decimals": 2, + "color": { + "mode": "fixed", + "fixedColor": "orange" + } + }, + "overrides": [] + }, + "options": { + "reduceOptions": { + "calcs": [ + "lastNotNull" + ], + "fields": "", + "values": false + }, + "colorMode": "background", + "graphMode": "none" + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "instant": true, + "expr": "histogram_quantile(0.95, sum by (le) (rate(gen_ai_client_operation_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range])))" + } + ], + "id": 4 + }, + { + "type": "timeseries", + "title": "Request rate by model", + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 4 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "reqpm", + "custom": { + "lineWidth": 2, + "fillOpacity": 8, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}}", + "expr": "sum by (gen_ai_request_model) (rate(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 60" + } + ], + "id": 5 + }, + { + "type": "timeseries", + "title": "Spend rate by model", + "description": "USD per hour, derived from the gen_ai.usage.cost histogram", + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 4 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "currencyUSD", + "custom": { + "lineWidth": 2, + "fillOpacity": 8, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}}", + "expr": "sum by (gen_ai_request_model) (rate(gen_ai_usage_cost_USD_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 3600" + } + ], + "id": 6 + }, + { + "type": "timeseries", + "title": "Tokens per minute by model and type", + "description": "gen_ai.client.token.usage split by the gen_ai.token.type attribute", + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 12 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "short", + "custom": { + "lineWidth": 2, + "fillOpacity": 8, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}} {{gen_ai_token_type}}", + "expr": "sum by (gen_ai_request_model, gen_ai_token_type) (rate(gen_ai_client_token_usage_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 60" + } + ], + "id": 7 + }, + { + "type": "timeseries", + "title": "p95 request duration by model", + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 12 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "s", + "custom": { + "lineWidth": 2, + "fillOpacity": 0, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}}", + "expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_client_operation_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))" + } + ], + "id": 8 + }, + { + "type": "timeseries", + "title": "p95 time to first token (streaming)", + "description": "gen_ai.server.time_to_first_token, recorded only for streaming requests", + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 20 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "s", + "custom": { + "lineWidth": 2, + "fillOpacity": 0, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}}", + "expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_server_time_to_first_token_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))" + } + ], + "id": 9 + }, + { + "type": "timeseries", + "title": "p95 provider generation time", + "description": "gen_ai.client.response.duration, upstream generation time excluding LiteLLM overhead", + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 20 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "s", + "custom": { + "lineWidth": 2, + "fillOpacity": 0, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}}", + "expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_client_response_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))" + } + ], + "id": 10 + } + ], + "preload": false, + "refresh": "30s", + "schemaVersion": 42, + "tags": [ + "litellm", + "genai", + "opentelemetry" + ], + "templating": { + "list": [ + { + "name": "datasource", + "label": "Prometheus", + "type": "datasource", + "query": "prometheus", + "current": {}, + "hide": 0 + }, + { + "name": "service", + "label": "Service", + "type": "query", + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "query": "label_values(gen_ai_client_operation_duration_seconds_count, service_name)", + "refresh": 2, + "includeAll": true, + "multi": true, + "current": { + "text": "All", + "value": "$__all" + } + }, + { + "name": "model", + "label": "Model", + "type": "query", + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "query": "label_values(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\"}, gen_ai_request_model)", + "refresh": 2, + "includeAll": true, + "multi": true, + "current": { + "text": "All", + "value": "$__all" + } + } + ] + }, + "time": { + "from": "now-1h", + "to": "now" + }, + "timepicker": {}, + "timezone": "browser", + "title": "LiteLLM GenAI (OpenTelemetry)", + "uid": "litellm-genai-otel", + "weekStart": "" +} diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/readme.md b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/readme.md new file mode 100644 index 00000000000..c51f0166462 --- /dev/null +++ b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/readme.md @@ -0,0 +1,35 @@ +# LiteLLM GenAI dashboard (OpenTelemetry metrics) + +Dashboard for the `gen_ai.*` metrics the OpenTelemetry v2 integration emits, as opposed to the `litellm_*` Prometheus metrics the other dashboards in this folder chart. + +Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source. Panels: request count, spend, token count, p95 duration, request rate by model, spend rate per hour by model, tokens per minute split by input and output, p95 duration by model, p95 time to first token, and p95 provider generation time. Template variables for data source, service, and model. + +## Pre-requisites + +Metrics are off by default. In the proxy environment: + +```shell +LITELLM_OTEL_V2=true +LITELLM_OTEL_INTEGRATION_ENABLE_METRICS=true +OTEL_EXPORTER="otlp_http" +OTEL_ENDPOINT="" +``` + +You also need the metric attribute filter, or the panels will plot flat lines at zero. LiteLLM's default attribute set includes per-request fields, so nearly every request lands in its own time series with a single sample, and `rate()` has nothing to compute over: + +```yaml title="config.yaml" +callback_settings: + otel: + attributes: + include_list: + - gen_ai.operation.name + - gen_ai.system + - gen_ai.request.model + - gen_ai.framework +``` + +See [Grafana Cloud](https://docs.litellm.ai/docs/observability/grafana_cloud) for the full setup, and [OpenTelemetry v2](https://docs.litellm.ai/docs/observability/opentelemetry_v2#metrics) for the metric reference. + +## Note on Grafana's AI Observability integration + +Grafana Cloud ships prebuilt GenAI dashboards that query these same metric names, so they look like a drop-in alternative to this one. They are not: twenty of their twenty-two panels filter on `telemetry_sdk_name="openlit"`, a label LiteLLM does not carry and cannot be configured to add, so those panels stay empty. diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/readme.md b/cookbook/litellm_proxy_server/grafana_dashboard/readme.md index 81235c308f2..a1564a406e0 100644 --- a/cookbook/litellm_proxy_server/grafana_dashboard/readme.md +++ b/cookbook/litellm_proxy_server/grafana_dashboard/readme.md @@ -2,6 +2,10 @@ This folder contains the `json` for creating Grafana Dashboards +## [LiteLLM GenAI Dashboard (OpenTelemetry)](./dashboard_genai_otel) + +Charts the `gen_ai.*` metrics from the OpenTelemetry v2 integration: spend, tokens, request rate, and latency percentiles by model. Separate from the dashboards below, which chart the `litellm_*` Prometheus metrics. + ## [LiteLLM v2 Dashboard](./dashboard_v2) grafana_1 diff --git a/helm/litellm-helm/Chart.lock b/helm/litellm-helm/Chart.lock index f13578d8d35..d626fbb472b 100644 --- a/helm/litellm-helm/Chart.lock +++ b/helm/litellm-helm/Chart.lock @@ -5,5 +5,5 @@ dependencies: - name: redis repository: oci://registry-1.docker.io/bitnamicharts version: 18.19.1 -digest: sha256:8660fe6287f9941d08c0902f3f13731079b8cecd2a5da2fbc54e5b7aae4a6f62 -generated: "2024-03-10T02:28:52.275022+05:30" +digest: sha256:38962e231f6596b93f82a8412bbe4cf5de696caecf5775dfbbd163383eb1c009 +generated: "2026-07-28T10:21:22.511401-07:00" diff --git a/helm/litellm-helm/Chart.yaml b/helm/litellm-helm/Chart.yaml index 0aef2442bfe..8ca217825b8 100644 --- a/helm/litellm-helm/Chart.yaml +++ b/helm/litellm-helm/Chart.yaml @@ -18,7 +18,7 @@ type: application # This is the chart version. This version number should be incremented each time you make changes # to the chart and its templates, including the app version. # Versions are expected to follow Semantic Versioning (https://semver.org/) -version: 1.1.0 +version: 1.1.1 # This is the version number of the application being deployed. This version number should be # incremented each time you make changes to the application. Versions are not expected to @@ -32,10 +32,10 @@ annotations: dependencies: - name: "postgresql" - version: ">=13.3.0" + version: "14.3.1" repository: oci://registry-1.docker.io/bitnamicharts condition: db.deployStandalone - name: redis - version: ">=18.0.0" + version: "18.19.1" repository: oci://registry-1.docker.io/bitnamicharts condition: redis.enabled diff --git a/helm/litellm-helm/README.md b/helm/litellm-helm/README.md index 0edc4d2504b..4e0884dd08c 100644 --- a/helm/litellm-helm/README.md +++ b/helm/litellm-helm/README.md @@ -130,6 +130,16 @@ Set `billingMetrics.caSecretName` only when the collector is a private or test o | `db.deployStandalone` | Deploy a standalone, single instance deployment of Postgres, using the Bitnami postgresql chart. This is useful for getting started but doesn't provide HA or (by default) data backups. | `true` | | `postgresql.*` | If `db.deployStandalone` is `true`, configuration passed to the Bitnami postgresql chart. See the [Bitnami Documentation](https://github.com/bitnami/charts/tree/main/bitnami/postgresql) for full configuration details. See [values.yaml](./values.yaml) for the default configuration. | See [values.yaml](./values.yaml) | | `postgresql.auth.*` | If `db.deployStandalone` is `true`, care should be taken to ensure the default `password` and `postgres-password` values are **NOT** used. | `NoTaGrEaTpAsSwOrD` | +| `postgresql.image.*` | If `db.deployStandalone` is `true`, the image for the bundled Postgres. Pinned to a `docker.io/bitnamilegacy` build because Bitnami retired the versioned tags under `docker.io/bitnami`. | `bitnamilegacy/postgresql:16.2.0-debian-12-r6` | +| `redis.image.*` | If `redis.enabled` is `true`, the image for the bundled Redis. Pinned to a `docker.io/bitnamilegacy` build for the same reason. | `bitnamilegacy/redis:7.2.4-debian-12-r9` | + +#### Bundled Postgres image + +Bitnami removed the versioned tags from `docker.io/bitnami` and republished the archived builds under `docker.io/bitnamilegacy`, so the image defaults that ship inside the `postgresql` and `redis` subcharts no longer pull. The chart pins both to the `bitnamilegacy` copies of the exact builds those subchart versions were released with, which keeps the on-disk data directory layout unchanged for existing installs. + +Keep `postgresql.image.tag` pinned. `docker.io/bitnami/postgresql` still publishes a floating `latest`, and pointing the bundled Postgres at a different major version starts the server against a data directory it cannot read (`database files are incompatible with server`). There is no in-place way back, so crossing a major version means dumping the database with the old image and restoring it into the new one. The chart refuses to render when the tag is empty or `latest`. + +Those images no longer receive updates. For anything beyond getting started, run Postgres outside the chart and point at it with `db.useExisting`. #### Example Postgres `db.useExisting` Secret diff --git a/helm/litellm-helm/templates/_helpers.tpl b/helm/litellm-helm/templates/_helpers.tpl index 387bc3d5dc4..8f2acb20fce 100644 --- a/helm/litellm-helm/templates/_helpers.tpl +++ b/helm/litellm-helm/templates/_helpers.tpl @@ -146,3 +146,18 @@ Get redis service port {{ .Values.redis.master.service.ports.redis }} {{- end -}} {{- end -}} + +{{/* +Reject an unpinned image tag for the bundled PostgreSQL. +A floating tag lets a chart upgrade start a newer PostgreSQL major against the +existing PersistentVolumeClaim. The server then refuses to start on a data +directory written by another major version, and the only way back is a dump +taken before the change, which by that point no longer exists. +*/}} +{{- define "litellm.validateBundledPostgresImageTag" -}} +{{- $tag := .Values.postgresql.image.tag | default "" | toString -}} +{{- $digest := .Values.postgresql.image.digest | default "" | toString -}} +{{- if and (eq $digest "") (or (eq $tag "") (eq $tag "latest")) -}} +{{- fail (printf "postgresql.image.tag must be pinned to an explicit version when db.deployStandalone is true (got %q). An unpinned tag can start a different PostgreSQL major against the existing data directory, which makes the database unreadable and is not recoverable in place. Crossing a major version requires a dump and restore." $tag) -}} +{{- end -}} +{{- end -}} diff --git a/helm/litellm-helm/templates/secret-dbcredentials.yaml b/helm/litellm-helm/templates/secret-dbcredentials.yaml index 8851f5802f2..8ab89a4579e 100644 --- a/helm/litellm-helm/templates/secret-dbcredentials.yaml +++ b/helm/litellm-helm/templates/secret-dbcredentials.yaml @@ -1,4 +1,5 @@ {{- if .Values.db.deployStandalone -}} +{{- include "litellm.validateBundledPostgresImageTag" . -}} apiVersion: v1 kind: Secret metadata: diff --git a/helm/litellm-helm/templates/tests/test-servicemonitor.yaml b/helm/litellm-helm/templates/tests/test-servicemonitor.yaml index c2a4f84ec21..ef8475339c3 100644 --- a/helm/litellm-helm/templates/tests/test-servicemonitor.yaml +++ b/helm/litellm-helm/templates/tests/test-servicemonitor.yaml @@ -10,7 +10,7 @@ metadata: spec: containers: - name: test - image: bitnami/kubectl:latest + image: docker.io/bitnamilegacy/kubectl:1.29.2-debian-12-r3 command: ['sh', '-c'] args: - | diff --git a/helm/litellm-helm/tests/bundled_db_images_tests.yaml b/helm/litellm-helm/tests/bundled_db_images_tests.yaml new file mode 100644 index 00000000000..8f0860c2721 --- /dev/null +++ b/helm/litellm-helm/tests/bundled_db_images_tests.yaml @@ -0,0 +1,94 @@ +suite: test bundled database images +templates: + - charts/postgresql/templates/primary/statefulset.yaml + - charts/redis/templates/master/application.yaml + - charts/redis/templates/configmap.yaml + - charts/redis/templates/health-configmap.yaml + - charts/redis/templates/scripts-configmap.yaml + - charts/redis/templates/secret.yaml + - secret-dbcredentials.yaml + - templates/tests/test-servicemonitor.yaml +tests: + - it: should pull the bundled postgres from a repository that still publishes the pinned tag + template: charts/postgresql/templates/primary/statefulset.yaml + set: + db.deployStandalone: true + asserts: + - equal: + path: spec.template.spec.containers[0].image + value: docker.io/bitnamilegacy/postgresql:16.2.0-debian-12-r6 + + - it: should pull the bundled postgres metrics exporter from the same repository + template: charts/postgresql/templates/primary/statefulset.yaml + set: + db.deployStandalone: true + postgresql.metrics.enabled: true + asserts: + - equal: + path: spec.template.spec.containers[1].image + value: docker.io/bitnamilegacy/postgres-exporter:0.15.0-debian-12-r14 + + - it: should run the bundled postgres init container from the same repository + template: charts/postgresql/templates/primary/statefulset.yaml + set: + db.deployStandalone: true + postgresql.volumePermissions.enabled: true + asserts: + - equal: + path: spec.template.spec.initContainers[0].image + value: docker.io/bitnamilegacy/os-shell:12-debian-12-r16 + + - it: should pull the bundled redis from a repository that still publishes the pinned tag + template: charts/redis/templates/master/application.yaml + set: + redis.enabled: true + asserts: + - equal: + path: spec.template.spec.containers[0].image + value: docker.io/bitnamilegacy/redis:7.2.4-debian-12-r9 + + - it: should reject a floating postgres tag that could cross a major version on an existing volume + template: secret-dbcredentials.yaml + set: + db.deployStandalone: true + postgresql.image.tag: latest + asserts: + - failedTemplate: + errorMessage: 'postgresql.image.tag must be pinned to an explicit version when db.deployStandalone is true (got "latest"). An unpinned tag can start a different PostgreSQL major against the existing data directory, which makes the database unreadable and is not recoverable in place. Crossing a major version requires a dump and restore.' + + - it: should reject an empty postgres tag + template: secret-dbcredentials.yaml + set: + db.deployStandalone: true + postgresql.image.tag: "" + asserts: + - failedTemplate: + errorMessage: 'postgresql.image.tag must be pinned to an explicit version when db.deployStandalone is true (got ""). An unpinned tag can start a different PostgreSQL major against the existing data directory, which makes the database unreadable and is not recoverable in place. Crossing a major version requires a dump and restore.' + + - it: should accept an empty postgres tag when the image is pinned by digest + template: secret-dbcredentials.yaml + set: + db.deployStandalone: true + postgresql.image.tag: "" + postgresql.image.digest: sha256:0d0e2f1a5b3c4d6e7f8091a2b3c4d5e6f708192a3b4c5d6e7f8091a2b3c4d5e6 + asserts: + - hasDocuments: + count: 1 + + - it: should run the servicemonitor test pod from a pinned image + template: templates/tests/test-servicemonitor.yaml + set: + serviceMonitor.enabled: true + asserts: + - equal: + path: spec.containers[0].image + value: docker.io/bitnamilegacy/kubectl:1.29.2-debian-12-r3 + + - it: should not constrain the postgres tag when the bundled database is not deployed + template: secret-dbcredentials.yaml + set: + db.deployStandalone: false + postgresql.image.tag: latest + asserts: + - hasDocuments: + count: 0 diff --git a/helm/litellm-helm/values.yaml b/helm/litellm-helm/values.yaml index 0529e74d6e4..7235bb0bd78 100644 --- a/helm/litellm-helm/values.yaml +++ b/helm/litellm-helm/values.yaml @@ -328,8 +328,32 @@ lifecycle: {} # Settings for Bitnami postgresql chart (if db.deployStandalone is true, ignored # otherwise) +# +# Bitnami retired the versioned tags under docker.io/bitnami and republished the +# archived builds under docker.io/bitnamilegacy, so the subchart's own image +# defaults no longer resolve. The repository below points at the same build the +# subchart was released with, which keeps the on-disk data directory layout +# identical for existing installs. +# +# Keep the tag pinned. docker.io/bitnami still publishes a floating `latest`, +# and starting a newer PostgreSQL major against an existing data directory +# leaves the server refusing to boot ("database files are incompatible with +# server") with no way back other than a dump taken beforehand. Crossing a major +# version is a dump-and-restore, not an image bump. The chart refuses to render +# an unpinned tag for this reason postgresql: architecture: standalone + image: + repository: bitnamilegacy/postgresql + tag: 16.2.0-debian-12-r6 + volumePermissions: + image: + repository: bitnamilegacy/os-shell + tag: 12-debian-12-r16 + metrics: + image: + repository: bitnamilegacy/postgres-exporter + tag: 0.15.0-debian-12-r14 auth: username: litellm database: litellm @@ -359,9 +383,36 @@ postgresql: # When `redis.sentinel.enabled` is set, the coordination block is rendered with # `sentinel_nodes` and `service_name` (from `redis.sentinel.masterSet`) instead # of host/port, because a plain Redis client cannot talk to the sentinel port +# +# The image repositories carry the same bitnamilegacy repoint as postgresql +# above; the versioned tags the subchart ships with are gone from +# docker.io/bitnami redis: enabled: false architecture: standalone + image: + repository: bitnamilegacy/redis + tag: 7.2.4-debian-12-r9 + sentinel: + image: + repository: bitnamilegacy/redis-sentinel + tag: 7.2.4-debian-12-r7 + metrics: + image: + repository: bitnamilegacy/redis-exporter + tag: 1.58.0-debian-12-r4 + volumePermissions: + image: + repository: bitnamilegacy/os-shell + tag: 12-debian-12-r16 + sysctl: + image: + repository: bitnamilegacy/os-shell + tag: 12-debian-12-r16 + kubectl: + image: + repository: bitnamilegacy/kubectl + tag: 1.29.2-debian-12-r3 coordination: # Set to false to keep the bundled Redis for response caching only and leave # `general_settings.coordination_redis` out of the rendered config. A diff --git a/litellm-proxy-extras/litellm_proxy_extras/replica_identity.py b/litellm-proxy-extras/litellm_proxy_extras/replica_identity.py new file mode 100644 index 00000000000..dc92e9dca6a --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/replica_identity.py @@ -0,0 +1,106 @@ +"""Optional post-migration step that raises Postgres REPLICA IDENTITY to FULL. + +Logical-replication consumers (Neon / lakehouse sync and similar) need FULL +replica identity to reconstruct the old row of an UPDATE or DELETE. Prisma +leaves every table it creates at the Postgres default, so the setting has to be +re-applied by hand after each migration run. Setting +``LITELLM_SET_REPLICA_IDENTITY_FULL`` makes every migration run re-assert it. + +The statement goes through the Prisma CLI rather than a Postgres driver because +``litellm-proxy-extras`` has no runtime dependencies, while the CLI is already +required for the migrations themselves. +""" + +import subprocess +import tempfile +from pathlib import Path + +from litellm_proxy_extras._logging import logger + +REPLICA_IDENTITY_FULL_ENV_VAR = "LITELLM_SET_REPLICA_IDENTITY_FULL" + +REPLICA_IDENTITY_FULL_SQL = r""" +DO $$ +DECLARE + target regclass; +BEGIN + SET LOCAL lock_timeout = '5s'; + FOR target IN + SELECT c.oid::regclass + FROM pg_class c + JOIN pg_namespace n ON n.oid = c.relnamespace + WHERE c.relkind = 'r' + AND c.relreplident <> 'f' + AND n.nspname = ANY (current_schemas(false)) + AND c.relname LIKE 'LiteLLM\_%' + LOOP + BEGIN + EXECUTE format('ALTER TABLE %s REPLICA IDENTITY FULL', target); + EXCEPTION WHEN lock_not_available THEN + RAISE WARNING 'REPLICA IDENTITY FULL skipped for %: table busy, retrying next run', target; + END; + END LOOP; +END +$$; +""" + + +def apply_replica_identity_full( + schema_path: str, + prisma_command: str, + prisma_env: dict[str, str], +) -> bool: + """Set REPLICA IDENTITY FULL on every LiteLLM table that is not already FULL. + + Never raises. Replication metadata is not needed to serve requests, so + every failure mode is reported and stepped over rather than taking down a + migration run that already succeeded: a database that refuses the ALTER + (most often because the runtime user does not own the tables), a missing + or unrunnable Prisma CLI, a read-only temp directory, or a timeout. + + Returns True when the statement was applied, False when it failed. + """ + logger.info("Applying REPLICA IDENTITY FULL to LiteLLM tables") + try: + with tempfile.TemporaryDirectory(prefix="litellm_replica_identity_") as tmp_dir: + sql_path = Path(tmp_dir) / "replica_identity_full.sql" + sql_path.write_text(REPLICA_IDENTITY_FULL_SQL) + subprocess.run( + [ + prisma_command, + "db", + "execute", + "--file", + str(sql_path), + "--schema", + schema_path, + ], + timeout=60, + check=True, + capture_output=True, + text=True, + env=prisma_env, + ) + except subprocess.CalledProcessError as e: + logger.error( + "Failed to set REPLICA IDENTITY FULL. Logical replication " + "consumers may reject updates to these tables. Grant table " + "ownership to the migration user, or apply " + "`ALTER TABLE ... REPLICA IDENTITY FULL` by hand. Error: %s", + e.stderr, + ) + return False + except subprocess.TimeoutExpired: + logger.error("Timed out setting REPLICA IDENTITY FULL on LiteLLM tables") + return False + except OSError as e: + logger.error( + "Could not run the REPLICA IDENTITY FULL statement. Logical " + "replication consumers may reject updates to these tables. " + "Error: %s", + e, + ) + return False + + logger.info("REPLICA IDENTITY FULL applied to LiteLLM tables") + return True diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py index 369b6561931..af822573322 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/utils.py +++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py @@ -10,6 +10,10 @@ from pathlib import Path from typing import Optional from litellm_proxy_extras._logging import logger +from litellm_proxy_extras.replica_identity import ( + REPLICA_IDENTITY_FULL_ENV_VAR, + apply_replica_identity_full, +) def str_to_bool(value: Optional[str]) -> bool: @@ -676,6 +680,39 @@ class ProxyExtrasDBManager: finally: os.chdir(original_dir) + @staticmethod + def apply_replica_identity_full_if_requested() -> bool: + """ + Re-assert REPLICA IDENTITY FULL on LiteLLM's tables when the operator + opted in via LITELLM_SET_REPLICA_IDENTITY_FULL. + + Prisma leaves new tables at the Postgres default, which logical + replication consumers reject, so the setting has to be re-applied after + every migration run rather than once by hand. + + Returns: + bool: True if the setting was applied, False if it was not + requested or could not be applied. + """ + if not str_to_bool(os.getenv(REPLICA_IDENTITY_FULL_ENV_VAR)): + return False + try: + schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma" + prisma_command = _get_prisma_command() + prisma_env = _get_prisma_env() + except OSError as e: + logger.error( + "Could not resolve the migrations directory for the REPLICA " + "IDENTITY FULL step, skipping it. Error: %s", + e, + ) + return False + return apply_replica_identity_full( + schema_path=schema_path, + prisma_command=prisma_command, + prisma_env=prisma_env, + ) + @staticmethod def setup_database( use_migrate: bool = False, use_v2_resolver: bool = False @@ -694,6 +731,15 @@ class ProxyExtrasDBManager: Returns: bool: True if setup was successful, False otherwise """ + migrated = ProxyExtrasDBManager._run_migrations( + use_migrate=use_migrate, use_v2_resolver=use_v2_resolver + ) + if migrated: + ProxyExtrasDBManager.apply_replica_identity_full_if_requested() + return migrated + + @staticmethod + def _run_migrations(use_migrate: bool, use_v2_resolver: bool) -> bool: if use_v2_resolver: logger.info("Using v2 migration resolver (--use_v2_migration_resolver)") return ProxyExtrasDBManager._setup_database_v2(use_migrate=use_migrate) diff --git a/litellm-rust/ADDING_A_PROVIDER.md b/litellm-rust/ADDING_A_PROVIDER.md index 5f933ec4fa8..857a744e014 100644 --- a/litellm-rust/ADDING_A_PROVIDER.md +++ b/litellm-rust/ADDING_A_PROVIDER.md @@ -1,10 +1,11 @@ # Adding a provider / route to litellm-rust -Three layers, same for every route (see `ocr` and `realtime` as references): +Everything for a route lives in `crates/core/src//`; `crates/core/src/messages` is the reference. A host (the axum gateway, the Python bridge) only calls the route's entrypoint. -1. **Transform contract (pure)** — `crates/core/src//transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) + types in `types.rs`. No network, env, or auth. -2. **Provider config (pure)** — `crates/providers/src///transformation.rs`: implement that trait as a `const __CONFIG`, mirroring the Python provider tree. Add parity unit tests. -3. **HTTP / transport (the host)** — `crates/providers/src/.rs` (e.g. `ocr.rs`, `realtime.rs`): the callable fn (`run_ocr`, `realtime`). It resolves the key, builds the auth header, builds URL + transforms via the config, then does the network call. This is the only layer allowed to do I/O. +1. **Entrypoint** — `mod.rs`: `pub async fn (request) -> CoreResult`, the Rust equivalent of `litellm.()`, plus a `_stream` variant when the route streams. It is the only thing a host touches. +2. **Transform contract** — `transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) with types in `types.rs`. +3. **Provider config** — `crates/core/src/providers///transformation.rs`: implement that trait as a `const __CONFIG`, mirroring the Python provider tree. Add parity unit tests. +4. **Prepare + handler** — `prepare.rs` resolves provider/model, credentials, auth headers, and URL, then transforms the request; `handler.rs` performs the provider call through the shared client in `client.rs` and transforms the response. ## Coding standards @@ -25,4 +26,4 @@ variants of it. The test for a good abstraction is that adding the next provider is a few declarative lines, not a new file of duplicated flow. Only diverge from the base when behavior is genuinely different, and say so explicitly in the PR. -**Calling:** the host invokes the route fn — the Python bridge calls `run_ocr`; the `ai-gateway` server calls `realtime`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`. +**Calling:** hosts invoke the core entrypoint — the Python bridge and the `ai-gateway` route service both call `litellm_core::messages::messages`. Never add a provider handler to `ai-gateway`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`. diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index 398eec4685c..36a5ad5a8f4 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -4,14 +4,30 @@ litellm-rust has exactly THREE crates. A crate is a LAYER, not a route. Routes ( ## Crates -| Crate | Role | Pure / I/O | -|-------|------|------------| -| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure | -| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O | -| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding | +| 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-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-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. | Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge. +## Where a route lives + +A top-level LiteLLM call is a module under `crates/core/src//`, shaped like `messages`: + +``` +core/src/messages/ + mod.rs # pub async fn messages(..) -> CoreResult<..> (+ messages_stream for SSE) + types.rs # request/response types, MessagesRequest + transformation.rs # the provider template trait + prepare.rs # provider resolution, auth headers, URL + handler.rs # the provider call + client.rs # the shared reqwest client +``` + +Handlers never live in `ai-gateway`. `ocr`, `audio_transcription`, and `realtime` are still hosted there from before this rule; they move to `core` as they are touched. + Adding a crate: default to a MODULE. New crate ONLY on a real trigger — separate artifact (binary/cdylib), proc-macro, shared foundation, or publishable standalone. A new provider or route is none of these. Adding a crate fails crates/core/tests/workspace_crate_allowlist.rs until you update its allowlist and this file — intentional. diff --git a/litellm-rust/CLAUDE.md b/litellm-rust/CLAUDE.md index 0659e63df39..fe6ceedbb86 100644 --- a/litellm-rust/CLAUDE.md +++ b/litellm-rust/CLAUDE.md @@ -23,21 +23,34 @@ the base when behavior is genuinely different, and say so explicitly in the PR. ## Crates (exactly three — see AGENTS.md) -`litellm-core` describes work; `litellm-ai-gateway` executes it; `litellm-python-bridge` -exposes it to the Python SDK. A crate is a **layer**, not a route — add modules, not crates. +`litellm-core` **is** the LiteLLM SDK in Rust: it makes the LLM call. +`litellm-ai-gateway` is an HTTP/WebSocket server in front of it, and +`litellm-python-bridge` exposes it to the Python SDK. A crate is a **layer**, not +a route — add modules, not crates. ## Core Boundary -`litellm-core` is the pure translation layer; the `litellm-ai-gateway` host executes work. +`litellm-core` owns the whole call. The Rust equivalent of `litellm.messages()` +is `litellm_core::messages::messages(request).await`: you call it, it does the +provider call, and you get a typed non-streaming response back. Route-level Rust structure mirrors LiteLLM's Python responsibilities: -- `core/src//` owns the route contract, shared types, and provider - template traits. For OCR, this means `core/src/ocr`. +- `core/src//` owns the route end to end: the public entrypoint fn named + after the route in `mod.rs`, the request/response types (`types.rs`), the + provider template trait (`transformation.rs`), the provider/auth/URL + resolution (`prepare.rs`), the HTTP client (`client.rs`), and the handler that + performs the call (`handler.rs`). `core/src/messages` is the reference. - `core/src/providers///transformation.rs` owns the - provider-specific transform. For Mistral OCR, this means - `core/src/providers/mistral/ocr/transformation.rs`. -- Network execution lives in the host crate `ai-gateway` (`ai-gateway/src/io/`), - never inside `core`. + provider-specific transform. For Anthropic Messages, this means + `core/src/providers/anthropic/messages/transformation.rs`. +- Handlers live in `core`, never in a host. `ai-gateway` must not contain a + route handler that talks to a provider; its axum route reads the HTTP request, + picks a deployment, and calls the `core` entrypoint. `python-bridge` marshals + Python objects and calls the same entrypoint. + +Streaming keeps the same shape: the route entrypoint has a `_stream` +variant in `core` that returns the upstream response so a host can splice it to +its own caller; the host still owns no provider logic. Call-hook and lifecycle instrumentation, including phase timing, usage accumulation, and callback payload construction, always lives in `core`. @@ -45,21 +58,31 @@ Hosts feed observed events into core and dispatch the completed payloads through their I/O logger; hosts must not own callback orchestration. Allowed in `core`: -- Pure request transforms -- Pure response transforms -- Pure stream chunk normalization +- The public entrypoint for a top-level LiteLLM call +- Request/response transforms and stream chunk normalization +- Provider resolution, auth header construction, and URL building +- The provider HTTP call itself, through a shared reused client with connect and + request timeouts - Shared data types and validation errors - Deterministic token/cost helper logic Not allowed in `core`: -- Network calls -- Environment variable or secret reads +- Serving HTTP: axum routes, extractors, and transport concerns stay in the host - Filesystem access -- Database or cache access -- Provider SDK signing or auth flows +- Database access +- Config file reading and rollout state - Logging callbacks, spend writes, or custom callbacks - Global mutable runtime state +Env reads in `core` are limited to credential fallback inside a route's +`prepare.rs` (the `env_lookup` closure), mirroring what the Python SDK does when +no key is passed. Everything else config-shaped is resolved by the host and +passed in. + +Routes still hosted in `ai-gateway` (`ocr`, `audio_transcription`, `realtime`) +predate this rule and are being moved into `core` route modules; do not add new +ones there, and prefer moving one when you touch it. + Python owns rollout state and fallback while Rust is being introduced. Rust paths must be off by default until parity tests prove equivalence with Python. A new provider/route may instead be implemented rust-only with no Python @@ -93,10 +116,10 @@ the first PR: - Preserve Python output shape intentionally. If a field is always serialized as `null` for Python parity, leave a short comment explaining that parity choice. -## Host I/O Rules +## Network I/O Rules -These rules apply when adding future crates or modules that execute network I/O, -such as `ai-gateway`, router hosts, or standalone servers: +These rules apply to every module that executes network I/O, whether it is a +`core` route handler or a host such as `ai-gateway`: - Set connect and full-request timeouts. No unbounded waits. - Reuse HTTP clients; do not construct clients per request. diff --git a/litellm-rust/README.md b/litellm-rust/README.md index 1646c90ad76..bcccf93300b 100644 --- a/litellm-rust/README.md +++ b/litellm-rust/README.md @@ -2,18 +2,31 @@ This workspace contains the staged Rust implementation for LiteLLM. -Rust starts as a pure transform core used by the existing Python host. Python -continues to own auth, configuration, network I/O, retries, routing, logging, +`litellm-core` is the LiteLLM SDK in Rust: one entrypoint per top-level call +that makes the LLM call and hands back a typed response, the same shape as +`litellm.messages()` in Python. + +```rust +let response = litellm_core::messages::messages(MessagesRequest { + model: "claude-sonnet-4-5", + body, + api_key: Some(key), + .. +}) +.await?; +``` + +Python continues to own configuration, retries, routing policy, logging, callbacks, spend tracking, and customer plugins until each Rust path has parity coverage and production evidence. ## Crates -| Crate | Role | Pure / I/O | -|-------|------|------------| -| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure | -| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O | -| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding | +| Crate | Role | +|-------|------| +| litellm-core | The SDK. Per-route entrypoints (`messages::messages()`), types, provider transforms (modules under `providers/`), provider resolution, auth, the provider HTTP call, and the router. | +| litellm-ai-gateway | The axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. | +| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. | Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge. @@ -21,16 +34,16 @@ Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm- ```text crates/ - core/ Route contracts, shared pure types, errors, and templates. - src/ocr/ - providers/ Provider-specific pure transforms. - src/mistral/ocr/transformation.rs + core/ The SDK: route modules + provider transforms. + src/messages/ mod.rs (entrypoint), types, transformation, prepare, handler, client + src/providers/anthropic/messages/transformation.rs + ai-gateway/ Axum server + WebSocket hosts; calls core entrypoints. python-bridge/ PyO3 bridge for Python LiteLLM. ``` -The folder shape should follow the Python provider tree: -`providers/src///transformation.rs`. The bridge should expose -one function per top-level route, starting with `ocr(payload)`. +The folder shape follows the Python provider tree: +`core/src/providers///transformation.rs`. The bridge exposes one +function per top-level route, mirroring the core entrypoints. ## Checks diff --git a/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md b/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md index ed44dc4c729..a1860d8a9c9 100644 --- a/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md +++ b/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md @@ -1,6 +1,6 @@ # Provider coding standards (litellm-rust) -Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MISTRAL_OCR_CONFIG`) is the reference; `messages` (`ANTHROPIC_MESSAGES_CONFIG`) is the next port. +Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages` (`core/src/messages`, `ANTHROPIC_MESSAGES_CONFIG`) is the reference: a route is a `core` module with a public entrypoint that makes the call and returns a typed response. ## Provider resolution @@ -16,10 +16,10 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST ## Boundaries -7. Layers never cross: `core` = pure transforms/types (no network, env, secrets, auth, logging, global mutable state); `ai-gateway` = all I/O, auth headers, HTTP/SSE, lifecycle hooks; `python-bridge` = thin PyO3 adapter. +7. Layers never cross: `core` = the call itself (entrypoint, types, transforms, provider resolution, auth headers, provider HTTP, lifecycle hooks); `ai-gateway` = serving HTTP/WS (routing, extractors, auth of *our* callers, streaming to the client); `python-bridge` = thin PyO3 adapter. Hosts call the core entrypoint; they never build a provider request. 8. Generic/route files contain zero provider-specific branches. A provider is one module under `core/src/providers///`; a route is a module, never a new crate. -9. Route entry point stays thin: `()` -> `prepare_*` -> `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing. Handlers validate and delegate; no business logic in them. -10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Env reads happen only at the host/config layer, with the `DEFAULT_*` fallback defined in `constants.rs`. +9. Route entry point stays thin: `core::::()` -> `prepare_*` -> handler (or `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing). Axum handlers validate and delegate to a service that calls the entrypoint; no business logic in them. +10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Config-shaped env reads happen at the host/config layer with the `DEFAULT_*` fallback defined in `constants.rs`; the only env read in `core` is the credential fallback in a route's `prepare.rs`. ## Types and errors @@ -33,7 +33,7 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST 16. Never log request/response bodies, base64 payloads, document contents, or secrets. Truncate and bound any upstream body before it crosses a host boundary. 17. Treat empty/whitespace credentials, URLs, and config values as absent at the host resolution layer. -18. Host I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS. +18. Network I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS. ## Tests and rollout diff --git a/litellm-rust/crates/ai-gateway/AGENTS.md b/litellm-rust/crates/ai-gateway/AGENTS.md index d9e6e1adde5..92567091cd3 100644 --- a/litellm-rust/crates/ai-gateway/AGENTS.md +++ b/litellm-rust/crates/ai-gateway/AGENTS.md @@ -1,7 +1,9 @@ # ai-gateway — folder architecture The Axum server that fronts the Rust gateway. It owns transport + config + auth -only; deployment selection lives in `core::router`, transforms in `core`/`providers`. +only; deployment selection lives in `core::router`, and the LLM call itself +(transforms, auth headers, provider HTTP) lives behind a `core` route entrypoint +such as `litellm_core::messages::messages`. No provider handler lives here. ``` src/ @@ -32,6 +34,11 @@ src/ args; it runs during extraction. Never re-implement the check per route. - **Handlers are thin.** A handler validates and delegates to its `service`. No business logic, no provider calls, no transforms in handlers. +- **Services call `core`, they don't reimplement it.** A `service` picks the + deployment and calls the `core` route entrypoint. Provider resolution, auth + headers, URL building, and the HTTP call are `core`'s job; a service that + builds a provider request itself is a bug (`routes/messages/service.rs` is + the reference). - **State is shared and cheap to clone.** Long-lived handles live behind `Arc` in `state.rs`; read env/config only in `main.rs` when building state. diff --git a/litellm-rust/crates/ai-gateway/README.md b/litellm-rust/crates/ai-gateway/README.md index f913beff6d5..7a6c620ee84 100644 --- a/litellm-rust/crates/ai-gateway/README.md +++ b/litellm-rust/crates/ai-gateway/README.md @@ -8,11 +8,11 @@ dials OpenAI upstream, and splices the two sockets frame-by-frame. `litellm-rust` is exactly three crates (a crate is a **layer**, not a route): -| Crate | Role | Pure / I/O | -|-------|------|------------| -| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. Builds requests/responses; no network. | Pure | -| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under `io/`) plus the Axum server binary (behind the `server` feature). | I/O | -| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding | +| 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-ai-gateway | The Axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. | +| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. | Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge. diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 74808cf1ce6..78af374bf70 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -29,18 +29,6 @@ pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500; #[cfg(feature = "server")] pub(crate) const DEFAULT_PROVIDER: &str = "openai"; -/// Full-request timeout ceiling for Anthropic Messages provider calls, in -/// seconds. Mirrors the Python Anthropic Messages default. The per-request -/// timeout from `litellm_params` still overrides this on the request builder. -pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600; - -/// Connect timeout for Anthropic Messages provider calls, in seconds. -pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10; - -/// Max characters of an upstream error body echoed across the host boundary -/// before truncation, so provider bodies are bounded and data-minimized. -pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256; - pub(crate) const DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10; pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300; @@ -48,10 +36,6 @@ pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300; #[cfg(feature = "server")] pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages"; -/// Provider name used by the Anthropic Messages route when a deployment's -/// provider model does not carry an explicit provider prefix. -pub(crate) const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; - /// Request headers owned by the gateway and never forwarded upstream. #[cfg(feature = "server")] pub(crate) const MESSAGES_HEADERS_NOT_FORWARDED: &[&str] = diff --git a/litellm-rust/crates/ai-gateway/src/io/messages.rs b/litellm-rust/crates/ai-gateway/src/io/messages.rs deleted file mode 100644 index 86170e45678..00000000000 --- a/litellm-rust/crates/ai-gateway/src/io/messages.rs +++ /dev/null @@ -1 +0,0 @@ -pub use crate::messages::{MessagesRequest, messages}; diff --git a/litellm-rust/crates/ai-gateway/src/io/mod.rs b/litellm-rust/crates/ai-gateway/src/io/mod.rs index 6129a808965..cce56dd2121 100644 --- a/litellm-rust/crates/ai-gateway/src/io/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/io/mod.rs @@ -1,5 +1,4 @@ pub mod audio_transcription; -pub mod messages; pub mod ocr; pub mod realtime; pub mod realtime_pool; diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index c44d661c29e..057db6457c4 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -4,7 +4,9 @@ //! without pulling in the HTTP server: //! //! - Call-type modules such as [`ocr`]: provider transforms, lifecycle hooks, -//! and provider I/O. Always available — no feature required. +//! and provider I/O. Always available — no feature required. These predate the +//! rule that a route's entrypoint and handler live in `litellm-core` (see +//! `litellm_core::messages`) and move there as they are touched. //! - [`io`]: compatibility exports and realtime WebSocket splice helpers. //! - The server modules ([`auth`], [`routes`], [`state`]) and anything pulling //! `axum` are gated behind the `server` feature, which the `litellm-ai-gateway` @@ -14,7 +16,6 @@ pub mod audio_transcription; mod client; pub mod io; -pub mod messages; pub mod ocr; /// GIL-activity tracking. Pure (atomics only); shared by the `server` routes and diff --git a/litellm-rust/crates/ai-gateway/src/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/messages/mod.rs deleted file mode 100644 index fd2dd546941..00000000000 --- a/litellm-rust/crates/ai-gateway/src/messages/mod.rs +++ /dev/null @@ -1,49 +0,0 @@ -use litellm_core::CoreResult; -use serde_json::Value; - -mod client; -mod common_utils; -mod handler; -mod prepare; -mod types; - -pub use types::MessagesRequest; - -use handler::{execute_messages_provider_call, execute_messages_provider_stream}; -use prepare::prepare_messages_call; - -pub async fn messages(request: MessagesRequest<'_>) -> CoreResult { - match execute_messages(request, false).await? { - MessagesResponse::Json(body) => Ok(body), - MessagesResponse::Stream(response) => { - drop(response); - Err(litellm_core::CoreError::InvalidResponse( - "non-streaming messages execution returned a stream".to_string(), - )) - } - } -} - -pub(crate) enum MessagesResponse { - Json(Value), - Stream(reqwest::Response), -} - -pub(crate) async fn execute_messages( - request: MessagesRequest<'_>, - stream: bool, -) -> CoreResult { - let prepared = prepare_messages_call(request)?; - if stream { - execute_messages_provider_stream(prepared) - .await - .map(MessagesResponse::Stream) - } else { - execute_messages_provider_call(prepared) - .await - .map(MessagesResponse::Json) - } -} - -#[cfg(test)] -mod tests; diff --git a/litellm-rust/crates/ai-gateway/src/messages/types.rs b/litellm-rust/crates/ai-gateway/src/messages/types.rs deleted file mode 100644 index 848fadb4b02..00000000000 --- a/litellm-rust/crates/ai-gateway/src/messages/types.rs +++ /dev/null @@ -1,24 +0,0 @@ -use std::time::Duration; - -use litellm_core::messages::transformation::AnthropicMessagesProviderConfig; -use serde_json::{Map, Value}; - -pub struct MessagesRequest<'a> { - pub model: &'a str, - pub body: Value, - pub api_key: Option<&'a str>, - pub api_base: Option<&'a str>, - pub custom_llm_provider: Option<&'a str>, - pub extra_headers: Option>, - pub timeout: Option, -} - -pub(crate) struct ProviderMessagesRequest { - pub(crate) provider: String, - pub(crate) model: String, - pub(crate) config: &'static dyn AnthropicMessagesProviderConfig, - pub(crate) url: String, - pub(crate) body: Value, - pub(crate) upstream_headers: Vec<(String, String)>, - pub(crate) timeout: Option, -} diff --git a/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md b/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md index 02c5f18c4f3..3eee43e7a2f 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md +++ b/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md @@ -19,7 +19,10 @@ async fn handle(...) -> impl IntoResponse { ... } When a route has business logic worth testing without axum, put it in a sibling `service` (a file, or a folder if the route grows). The route file stays the **axum surface** (router + handler + any socket/SSE adapter); `service` is plain -Rust with **no axum types**. `realtime/` is the example: +Rust with **no axum types**, and its job is to pick the deployment and call the +`core` route entrypoint (see `messages/service.rs` calling +`litellm_core::messages::messages`). Never build a provider request, resolve a +key, or perform the provider call here. `realtime/` is the older example: ``` realtime/ mod.rs # axum surface: router() + handler + the WS<->events adapter @@ -33,6 +36,8 @@ genuinely gets hard to read. `crate::auth::RequireMasterKey` to its arguments; it runs during extraction. Never re-implement the check per route. - **Handlers contain no business logic; `service` contains no axum types.** +- **No provider handlers in this crate.** Transforms, auth headers, and the + provider HTTP call live in `core/src//`. - A route owns its paths in its own `router()`; `mod.rs` only merges. - Cross-cutting concerns (logging, CORS, timeouts) → Tower layers in `mod.rs`, not duplicated in handlers. diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs index 75ed26e5be8..5f4c5fe8de4 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -1,12 +1,12 @@ use std::sync::Arc; +use litellm_core::constants::ANTHROPIC_MESSAGES_PROVIDER; +use litellm_core::messages::types::MessagesRequest; +use litellm_core::messages::{messages, messages_stream}; use litellm_core::router::Router; use litellm_core::{CoreError, CoreResult}; use serde_json::{Map, Value}; -use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; -use crate::messages::{MessagesRequest, execute_messages}; - pub(crate) enum MessagesResponse { Json(Value), Stream(reqwest::Response), @@ -52,13 +52,14 @@ pub async fn run( extra_headers, timeout: None, }; - let stream = request.body.get("stream").and_then(Value::as_bool) == Some(true); - execute_messages(request, stream) - .await - .map(|response| match response { - crate::messages::MessagesResponse::Json(body) => MessagesResponse::Json(body), - crate::messages::MessagesResponse::Stream(upstream) => { - MessagesResponse::Stream(upstream) - } + if request.body.get("stream").and_then(Value::as_bool) == Some(true) { + return messages_stream(request).await.map(MessagesResponse::Stream); + } + + let response = messages(request).await?; + serde_json::to_value(response) + .map(MessagesResponse::Json) + .map_err(|err| { + CoreError::InvalidResponse(format!("failed to serialize messages response: {err}")) }) } diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 8740dccaf01..aee8b4937ef 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -1,3 +1,7 @@ -litellm-core is the PURE translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. No network, no I/O, no env reads. +litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src//` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back. -Routes (ocr, realtime) and providers (mistral, openai) are modules, not crates. +A route module owns everything the call needs: types, the provider template trait, provider transforms (under `providers/`), provider/auth/URL resolution, and the handler that performs the HTTP call. Handlers belong here, never in a host crate. + +Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback dispatch. Env reads are limited to credential fallback in a route's `prepare.rs`. + +Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates. diff --git a/litellm-rust/crates/core/CLAUDE.md b/litellm-rust/crates/core/CLAUDE.md index 20873878967..5d36305ded5 100644 --- a/litellm-rust/crates/core/CLAUDE.md +++ b/litellm-rust/crates/core/CLAUDE.md @@ -4,20 +4,28 @@ Rules for `litellm-rust/crates/core`. ## Responsibility -`core` owns shared data types, typed errors, and deterministic helper contracts. -It must stay pure and host-independent. +`core` is the LiteLLM SDK in Rust: it makes the LLM call. Every top-level +LiteLLM call has a public entrypoint here, named after the route +(`messages::messages()` is the Rust equivalent of `litellm.messages()`), and +calling it returns a typed non-streaming response. Allowed: +- The public entrypoint for a route, plus its `_stream` variant when the + route supports streaming. +- Provider resolution, auth header construction, URL building, and the provider + HTTP call (shared reused client, connect + request timeouts). - Shared request/response structs. - Typed errors with stable, non-sensitive messages. - Deterministic validation helpers. - Serialization helpers that intentionally mirror Python output shape. - Route templates that match Python base config responsibilities, such as - `ocr::transformation::OcrProviderConfig`. + `messages::transformation::AnthropicMessagesProviderConfig`. Not allowed: -- Network, filesystem, database, cache, or environment access. -- Secret reads or auth/header construction. +- Serving HTTP: axum routers, extractors, and other transport concerns. +- Filesystem, database, or cache access. +- Config file reading or rollout state; the host resolves those and passes them + in. Env reads are limited to credential fallback in a route's `prepare.rs`. - Logging callbacks, tracing spans, spend writes, or customer callbacks. - Provider-specific branching that belongs in `providers`. - Panics for user/provider-controlled input. @@ -33,10 +41,21 @@ typed field on a struct, not a raw string threaded through the API. ## Structure -Use route names directly under `src/`: `ocr`, future `messages`, +Use route names directly under `src/`: `messages`, `ocr`, future `chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not invent broad names like `engine` for route contracts. +`src/messages` is the reference shape for a route module: + +``` +mod.rs pub async fn messages(..) (+ messages_stream) +types.rs request/response types +transformation.rs the provider template trait +prepare.rs provider resolution, auth headers, URL +handler.rs the provider call +client.rs the shared reqwest client +``` + ## Parity Rules - Every shared type used by a provider transform needs unit tests for diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 65c6db7412c..ab8050734f2 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -7,6 +7,7 @@ repository.workspace = true [dependencies] rand.workspace = true +reqwest.workspace = true serde.workspace = true serde_json.workspace = true thiserror.workspace = true @@ -30,5 +31,4 @@ bedrock-auth = [ ] [dev-dependencies] -reqwest.workspace = true tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index 5826a5bc9c1..caada1d98b0 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -1,3 +1,19 @@ pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com"; pub const OPENAI_RESPONSES_DEFAULT_API_BASE: &str = "https://api.openai.com/v1"; pub const OPENAI_RESPONSES_PATH: &str = "/responses"; + +/// Full-request timeout ceiling for Anthropic Messages provider calls, in +/// seconds. Mirrors the Python Anthropic Messages default. The per-request +/// timeout from the caller still overrides this on the request builder. +pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600; + +/// Connect timeout for Anthropic Messages provider calls, in seconds. +pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10; + +/// Max characters of an upstream error body echoed across the call boundary +/// before truncation, so provider bodies are bounded and data-minimized. +pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256; + +/// Provider name used for Anthropic Messages when a deployment's provider model +/// does not carry an explicit provider prefix. +pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; diff --git a/litellm-rust/crates/ai-gateway/src/messages/client.rs b/litellm-rust/crates/core/src/messages/client.rs similarity index 100% rename from litellm-rust/crates/ai-gateway/src/messages/client.rs rename to litellm-rust/crates/core/src/messages/client.rs diff --git a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs similarity index 83% rename from litellm-rust/crates/ai-gateway/src/messages/common_utils.rs rename to litellm-rust/crates/core/src/messages/common_utils.rs index 68ecc3f17c1..9dcfcaa71e3 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -1,11 +1,11 @@ -use litellm_core::CoreResult; -use litellm_core::error::{CoreError, json_type_name}; -use litellm_core::messages::transformation::AnthropicMessagesProviderConfig; -use litellm_core::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG; -use litellm_core::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG; use serde_json::{Map, Value}; use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS; +use crate::error::{CoreError, CoreResult, json_type_name}; +use crate::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG; +use crate::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG; + +use super::transformation::AnthropicMessagesProviderConfig; pub(super) fn truncate_error_body(body: &str) -> String { if body.chars().count() <= MESSAGES_ERROR_BODY_MAX_CHARS { diff --git a/litellm-rust/crates/ai-gateway/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs similarity index 84% rename from litellm-rust/crates/ai-gateway/src/messages/handler.rs rename to litellm-rust/crates/core/src/messages/handler.rs index 90c12367f50..1c895f66eba 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -1,15 +1,13 @@ -use litellm_core::CoreResult; -use litellm_core::error::CoreError; -use serde_json::Value; +use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; +use crate::error::{CoreError, CoreResult}; use super::client::http_client; use super::common_utils::truncate_error_body; -use super::types::ProviderMessagesRequest; -use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; +use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest}; pub(super) async fn execute_messages_provider_call( request: ProviderMessagesRequest, -) -> CoreResult { +) -> CoreResult { let mut request_builder = http_client().post(&request.url).json(&request.body); for (key, value) in &request.upstream_headers { request_builder = request_builder.header(key, value); @@ -39,12 +37,7 @@ pub(super) async fn execute_messages_provider_call( let response = serde_json::from_str(&text).map_err(|err| { CoreError::InvalidResponse(format!("invalid messages response JSON: {err}")) })?; - let transformed = request - .config - .transform_response(&request.model, response)?; - serde_json::to_value(transformed).map_err(|err| { - CoreError::InvalidResponse(format!("failed to serialize messages response: {err}")) - }) + request.config.transform_response(&request.model, response) } pub(super) async fn execute_messages_provider_stream( diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index ec2fbb969a6..acb36d89daf 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -1,2 +1,32 @@ +//! The Anthropic Messages call, the Rust equivalent of Python's +//! `litellm.messages()`. +//! +//! [`messages`] is the top-level entrypoint: give it a model, a body, and +//! credentials, and it resolves the provider, transforms the request, calls the +//! provider, and returns a typed non-streaming response. [`messages_stream`] +//! is the streaming variant; it hands the raw upstream response back so a host +//! can splice the event stream to its own caller. + +mod client; +mod common_utils; +mod handler; +mod prepare; pub mod transformation; pub mod types; + +use crate::error::CoreResult; + +use handler::{execute_messages_provider_call, execute_messages_provider_stream}; +use prepare::prepare_messages_call; +use types::{AnthropicMessagesResponse, MessagesRequest}; + +pub async fn messages(request: MessagesRequest<'_>) -> CoreResult { + execute_messages_provider_call(prepare_messages_call(request)?).await +} + +pub async fn messages_stream(request: MessagesRequest<'_>) -> CoreResult { + execute_messages_provider_stream(prepare_messages_call(request)?).await +} + +#[cfg(test)] +mod tests; diff --git a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs similarity index 92% rename from litellm-rust/crates/ai-gateway/src/messages/prepare.rs rename to litellm-rust/crates/core/src/messages/prepare.rs index 9a027490eb6..94b5b1eaed7 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -1,9 +1,8 @@ -use litellm_core::CoreError; -use litellm_core::CoreResult; -use litellm_core::messages::transformation::MessagesAuthStrategy; -use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; +use crate::error::{CoreError, CoreResult}; +use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers}; +use super::transformation::MessagesAuthStrategy; use super::types::{MessagesRequest, ProviderMessagesRequest}; pub(super) fn prepare_messages_call( diff --git a/litellm-rust/crates/ai-gateway/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs similarity index 97% rename from litellm-rust/crates/ai-gateway/src/messages/tests.rs rename to litellm-rust/crates/core/src/messages/tests.rs index 23a53e98045..9fc1763683b 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -1,14 +1,16 @@ use std::time::Duration; -use litellm_core::error::CoreError; use serde_json::{Map, Value, json}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; +use crate::error::CoreError; + use super::common_utils::{ has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body, }; -use super::{MessagesRequest, messages}; +use super::messages; +use super::types::MessagesRequest; async fn read_http_request(socket: &mut TcpStream) -> String { let mut request = Vec::new(); @@ -152,8 +154,8 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through() .await .expect("messages request succeeds"); - assert_eq!(response["content"][0]["text"], "hi"); - assert_eq!(response["stop_reason"], "end_turn"); + assert_eq!(response.content[0]["text"], "hi"); + assert_eq!(response.stop_reason.as_deref(), Some("end_turn")); let request = server.await.expect("server task completes"); let (head, body) = request.split_once("\r\n\r\n").expect("has body"); @@ -208,8 +210,8 @@ async fn messages_round_trip_builds_native_anthropic_request() { .await .expect("messages request succeeds"); - assert_eq!(response["content"][0]["text"], "hi"); - assert_eq!(response["stop_reason"], "end_turn"); + assert_eq!(response.content[0]["text"], "hi"); + assert_eq!(response.stop_reason.as_deref(), Some("end_turn")); let request = server.await.expect("server task completes"); let (head, _) = request.split_once("\r\n\r\n").expect("has body"); diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index 11fe17ea40f..b9f807c29fd 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -1,6 +1,30 @@ +use std::time::Duration; + use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use super::transformation::AnthropicMessagesProviderConfig; + +pub struct MessagesRequest<'a> { + pub model: &'a str, + pub body: Value, + pub api_key: Option<&'a str>, + pub api_base: Option<&'a str>, + pub custom_llm_provider: Option<&'a str>, + pub extra_headers: Option>, + pub timeout: Option, +} + +pub(super) struct ProviderMessagesRequest { + pub(super) provider: String, + pub(super) model: String, + pub(super) config: &'static dyn AnthropicMessagesProviderConfig, + pub(super) url: String, + pub(super) body: Value, + pub(super) upstream_headers: Vec<(String, String)>, + pub(super) timeout: Option, +} + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[serde(untagged)] pub enum SystemPrompt { diff --git a/litellm-rust/crates/python-bridge/AGENTS.md b/litellm-rust/crates/python-bridge/AGENTS.md index d6d3d90e6ab..ad3cddfa5fd 100644 --- a/litellm-rust/crates/python-bridge/AGENTS.md +++ b/litellm-rust/crates/python-bridge/AGENTS.md @@ -1,3 +1,3 @@ -litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over litellm-ai-gateway. +litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over the litellm-core route entrypoints (e.g. `litellm_core::messages::messages`). -Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call into litellm-ai-gateway. +Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call the core entrypoint. diff --git a/litellm-rust/crates/python-bridge/CLAUDE.md b/litellm-rust/crates/python-bridge/CLAUDE.md index e5d021ec25b..3ce8b8c639a 100644 --- a/litellm-rust/crates/python-bridge/CLAUDE.md +++ b/litellm-rust/crates/python-bridge/CLAUDE.md @@ -11,11 +11,11 @@ Python-compatible dictionaries. ## Bridge Shape - Prefer one stable method per top-level LiteLLM route, for example - `ocr(payload)`. + `messages(...)`, calling the matching `litellm-core` entrypoint. - Do not add one exported PyO3 function per provider helper unless there is a measured reason. -- Provider dispatch belongs in Rust route modules such as - `litellm_providers::ocr`, not in this PyO3 crate. +- Provider dispatch belongs in the `litellm-core` route module (e.g. + `litellm_core::messages`), not in this PyO3 crate. - Python owns rollout state and fallback. Rust should return errors; Python decides whether to raise or fall back. For a rust-only provider/route (no Python reference), the Python side is a thin dispatch that calls Rust and diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index ee9bdd0b81f..f0cc26a0cca 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -4,10 +4,11 @@ use std::time::Duration; use litellm_ai_gateway::io::audio_transcription::{ AudioTranscriptionRequest, audio_transcription as run_audio_transcription, }; -use litellm_ai_gateway::io::messages::{MessagesRequest, messages as run_messages}; use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr}; use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use litellm_core::error::CoreError; +use litellm_core::messages::messages as run_messages; +use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest}; use pyo3::exceptions::{PyRuntimeError, PyValueError}; use pyo3::prelude::*; use pyo3::types::{PyAny, PyDict}; @@ -35,6 +36,15 @@ fn json_to_py(py: Python<'_>, value: Value) -> PyResult> { Ok(json.call_method1("loads", (encoded,))?.unbind()) } +fn messages_response_to_py( + py: Python<'_>, + response: AnthropicMessagesResponse, +) -> PyResult> { + let value = + serde_json::to_value(response).map_err(|err| PyValueError::new_err(err.to_string()))?; + json_to_py(py, value) +} + fn core_error_to_pyerr(err: CoreError) -> PyErr { match err { CoreError::Auth(message) => PyValueError::new_err(message), @@ -382,7 +392,7 @@ fn messages( }); match result { - Ok(value) => json_to_py(py, value), + Ok(response) => messages_response_to_py(py, response), Err(err) => Err(core_error_to_pyerr(err)), } } @@ -404,7 +414,7 @@ fn amessages( marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?; pyo3_async_runtimes::tokio::future_into_py(py, async move { - let value = run_messages(MessagesRequest { + let response = run_messages(MessagesRequest { model: &model, body, api_key: api_key.as_deref(), @@ -416,7 +426,7 @@ fn amessages( .await .map_err(core_error_to_pyerr)?; - Python::attach(|py| json_to_py(py, value)) + Python::attach(|py| messages_response_to_py(py, response)) }) } diff --git a/litellm/_redis.py b/litellm/_redis.py index fe5c5cdabe9..9e3b247f577 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -61,23 +61,51 @@ def _get_redis_kwargs(): return available_args -def _get_redis_url_kwargs(client=None): +def _init_arg_names(cls: type) -> frozenset[str]: + """Every ``__init__`` parameter accepted anywhere in a class's MRO. + + Keyword-only parameters are included, and the MRO is walked because redis-py splits a + connection's parameters between ``AbstractConnection`` and its concrete subclasses. + """ + return frozenset( + name + for klass in inspect.getmro(cls) + if klass is not object + for spec in (inspect.getfullargspec(klass.__init__),) + for name in spec.args + spec.kwonlyargs + ) + + +def _get_redis_url_kwargs(client: Optional[type] = None) -> tuple[str, ...]: + """Connection kwargs that redis-py forwards from ``from_url`` down to the connection. + + ``from_url`` is declared as ``(cls, url, **kwargs)``, so introspecting it yields no + connection kwargs at all. What it really does is hand its kwargs to the connection + class, so that class's signature is the allowlist. + + Taking the client's signature instead would be wrong in both directions: it omits + nothing useful, but it admits client-only parameters such as + ``single_connection_client`` and ``auto_close_connection_pool``, plus the ``ssl_*`` + family that only ``SSLConnection`` accepts. Those reach ``AbstractConnection`` and + raise ``TypeError`` the first time a connection is created. TLS on a url config is + selected by the ``rediss://`` scheme, which picks ``SSLConnection`` on its own. + """ if client is None: - client = redis.Redis.from_url - arg_spec = inspect.getfullargspec(redis.Redis.from_url) + client = redis.Redis + connection_cls = async_redis.Connection if client is async_redis.Redis else redis.Connection + + exclude_args = frozenset( + { + "self", + "connection_pool", + "retry", + } + ) # Only allow primitive arguments - exclude_args = { - "self", - "connection_pool", - "retry", - } + include_args = ("url", "max_connections") - include_args = ["url"] - - available_args = [x for x in arg_spec.args if x not in exclude_args] + include_args - - return available_args + return tuple(x for x in _init_arg_names(connection_cls) if x not in exclude_args) + include_args def _get_redis_cluster_kwargs(client=None): @@ -614,7 +642,7 @@ def get_redis_async_client( if "url" in redis_kwargs and redis_kwargs["url"] is not None: if connection_pool is not None: return async_redis.Redis(connection_pool=connection_pool) - args = _get_redis_url_kwargs(client=async_redis.Redis.from_url) + args = _get_redis_url_kwargs(client=async_redis.Redis) url_kwargs = {} for arg in redis_kwargs: if arg in args: @@ -662,10 +690,10 @@ def get_redis_connection_pool( return None if "url" in redis_kwargs and redis_kwargs["url"] is not None: - pool_kwargs = { - "timeout": REDIS_CONNECTION_POOL_TIMEOUT, - "url": redis_kwargs["url"], - } + allowed_args = _get_redis_url_kwargs(client=async_redis.Redis) + pool_kwargs = {k: v for k, v in redis_kwargs.items() if k in allowed_args and k != "max_connections"} + pool_kwargs["timeout"] = REDIS_CONNECTION_POOL_TIMEOUT + pool_kwargs["url"] = redis_kwargs["url"] if "max_connections" in redis_kwargs: try: pool_kwargs["max_connections"] = int(redis_kwargs["max_connections"]) diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index 4c23ecfed54..f04edf2579b 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -568,6 +568,7 @@ def _build_streaming_logging_obj( logging_obj.custom_llm_provider = "a2a_agent" logging_obj.model_call_details["model"] = model logging_obj.model_call_details["custom_llm_provider"] = "a2a_agent" + logging_obj.model_call_details["call_type"] = logging_obj.call_type if agent_id: logging_obj.model_call_details["agent_id"] = agent_id diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 2fcb8455e90..eef4cf8d87f 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -1,5 +1,6 @@ import json -from typing import Any, Iterator, List, Literal, Optional, Tuple +from dataclasses import dataclass +from typing import Any, Iterable, Iterator, List, Literal, Optional, Tuple import litellm from litellm._logging import verbose_logger @@ -24,20 +25,20 @@ async def calculate_batch_cost_and_usage( deployment-specific pricing (e.g. input_cost_per_token_batches) is used instead of the global cost map. """ - batch_cost = _batch_cost_calculator( + if ( + custom_llm_provider == "vertex_ai" + and model_name + and getattr(litellm, "disable_vertex_batch_output_transformation", False) + ): + batch_cost, batch_usage = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name) + return batch_cost, batch_usage, [model_name] + + return _aggregate_batch_cost_usage_models( + entries=file_content_dictionary, custom_llm_provider=custom_llm_provider, - file_content_dictionary=file_content_dictionary, model_name=model_name, model_info=model_info, ) - batch_usage = _get_batch_job_total_usage_from_file_content( - file_content_dictionary=file_content_dictionary, - custom_llm_provider=custom_llm_provider, - model_name=model_name, - ) - batch_models = _get_batch_models_from_file_content(file_content_dictionary, model_name, custom_llm_provider) - - return batch_cost, batch_usage, batch_models async def _handle_completed_batch( @@ -46,7 +47,9 @@ async def _handle_completed_batch( model_name: Optional[str] = None, litellm_params: Optional[dict] = None, ) -> Tuple[float, Usage, List[str]]: - """Helper function to process a completed batch and handle logging + """Fetch a completed batch's output file and aggregate its cost, usage, and + models in a single pass over the JSONL lines, so the parsed file content is + never materialized in memory. Args: batch: The batch object @@ -54,75 +57,109 @@ async def _handle_completed_batch( model_name: Optional model name litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.) """ - # Get batch results - file_content_dictionary = await _get_batch_output_file_content_as_dictionary( - batch, custom_llm_provider, litellm_params=litellm_params - ) + file_content = await _fetch_batch_output_file_content(batch, custom_llm_provider, litellm_params=litellm_params) - # Calculate costs and usage - batch_cost = _batch_cost_calculator( - custom_llm_provider=custom_llm_provider, - file_content_dictionary=file_content_dictionary, - model_name=model_name, - ) - batch_usage = _get_batch_job_total_usage_from_file_content( - file_content_dictionary=file_content_dictionary, - custom_llm_provider=custom_llm_provider, - model_name=model_name, - ) - - batch_models = _get_batch_models_from_file_content(file_content_dictionary, model_name, custom_llm_provider) - - return batch_cost, batch_usage, batch_models - - -def _get_batch_models_from_file_content( - file_content_dictionary: List[dict], - model_name: Optional[str] = None, - custom_llm_provider: str = "openai", -) -> List[str]: - """ - Get the models from the file content - """ - if model_name: - return [model_name] - batch_models = [] - for _item in file_content_dictionary: - if _batch_response_was_successful(_item, custom_llm_provider): - _response_body = _get_response_from_batch_job_output_file(_item, custom_llm_provider) - _model = _response_body.get("model") - if _model: - batch_models.append(_model) - return batch_models - - -def _batch_cost_calculator( - file_content_dictionary: List[dict], - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai", - model_name: Optional[str] = None, - model_info: Optional[ModelInfo] = None, -) -> float: - """ - Calculate the cost of a batch based on the output file id - """ if ( custom_llm_provider == "vertex_ai" and model_name and getattr(litellm, "disable_vertex_batch_output_transformation", False) ): - batch_cost, _ = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name) - verbose_logger.debug("vertex_ai_total_cost=%s", batch_cost) - return batch_cost + batch_cost, batch_usage = calculate_vertex_ai_batch_cost_and_usage( + _get_file_content_as_dictionary(file_content), model_name + ) + return batch_cost, batch_usage, [model_name] - # For other providers, use the existing logic - total_cost = _get_batch_job_cost_from_file_content( - file_content_dictionary=file_content_dictionary, + return _aggregate_batch_cost_usage_models( + entries=_iter_batch_input_entries(file_content), custom_llm_provider=custom_llm_provider, model_name=model_name, - model_info=model_info, ) - verbose_logger.debug("total_cost=%s", total_cost) - return total_cost + + +@dataclass(frozen=True, slots=True) +class _BatchOutputLineStats: + cost: float + prompt_tokens: int + completion_tokens: int + total_tokens: int + cache_read_tokens: int + cache_creation_tokens: int + model: Optional[str] + + +def _iter_successful_output_line_stats( + entries: Iterable[dict], + custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"], + model_name: Optional[str], + model_info: Optional[ModelInfo], +) -> Iterator[_BatchOutputLineStats]: + from litellm.cost_calculator import batch_cost_calculator + + for entry in entries: + if not _batch_response_was_successful(entry, custom_llm_provider): + continue + response_body = _get_response_from_batch_job_output_file(entry, custom_llm_provider) + usage = _get_batch_job_usage_from_response_body(response_body, custom_llm_provider) + prompt_details = _parse_prompt_tokens_details(usage) + raw_model = response_body.get("model") + response_model = raw_model if isinstance(raw_model, str) and raw_model else None + if model_info is not None or custom_llm_provider in ("anthropic", "bedrock"): + if custom_llm_provider == "bedrock" and model_name: + cost_model = model_name + else: + cost_model = response_model or model_name or "" + prompt_cost, completion_cost = batch_cost_calculator( + usage=usage, + model=cost_model, + custom_llm_provider=custom_llm_provider, + model_info=model_info, + ) + line_cost = prompt_cost + completion_cost + else: + line_cost = litellm.completion_cost( + completion_response=response_body, + custom_llm_provider=custom_llm_provider, + call_type=CallTypes.aretrieve_batch.value, + ) + yield _BatchOutputLineStats( + cost=line_cost, + prompt_tokens=usage.prompt_tokens, + completion_tokens=usage.completion_tokens, + total_tokens=usage.total_tokens, + cache_read_tokens=prompt_details["cache_hit_tokens"], + cache_creation_tokens=prompt_details["cache_creation_tokens"], + model=response_model, + ) + + +def _aggregate_batch_cost_usage_models( + entries: Iterable[dict], + custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"], + model_name: Optional[str] = None, + model_info: Optional[ModelInfo] = None, +) -> Tuple[float, Usage, List[str]]: + """Aggregate cost, usage, and models from batch output entries in a single + pass, holding one small stats record per line instead of the parsed file.""" + line_stats = tuple(_iter_successful_output_line_stats(entries, custom_llm_provider, model_name, model_info)) + + cache_token_params = { + key: tokens + for key, tokens in ( + ("cache_read_input_tokens", sum(stats.cache_read_tokens for stats in line_stats)), + ("cache_creation_input_tokens", sum(stats.cache_creation_tokens for stats in line_stats)), + ) + if tokens > 0 + } + batch_usage = Usage( + total_tokens=sum(stats.total_tokens for stats in line_stats), + prompt_tokens=sum(stats.prompt_tokens for stats in line_stats), + completion_tokens=sum(stats.completion_tokens for stats in line_stats), + **cache_token_params, + ) + batch_models = [model_name] if model_name else [stats.model for stats in line_stats if stats.model] + total_cost = sum((stats.cost for stats in line_stats), 0.0) + verbose_logger.debug("batch output aggregate: cost=%s usage=%s models=%s", total_cost, batch_usage, batch_models) + return total_cost, batch_usage, batch_models def calculate_vertex_ai_batch_cost_and_usage( @@ -193,13 +230,13 @@ def calculate_vertex_ai_batch_cost_and_usage( ) -async def _get_batch_output_file_content_as_dictionary( +async def _fetch_batch_output_file_content( batch: Batch, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai", litellm_params: Optional[dict] = None, -) -> List[dict]: +) -> bytes: """ - Get the batch output file content as a list of dictionaries + Fetch the batch output file and return its raw JSONL bytes Args: batch: The batch object @@ -212,9 +249,6 @@ async def _get_batch_output_file_content_as_dictionary( _is_base64_encoded_unified_file_id, ) - if custom_llm_provider == "vertex_ai": - raise ValueError("Vertex AI does not support file content retrieval") - if batch.output_file_id is None: raise ValueError("Output file id is None cannot retrieve file content") @@ -240,7 +274,7 @@ async def _get_batch_output_file_content_as_dictionary( file_content_kwargs.update(credentials) _file_content = await afile_content(**file_content_kwargs) # type: ignore[reportArgumentType] - return _get_file_content_as_dictionary(_file_content.content) + return _file_content.content def _extract_file_access_credentials(litellm_params: Optional[dict]) -> dict: @@ -270,6 +304,8 @@ def _extract_file_access_credentials(litellm_params: Optional[dict]) -> dict: "vertex_project", "vertex_location", "vertex_credentials", + "gcs_bucket_name", + "bucket_name", "timeout", "max_retries", ] @@ -284,17 +320,7 @@ def _get_file_content_as_dictionary(file_content: bytes) -> List[dict]: """ Get the file content as a list of dictionaries from JSON Lines format """ - try: - _file_content_str = file_content.decode("utf-8") - # Split by newlines and parse each line as a separate JSON object - json_objects = [] - for line in _file_content_str.strip().split("\n"): - if line: # Skip empty lines - json_objects.append(json.loads(line)) - verbose_logger.debug("json_objects=%s", json.dumps(json_objects, indent=4)) - return json_objects - except Exception as e: - raise e + return list(_iter_batch_input_entries(file_content)) def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]: @@ -361,101 +387,6 @@ def _count_entry_tokens( return 0 -def _get_batch_job_cost_from_file_content( - file_content_dictionary: List[dict], - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai", - model_name: Optional[str] = None, - model_info: Optional[ModelInfo] = None, -) -> float: - """ - Get the cost of a batch job from the file content - """ - from litellm.cost_calculator import batch_cost_calculator - - try: - total_cost: float = 0.0 - # parse the file content as json - verbose_logger.debug("file_content_dictionary=%s", json.dumps(file_content_dictionary, indent=4)) - for _item in file_content_dictionary: - if _batch_response_was_successful(_item, custom_llm_provider): - _response_body = _get_response_from_batch_job_output_file(_item, custom_llm_provider) - if model_info is not None or custom_llm_provider in ("anthropic", "bedrock"): - usage = _get_batch_job_usage_from_response_body(_response_body, custom_llm_provider) - # Bedrock batch output lines report a short internal model id - # (e.g. "claude-sonnet-4-6") that is not in the cost map; use the - # deployment model name for pricing when available. - if custom_llm_provider == "bedrock" and model_name: - model = model_name - else: - model = _response_body.get("model") or model_name or "" - prompt_cost, completion_cost = batch_cost_calculator( - usage=usage, - model=model, - custom_llm_provider=custom_llm_provider, - model_info=model_info, - ) - total_cost += prompt_cost + completion_cost - else: - total_cost += litellm.completion_cost( - completion_response=_response_body, - custom_llm_provider=custom_llm_provider, - call_type=CallTypes.aretrieve_batch.value, - ) - verbose_logger.debug("total_cost=%s", total_cost) - return total_cost - except Exception as e: - verbose_logger.error("error in _get_batch_job_cost_from_file_content", e) - raise e - - -def _get_batch_job_total_usage_from_file_content( - file_content_dictionary: List[dict], - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai", - model_name: Optional[str] = None, -) -> Usage: - """ - Get the tokens of a batch job from the file content - """ - if ( - custom_llm_provider == "vertex_ai" - and model_name - and getattr(litellm, "disable_vertex_batch_output_transformation", False) - ): - _, batch_usage = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name) - return batch_usage - - # For other providers, use the existing logic - total_tokens: int = 0 - prompt_tokens: int = 0 - completion_tokens: int = 0 - cache_read_tokens: int = 0 - cache_creation_tokens: int = 0 - for _item in file_content_dictionary: - if _batch_response_was_successful(_item, custom_llm_provider): - _response_body = _get_response_from_batch_job_output_file(_item, custom_llm_provider) - usage: Usage = _get_batch_job_usage_from_response_body(_response_body, custom_llm_provider) - total_tokens += usage.total_tokens - prompt_tokens += usage.prompt_tokens - completion_tokens += usage.completion_tokens - prompt_details = _parse_prompt_tokens_details(usage) - cache_read_tokens += prompt_details["cache_hit_tokens"] - cache_creation_tokens += prompt_details["cache_creation_tokens"] - cache_token_params = { - key: tokens - for key, tokens in ( - ("cache_read_input_tokens", cache_read_tokens), - ("cache_creation_input_tokens", cache_creation_tokens), - ) - if tokens > 0 - } - return Usage( - total_tokens=total_tokens, - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - **cache_token_params, - ) - - def _count_prompt_or_input_tokens(model: str, value: Any) -> int: """Token-count a ``prompt`` / ``input`` field that the OpenAI batch schema allows in four shapes: diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index dd1c152a421..9e0f022262b 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -17,7 +17,8 @@ import json import time from collections.abc import Awaitable, Callable, Sequence from datetime import timedelta -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast +from contextvars import ContextVar +from typing import TYPE_CHECKING, Any, List, Optional, Tuple, TypeVar, Union, cast import litellm from litellm._logging import print_verbose, verbose_logger @@ -168,24 +169,97 @@ class RedisCircuitBreaker: self._state = self.CLOSED +_RedisCallResult = TypeVar("_RedisCallResult") + + +_swallowed_redis_failures: ContextVar[int] = ContextVar("litellm_swallowed_redis_failures", default=0) + + +@functools.lru_cache(maxsize=1) +def _redis_health_error_types() -> tuple[type, ...]: + """Exception types that mean the Redis backend itself is unhealthy. + + Command and data errors say nothing about connectivity: an INCR against a non-numeric + value or an undecodable cached entry is a request problem, and counting those would let + a caller trip the shared breaker on demand, dropping rate limiting to per-process + counters that spreading traffic across replicas can outrun. + + Imported lazily because this module is reachable from a base ``import litellm`` while + redis is not a base dependency. + """ + from redis.exceptions import BusyLoadingError, ClusterDownError + from redis.exceptions import ConnectionError as RedisConnectionError + from redis.exceptions import TimeoutError as RedisTimeoutError + + return (RedisConnectionError, RedisTimeoutError, BusyLoadingError, ClusterDownError, OSError, asyncio.TimeoutError) + + +def _is_redis_health_failure(exc: BaseException) -> bool: + """True when ``exc`` indicates Redis is unreachable rather than the request being bad.""" + try: + return isinstance(exc, _redis_health_error_types()) + except ImportError: + return True + + +def _record_swallowed_redis_failure(breaker: RedisCircuitBreaker, exc: BaseException) -> None: + """Record a Redis failure that the calling method is about to swallow. + + The marker is a ContextVar rather than a counter on the breaker because breakers are + shared by every concurrent caller. A plain shared counter cannot tell "my call failed" + from "some other in-flight call failed", so a success overlapping someone else's + failure would be discarded and a Redis that is answering would still be evicted. + asyncio gives each task its own copy of the context, so this is per-call. + """ + if not _is_redis_health_failure(exc): + return + breaker.record_failure() + _swallowed_redis_failures.set(_swallowed_redis_failures.get() + 1) + + +async def _run_under_circuit_breaker( + breaker: RedisCircuitBreaker, + name: str, + call: Callable[[], Awaitable[_RedisCallResult]], +) -> _RedisCallResult: + """Run one Redis coroutine under a circuit breaker. + + Shared by the method decorator and the Lua script executor so both feed the same + health signal. Success is recorded only when nothing failed while ``call`` ran, + because several Redis methods catch their own connection errors and return a default. + """ + if breaker.is_open(): + raise Exception(f"Redis circuit breaker is open — skipping {name}") + swallowed_before = _swallowed_redis_failures.get() + try: + result = await call() + except Exception as e: + if _is_redis_health_failure(e): + breaker.record_failure() + raise + if _swallowed_redis_failures.get() == swallowed_before: + breaker.record_success() + return result + + def _redis_circuit_breaker_guard(method): # type: ignore """ Decorator for RedisCache async methods. Checks the circuit breaker before each call; records success/failure after. Does not apply to ping/disconnect/test_connection (health/teardown must always run). + + A returning method is not proof of a healthy Redis: several methods catch their own + connection errors and return a default so callers degrade rather than fail. Counting + those as successes reset the failure streak on every request, so the breaker could + never open and Redis was never taken out of the pool. Success is therefore recorded + only when no failure was registered while the method ran. """ @functools.wraps(method) async def wrapper(self, *args, **kwargs): # type: ignore - if self._circuit_breaker.is_open(): - raise Exception(f"Redis circuit breaker is open — skipping {method.__name__}") - try: - result = await method(self, *args, **kwargs) - self._circuit_breaker.record_success() - return result - except Exception: - self._circuit_breaker.record_failure() - raise + return await _run_under_circuit_breaker( + self._circuit_breaker, method.__name__, lambda: method(self, *args, **kwargs) + ) return wrapper @@ -551,13 +625,16 @@ class RedisCache(BaseCache): ) async def run_script(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any: - executor: Optional[Callable[..., Awaitable[Any]]] = litellm.in_memory_llm_clients_cache.get_cache( - key=script_cache_key - ) - if executor is None: - executor = self._register_script_for_current_loop(script) - litellm.in_memory_llm_clients_cache.set_cache(key=script_cache_key, value=executor) - return await executor(keys=keys, args=args, client=client) + async def execute() -> object: + executor: Optional[Callable[..., Awaitable[Any]]] = litellm.in_memory_llm_clients_cache.get_cache( + key=script_cache_key + ) + if executor is None: + executor = self._register_script_for_current_loop(script) + litellm.in_memory_llm_clients_cache.set_cache(key=script_cache_key, value=executor) + return await executor(keys=keys, args=args, client=client) + + return await _run_under_circuit_breaker(self._circuit_breaker, "run_script", execute) return run_script @@ -674,6 +751,7 @@ class RedisCache(BaseCache): str(e), value, ) + _record_swallowed_redis_failure(self._circuit_breaker, e) async def _pipeline_helper( self, @@ -758,6 +836,7 @@ class RedisCache(BaseCache): str(e), cache_value, ) + _record_swallowed_redis_failure(self._circuit_breaker, e) async def _set_cache_sadd_helper( self, @@ -842,6 +921,7 @@ class RedisCache(BaseCache): str(e), value, ) + _record_swallowed_redis_failure(self._circuit_breaker, e) @_redis_circuit_breaker_guard async def batch_cache_write(self, key, value, **kwargs): @@ -1106,6 +1186,7 @@ class RedisCache(BaseCache): ) ) print_verbose(f"litellm.caching.caching: async get() - Got exception from REDIS: {str(e)}") + _record_swallowed_redis_failure(self._circuit_breaker, e) @_redis_circuit_breaker_guard async def async_batch_get_cache( @@ -1177,6 +1258,7 @@ class RedisCache(BaseCache): ) ) verbose_logger.error(f"Error occurred in async batch get cache - {str(e)}") + _record_swallowed_redis_failure(self._circuit_breaker, e) return key_value_dict def sync_ping(self) -> bool: @@ -1432,6 +1514,7 @@ class RedisCache(BaseCache): return ttl except Exception as e: verbose_logger.debug(f"Redis TTL Error: {e}") + _record_swallowed_redis_failure(self._circuit_breaker, e) return None @_redis_circuit_breaker_guard diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 108928871b0..8b831b55da3 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -519,7 +519,13 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac """ This log gets called after the MCP tool call is made. - Useful if you want to modiy the standard logging payload after the MCP tool call is made. + Useful if you want to modify the standard logging payload after the MCP tool call is made. + + To change what the caller sends back to the MCP client, mutate ``response_obj`` + in place: every call site discards the returned object, because the + dispatcher unwraps it to ``mcp_tool_call_response`` (a raw content list, not + a ``CallToolResult``) which the tool-call paths cannot forward. Guardrails + that mask or reject tool output should use ``post_mcp_call`` instead. """ return None diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 12465377b51..60eb960be4e 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -27,6 +27,7 @@ from litellm.integrations.opentelemetry_utils.gen_ai_semconv import ( OTELSemconvCategory, parse_semconv_opt_in, ) +from litellm.integrations.otel.model.semconv import Metric from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.secret_redaction import redact_string from litellm.secret_managers.main import get_secret_bool, str_to_bool @@ -117,6 +118,7 @@ TOKEN_TYPE_ATTRIBUTE: str = "gen_ai.token.type" VALID_METRIC_ATTRIBUTE_NAMES: FrozenSet[str] = frozenset( ( "gen_ai.operation.name", + "gen_ai.provider.name", "gen_ai.system", "gen_ai.request.model", "gen_ai.framework", @@ -597,32 +599,32 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): meter = meter_provider.get_meter(__name__) self._operation_duration_histogram = meter.create_histogram( - name="gen_ai.client.operation.duration", # Replace with semconv constant in otel 1.38 + name=Metric.OPERATION_DURATION, description="GenAI operation duration", unit="s", ) self._token_usage_histogram = meter.create_histogram( - name="gen_ai.client.token.usage", # Replace with semconv constant in otel 1.38 + name=Metric.TOKEN_USAGE, description="GenAI token usage", unit="{token}", ) self._cost_histogram = meter.create_histogram( - name="gen_ai.client.token.cost", + name=Metric.TOKEN_COST, description="GenAI request cost", unit="USD", ) self._time_to_first_token_histogram = meter.create_histogram( - name="gen_ai.client.response.time_to_first_token", + name=Metric.TIME_TO_FIRST_TOKEN, description="Time to first token for streaming requests", unit="s", ) self._time_per_output_token_histogram = meter.create_histogram( - name="gen_ai.client.response.time_per_output_token", + name=Metric.TIME_PER_OUTPUT_TOKEN, description="Average time per output token (generation time / completion tokens)", unit="s", ) self._response_duration_histogram = meter.create_histogram( - name="gen_ai.client.response.duration", + name=Metric.RESPONSE_DURATION, description="Total LLM API generation time (excludes LiteLLM overhead)", unit="s", ) @@ -2980,10 +2982,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def _get_metric_reader(self): """ Get the appropriate metric reader based on the configuration. + + Histograms keep the SDK's default cumulative temporality: Prometheus-backed + OTLP receivers reject delta histograms and drop the whole batch, while + backends that prefer delta still accept cumulative. """ - from opentelemetry.sdk.metrics import Histogram from opentelemetry.sdk.metrics.export import ( - AggregationTemporality, ConsoleMetricExporter, PeriodicExportingMetricReader, ) @@ -3014,7 +3018,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): exporter = OTLPMetricExporter( endpoint=normalized_endpoint, headers=_split_otel_headers, - preferred_temporality={Histogram: AggregationTemporality.DELTA}, ) return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) @@ -3032,7 +3035,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): exporter = OTLPMetricExporter( endpoint=normalized_endpoint, headers=_split_otel_headers, - preferred_temporality={Histogram: AggregationTemporality.DELTA}, ) return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) diff --git a/litellm/integrations/otel/README.md b/litellm/integrations/otel/README.md index 3038bdb90b2..338afe04e5e 100644 --- a/litellm/integrations/otel/README.md +++ b/litellm/integrations/otel/README.md @@ -222,7 +222,29 @@ lives in [`plumbing/`](./plumbing): otherwise the operator's globally configured `MeterProvider` is reused so its readers/exporters receive them alongside the server metrics, and one is built and registered as the global only when none is set (mirroring how V2 owns trace - export). + export). A **failed** call records `gen_ai.client.operation.duration` too, + carrying the semconv `error.type` (the mapped provider exception's class name), + so the histogram covers the whole traffic and failure-rate panels are buildable; + the other five instruments describe a completed generation and are skipped + rather than filled with a fabricated zero. `error.type` is stamped after the + cardinality filter, so an `otel.attributes` list cannot merge the failure series + back into the success series. A proxy-gate rejection (auth / rate limit) records + nothing, for the same reason it gets no span: no upstream call happened. + Both paths cap their attributes at `METRIC_ATTRIBUTE_CEILING` before the + operator's own `otel.attributes` filter runs, so the filter can narrow the set + but never widen it. The ceiling is what keeps series count bounded by the + deployment's own key/team/user/deployment count instead of by its traffic: a + label value that moves per request mints a time series per request, which both + bills per request on a hosted backend and leaves a histogram that cannot be + aggregated. So client-supplied and per-request metadata (`requester_metadata`, + `spend_logs_metadata`, `user_api_key_end_user_id`, `requester_ip_address`) is + metric-ineligible and stays on the span, where cardinality is free, and the + `hidden_params` label carries only `model_id`, the deployment identity a + per-deployment panel joins on. `api_base` is excluded despite naming the same + deployment, because it is a documented per-call parameter and so is caller-chosen + in SDK use. Because the shared validator accepts every span attribute name, a + filter that names a metric-ineligible one logs a warning once when the filter + resolves rather than silently emitting nothing for it. - [`events.py`](./plumbing/events.py) — GenAI client events. Gated on `enable_events` (`LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS`), a failed LLM call records the semconv `gen_ai.client.operation.exception` log event at severity @@ -267,7 +289,10 @@ lives in [`plumbing/`](./plumbing): - **A new attribute vocabulary for a backend**: add a mapper in `mappers/` (a class with a `map(data) -> AttributeMap` method, typically built from - `key -> extractor` tables) and register it in `mappers/__init__._MAPPER_BY_NAME`. + `key -> extractor` tables) and register it in `mappers/__init__._PLAIN_MAPPERS`. + If it spells declared tool definitions out per index, register it in + `_TOOL_DEFINITION_MAPPERS` instead and take the shared attribute budget in its + constructor, so the family stays bounded span-wide rather than per vocabulary. - **A new integration**: add a preset in `presets/` that returns an `OpenTelemetryV2Config`, and register it in `presets/__init__.PRESET_BY_CALLBACK`. If it supports dynamic credentials, add a header builder to diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index b33973f0676..b8908696547 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -281,13 +281,29 @@ class OpenTelemetryV2(CustomLogger): self._record_metrics(kwargs, response_obj, start_time, end_time) def _record_metrics(self, kwargs, response_obj, start_time, end_time) -> None: - """Record the GenAI metrics for a successful LLM call. Best-effort: a - recording failure (e.g. a malformed payload) must never break the span - close or the request itself.""" + """Record the GenAI metrics for a successful LLM call.""" + self._guarded_record(lambda recorder: recorder.record(kwargs, response_obj, start_time, end_time)) + + def _record_failure_metrics(self, kwargs, start_time, end_time) -> None: + """Record the GenAI metrics for a failed LLM call, so the duration + histogram covers the whole traffic rather than only what survived. + + A synthetic proxy-gate log (auth / rate-limit rejection) is skipped for the + same reason it gets no span: no upstream call happened, so its duration is + not a GenAI operation's duration and would pull the histogram down.""" + if LLMCallEvent.from_dict(kwargs).is_no_upstream_call: + return + self._guarded_record(lambda recorder: recorder.record_failure(kwargs, start_time, end_time)) + + def _guarded_record(self, record: "Callable[[GenAIMetricRecorder], None]") -> None: + """Run one metric recording. Best-effort: a recording failure (e.g. a + malformed payload) must never break the span close or the request itself. A + misconfigured attribute filter is operator-fixable, so it is surfaced once + at ERROR instead of being swallowed.""" if self._metrics_recorder is None: return try: - self._metrics_recorder.record(kwargs, response_obj, start_time, end_time) + record(self._metrics_recorder) except ValueError as exc: if not self._metric_filter_error_logged: verbose_logger.error( @@ -304,6 +320,7 @@ class OpenTelemetryV2(CustomLogger): if self._emit_mcp_list_tools(kwargs, start_time, end_time): return self._close_llm_call(kwargs, start_time, end_time) + self._record_failure_metrics(kwargs, start_time, end_time) def _seed_identity_baggage(self, identity: RequestIdentity, model: str | None, context: Context) -> Context: """Seed authenticated request-identity Baggage onto ``context`` so the Baggage diff --git a/litellm/integrations/otel/mappers/__init__.py b/litellm/integrations/otel/mappers/__init__.py index b0c1d7019db..55504c5e8bf 100644 --- a/litellm/integrations/otel/mappers/__init__.py +++ b/litellm/integrations/otel/mappers/__init__.py @@ -18,13 +18,19 @@ from litellm.integrations.otel.mappers.langfuse import LangfuseMapper from litellm.integrations.otel.mappers.langtrace import LangtraceMapper from litellm.integrations.otel.mappers.legacy import LegacyMapper from litellm.integrations.otel.mappers.openinference import OpenInferenceMapper +from litellm.integrations.otel.mappers.utils import tool_attr_budget from litellm.integrations.otel.mappers.weave import WeaveMapper -# Registry keyed by ``config.mapper_names`` entries. -_MAPPER_BY_NAME: dict[str, Callable[[], AttributeMapper]] = { +# Registries keyed by ``config.mapper_names`` entries, split by whether the +# vocabulary spells declared tool definitions out per index. Those share one +# span-wide attribute ceiling, so resolution has to know how many of them are +# active before it can build them. +_TOOL_DEFINITION_MAPPERS: dict[str, Callable[[int], AttributeMapper]] = { "genai": GenAIMapper, "legacy": LegacyMapper, "openinference": OpenInferenceMapper, +} +_PLAIN_MAPPERS: dict[str, Callable[[], AttributeMapper]] = { "langfuse": LangfuseMapper, "weave": WeaveMapper, "langtrace": LangtraceMapper, @@ -33,13 +39,19 @@ _MAPPER_BY_NAME: dict[str, Callable[[], AttributeMapper]] = { def resolve_mappers(names: Iterable[str]) -> list[AttributeMapper]: """Resolve mapper names to instances. Unknown names raise ``ValueError``.""" - out: list[AttributeMapper] = [] - for name in names: - factory = _MAPPER_BY_NAME.get(name) - if factory is None: - raise ValueError(f"unknown mapper name {name!r}; known: {sorted(_MAPPER_BY_NAME)}") - out.append(factory()) - return out + ordered = tuple(names) + for name in ordered: + if name not in _TOOL_DEFINITION_MAPPERS and name not in _PLAIN_MAPPERS: + known = sorted((*_TOOL_DEFINITION_MAPPERS, *_PLAIN_MAPPERS)) + raise ValueError(f"unknown mapper name {name!r}; known: {known}") + # Distinct vocabularies each write the tool family under their own keys, so + # the ceiling is split by how many of them are configured. Repeating a name + # rewrites the same keys, so only distinct ones count. + budget = tool_attr_budget(len({*ordered} & _TOOL_DEFINITION_MAPPERS.keys())) + return [ + _TOOL_DEFINITION_MAPPERS[name](budget) if name in _TOOL_DEFINITION_MAPPERS else _PLAIN_MAPPERS[name]() + for name in ordered + ] __all__ = [ diff --git a/litellm/integrations/otel/mappers/genai.py b/litellm/integrations/otel/mappers/genai.py index f568afa9e3e..70414734b72 100644 --- a/litellm/integrations/otel/mappers/genai.py +++ b/litellm/integrations/otel/mappers/genai.py @@ -11,10 +11,11 @@ from typing import Callable from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, SpanData from litellm.integrations.otel.mappers.utils import ( + MAX_TOOL_DEFINITION_ATTRS_PER_SPAN, collect, - drop_none, output_messages, serialize_messages, + tool_definition_attrs, ) from litellm.integrations.otel.model.payloads import ( GuardrailSpanData, @@ -135,6 +136,9 @@ class GenAIMapper: LiteLLM.SERVICE_CALL_TYPE: lambda d: d.call_type, } + def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None: + self._tool_attr_budget = tool_attr_budget + def map(self, data: SpanData) -> AttributeMap: match data: case LLMCallSpanData(): @@ -150,18 +154,18 @@ class GenAIMapper: case _: return {} - @classmethod - def _llm_call(cls, data: LLMCallSpanData) -> AttributeMap: - attrs = collect(cls._LLM_CALL_ATTRS, data) - attrs.update( - drop_none( - { - f"gen_ai.tool.{idx}.{suffix}": extract(tool) - for idx, tool in enumerate(data.tools) - for suffix, extract in cls._TOOL_ATTRS.items() - } + def _llm_call(self, data: LLMCallSpanData) -> AttributeMap: + attrs = collect(self._LLM_CALL_ATTRS, data) + if data.tools: + attrs[LiteLLM.TOOLS_DECLARED] = len(data.tools) + attrs.update( + tool_definition_attrs( + lambda idx, suffix: f"gen_ai.tool.{idx}.{suffix}", + data.tools, + self._TOOL_ATTRS, + self._tool_attr_budget, + ) ) - ) return attrs @classmethod diff --git a/litellm/integrations/otel/mappers/legacy.py b/litellm/integrations/otel/mappers/legacy.py index 57dc7ed3632..f17a62828e7 100644 --- a/litellm/integrations/otel/mappers/legacy.py +++ b/litellm/integrations/otel/mappers/legacy.py @@ -12,7 +12,11 @@ Like ``GenAIMapper``, each span kind declares its schema as a flat from typing import Callable, Final from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, SpanData -from litellm.integrations.otel.mappers.utils import collect, drop_none +from litellm.integrations.otel.mappers.utils import ( + MAX_TOOL_DEFINITION_ATTRS_PER_SPAN, + collect, + tool_definition_attrs, +) from litellm.integrations.otel.model.payloads import ( LLMCallSpanData, ServiceSpanData, @@ -63,6 +67,9 @@ class LegacyMapper: _LEGACY_ERROR: lambda d: d.error.message if d.error is not None and d.error.message else None, } + def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None: + self._tool_attr_budget = tool_attr_budget + def map(self, data: SpanData) -> AttributeMap: match data: case LLMCallSpanData(): @@ -72,16 +79,14 @@ class LegacyMapper: case _: return {} - @classmethod - def _llm_call(cls, data: LLMCallSpanData) -> AttributeMap: - attrs = collect(cls._LLM_CALL_ATTRS, data) + def _llm_call(self, data: LLMCallSpanData) -> AttributeMap: + attrs = collect(self._LLM_CALL_ATTRS, data) attrs.update( - drop_none( - { - f"llm.request.functions.{idx}.{suffix}": extract(tool) - for idx, tool in enumerate(data.tools) - for suffix, extract in cls._TOOL_ATTRS.items() - } + tool_definition_attrs( + lambda idx, suffix: f"llm.request.functions.{idx}.{suffix}", + data.tools, + self._TOOL_ATTRS, + self._tool_attr_budget, ) ) return attrs diff --git a/litellm/integrations/otel/mappers/openinference.py b/litellm/integrations/otel/mappers/openinference.py index dab9a616979..87c4d0d6484 100644 --- a/litellm/integrations/otel/mappers/openinference.py +++ b/litellm/integrations/otel/mappers/openinference.py @@ -13,9 +13,11 @@ from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, Span from litellm.integrations.otel.mappers.utils import ( collect, drop_none, + MAX_TOOL_DEFINITION_ATTRS_PER_SPAN, json_if, message_content, output_messages, + tool_definition_attrs, ) from litellm.integrations.otel.model.payloads import ( LLMCallSpanData, @@ -70,6 +72,9 @@ class OpenInferenceMapper: ), } + def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None: + self._tool_attr_budget = tool_attr_budget + def map(self, data: SpanData) -> AttributeMap: match data: case LLMCallSpanData(): @@ -77,14 +82,13 @@ class OpenInferenceMapper: case _: return {} - @classmethod - def _llm_call(cls, data: LLMCallSpanData) -> AttributeMap: + def _llm_call(self, data: LLMCallSpanData) -> AttributeMap: return { - **collect(cls._LLM_CALL_ATTRS, data), - **collect(cls._BLOB_ATTRS, data), - **cls._messages("llm.input_messages", "input.value", data.messages_in), - **cls._messages("llm.output_messages", "output.value", output_messages(data)), - **cls._tools(data), + **collect(self._LLM_CALL_ATTRS, data), + **collect(self._BLOB_ATTRS, data), + **self._messages("llm.input_messages", "input.value", data.messages_in), + **self._messages("llm.output_messages", "output.value", output_messages(data)), + **self._tools(data), } @staticmethod @@ -108,12 +112,10 @@ class OpenInferenceMapper: attrs[value_key] = json.dumps([{"role": role, "content": content} for role, content in parsed]) return attrs - @classmethod - def _tools(cls, data: LLMCallSpanData) -> AttributeMap: - return drop_none( - { - f"llm.tools.{idx}.{suffix}": extract(tool) - for idx, tool in enumerate(data.tools) - for suffix, extract in cls._TOOL_ATTRS.items() - } + def _tools(self, data: LLMCallSpanData) -> AttributeMap: + return tool_definition_attrs( + lambda idx, suffix: f"llm.tools.{idx}.{suffix}", + data.tools, + self._TOOL_ATTRS, + self._tool_attr_budget, ) diff --git a/litellm/integrations/otel/mappers/utils.py b/litellm/integrations/otel/mappers/utils.py index a91e59e4ab8..cbdb60f42c9 100644 --- a/litellm/integrations/otel/mappers/utils.py +++ b/litellm/integrations/otel/mappers/utils.py @@ -6,10 +6,34 @@ they live in one place. """ import json -from typing import Callable, Mapping, Sequence +from typing import Callable, Final, Mapping, Sequence from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue -from litellm.integrations.otel.model.payloads import LLMCallSpanData +from litellm.integrations.otel.model.payloads import LLMCallSpanData, ToolDefinition + +DEFAULT_SPAN_ATTRIBUTE_LIMIT: Final = 128 +"""The OTel SDK's default per-span attribute count limit.""" + +MAX_TOOL_DEFINITION_ATTRS_PER_SPAN: Final = DEFAULT_SPAN_ATTRIBUTE_LIMIT // 4 +"""Span-wide ceiling on attributes spent spelling out declared tool definitions. + +Tool definitions are an unbounded attribute family: one entry per declared +tool, per field, per active vocabulary. Agentic clients declare hundreds, which +overruns the span attribute limit. That limit evicts oldest-first, so an +uncapped family silently destroys the core ``gen_ai.*`` attributes written +before it. + +The ceiling is span-wide rather than per-mapper because several vocabularies +can be active at once and each spells the same tools out under its own keys, so +a per-mapper allowance multiplies by the number of vocabularies and reaches the +limit again. Reserving a quarter of the span for tool detail leaves the rest to +core telemetry no matter how many vocabularies are configured. +""" + + +def tool_attr_budget(vocabularies: int) -> int: + """Split the span-wide tool-definition ceiling across active vocabularies.""" + return MAX_TOOL_DEFINITION_ATTRS_PER_SPAN // max(vocabularies, 1) def drop_none(values: Mapping[str, AttrValue | None]) -> AttributeMap: @@ -17,6 +41,29 @@ def drop_none(values: Mapping[str, AttrValue | None]) -> AttributeMap: return {k: v for k, v in values.items() if v is not None} +def tool_definition_attrs( + key_for: Callable[[int, str], str], + tools: Sequence[ToolDefinition], + extractors: Mapping[str, Callable[[ToolDefinition], AttrValue | None]], + attr_budget: int, +) -> AttributeMap: + """Per-index attributes for as many tools as ``attr_budget`` affords. + + ``key_for`` builds a vocabulary's key from the tool's index and the field + name, so each mapper keeps its own naming while sharing the budget. One tool + always keeps its detail, so the family stays legible even when many + vocabularies split the ceiling. + """ + max_tools = max(attr_budget // max(len(extractors), 1), 1) + return drop_none( + { + key_for(idx, suffix): extract(tool) + for idx, tool in enumerate(tools[:max_tools]) + for suffix, extract in extractors.items() + } + ) + + def collect(table: Mapping[str, Callable], source: object) -> AttributeMap: """Apply an extractor table to ``source``, dropping ``None`` results.""" return drop_none({key: extract(source) for key, extract in table.items()}) diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index 44b2f7e0488..3b38b3d3655 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -6,17 +6,30 @@ without a semconv equivalent lives under the ``litellm.*`` vendor namespace. from enum import Enum from typing import Final +from litellm._logging import verbose_logger + class GenAIOperation(str, Enum): - """Values for ``gen_ai.operation.name``.""" + """Values for ``gen_ai.operation.name``. + + The first block is the convention's own vocabulary. The ``LITELLM_`` members + are vendor values for operations the convention names nothing for; its note + on this attribute directs instrumentation to use a system-specific name in + exactly that case, the same allowance :func:`resolve_provider` relies on for + unmapped providers. They stay under the ``litellm.`` prefix so a value the + convention adds later can never collide with one of ours. + """ CHAT = "chat" TEXT_COMPLETION = "text_completion" EMBEDDINGS = "embeddings" GENERATE_CONTENT = "generate_content" + RETRIEVAL = "retrieval" # vector-store search / RAG query spans CREATE_AGENT = "create_agent" # reserved for future agent spans - INVOKE_AGENT = "invoke_agent" # reserved for future agent spans + INVOKE_AGENT = "invoke_agent" # agent (A2A) message spans EXECUTE_TOOL = "execute_tool" # MCP tool-call spans + LITELLM_VECTOR_STORE_MANAGEMENT = "litellm.vector_store_management" + LITELLM_VECTOR_STORE_FILE_MANAGEMENT = "litellm.vector_store_file_management" class GenAIProvider(str, Enum): @@ -49,11 +62,17 @@ class MCPMethod(str, Enum): class GenAI: - """Canonical OTel GenAI span-attribute keys.""" + """Canonical OTel GenAI attribute keys. + + ``SYSTEM`` is the one exception: the convention deprecated it in favor of + ``PROVIDER_NAME``, and it survives here only so already-shipped series keep + resolving for consumers that query it. Nothing new should use it. + """ # request OPERATION_NAME: Final = "gen_ai.operation.name" PROVIDER_NAME: Final = "gen_ai.provider.name" + SYSTEM: Final = "gen_ai.system" REQUEST_MODEL: Final = "gen_ai.request.model" REQUEST_TEMPERATURE: Final = "gen_ai.request.temperature" REQUEST_TOP_P: Final = "gen_ai.request.top_p" @@ -233,6 +252,7 @@ class LiteLLM: # ``litellm_params.model``), distinct from the user-facing ``gen_ai.request.model``. PROVIDER_MODEL: Final = "litellm.provider.model" REQUEST_STREAMING: Final = "litellm.request.streaming" + TOOLS_DECLARED: Final = "litellm.request.tools.declared" GUARDRAIL_NAME: Final = "litellm.guardrail.name" GUARDRAIL_MODE: Final = "litellm.guardrail.mode" GUARDRAIL_STATUS: Final = "litellm.guardrail.status" @@ -257,13 +277,27 @@ class LiteLLM: class Metric: - """GenAI metric instrument names.""" + """GenAI metric instrument names. + + Every name here that a convention or a backend defines uses that name, so a + consumer charting GenAI telemetry finds litellm's series where it looks for + them. ``TOKEN_USAGE``, ``OPERATION_DURATION``, ``TIME_TO_FIRST_TOKEN`` and + ``TIME_PER_OUTPUT_TOKEN`` are semconv instruments, defined in the GenAI + conventions; the ``gen_ai.client.response.*`` spellings litellm used for the + latter two are not conventions at all, so nothing downstream could chart + them. Cost has no semconv instrument, so it takes ``gen_ai.usage.cost``, the + name backends already query for spend. + + ``RESPONSE_DURATION`` keeps its vendor spelling deliberately: the closest + convention, ``gen_ai.server.request.duration``, would collide in meaning with + ``OPERATION_DURATION``, which litellm already emits for the whole operation. + """ TOKEN_USAGE: Final = "gen_ai.client.token.usage" OPERATION_DURATION: Final = "gen_ai.client.operation.duration" - TOKEN_COST: Final = "gen_ai.client.token.cost" - TIME_TO_FIRST_TOKEN: Final = "gen_ai.client.response.time_to_first_token" - TIME_PER_OUTPUT_TOKEN: Final = "gen_ai.client.response.time_per_output_token" + TOKEN_COST: Final = "gen_ai.usage.cost" + TIME_TO_FIRST_TOKEN: Final = "gen_ai.server.time_to_first_token" + TIME_PER_OUTPUT_TOKEN: Final = "gen_ai.server.time_per_output_token" RESPONSE_DURATION: Final = "gen_ai.client.response.duration" @@ -301,6 +335,35 @@ _OPERATION_BY_CALL_TYPE: dict[str, GenAIOperation] = { "responses": GenAIOperation.CHAT, "aresponses": GenAIOperation.CHAT, "call_mcp_tool": GenAIOperation.EXECUTE_TOOL, + "vector_store_search": GenAIOperation.RETRIEVAL, + "avector_store_search": GenAIOperation.RETRIEVAL, + "query": GenAIOperation.RETRIEVAL, + "aquery": GenAIOperation.RETRIEVAL, + "send_message": GenAIOperation.INVOKE_AGENT, + "asend_message": GenAIOperation.INVOKE_AGENT, + "asend_message_streaming": GenAIOperation.INVOKE_AGENT, + "vector_store_create": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT, + "avector_store_create": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT, + "vector_store_retrieve": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT, + "avector_store_retrieve": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT, + "vector_store_list": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT, + "avector_store_list": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT, + "vector_store_update": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT, + "avector_store_update": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT, + "vector_store_delete": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT, + "avector_store_delete": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT, + "vector_store_file_create": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "avector_store_file_create": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "vector_store_file_list": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "avector_store_file_list": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "vector_store_file_retrieve": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "avector_store_file_retrieve": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "vector_store_file_content": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "avector_store_file_content": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "vector_store_file_update": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "avector_store_file_update": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "vector_store_file_delete": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, + "avector_store_file_delete": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT, } @@ -317,7 +380,21 @@ def resolve_provider(custom_llm_provider: str | None) -> str: def resolve_operation(call_type: str | None) -> GenAIOperation: - """Map a litellm ``call_type`` to a ``gen_ai.operation.name`` value.""" + """Map a litellm ``call_type`` to a ``gen_ai.operation.name`` value. + + An unmapped call type still falls back to ``chat`` so every series keeps an + operation label, but it logs at debug rather than falling through silently: + a new call type mislabelled as ``chat`` mixes its latency and cost into + everyone's chat charts, which is invisible until someone reads the numbers. + """ if not call_type: return GenAIOperation.CHAT - return _OPERATION_BY_CALL_TYPE.get(call_type.lower(), GenAIOperation.CHAT) + mapped = _OPERATION_BY_CALL_TYPE.get(call_type.lower()) + if mapped is not None: + return mapped + verbose_logger.debug( + "otel: call_type %r has no gen_ai.operation.name mapping; labelling it %r. Add it to _OPERATION_BY_CALL_TYPE.", + call_type, + GenAIOperation.CHAT.value, + ) + return GenAIOperation.CHAT diff --git a/litellm/integrations/otel/model/utils.py b/litellm/integrations/otel/model/utils.py index f37afc97879..ab54a558a9a 100644 --- a/litellm/integrations/otel/model/utils.py +++ b/litellm/integrations/otel/model/utils.py @@ -1,9 +1,11 @@ """Shared, OpenTelemetry-free helpers for the otel integration. -Generic value coercion (for reading heterogeneous logging-payload dicts), time -conversion, and header parsing — pulled out of the individual modules so they -live in one place. Deliberately free of any ``opentelemetry`` import so the -OTel-free sources of truth (payloads, semconv, spans, config) can use it too. +Generic value coercion (for reading heterogeneous logging-payload dicts) and +time conversion — pulled out of the individual modules so they live in one +place. Deliberately free of any ``opentelemetry`` import so the OTel-free +sources of truth (payloads, semconv, spans, config) can use it too. OTLP header +parsing lives in :mod:`litellm.integrations.otel.plumbing.providers` instead, +because it delegates to the OTel SDK's own W3C Baggage parser. """ from datetime import datetime @@ -89,15 +91,3 @@ def to_seconds(value: datetime | float | int | str | None) -> float | None: except ValueError: continue return None - - -def parse_headers(raw: str | None) -> dict[str, str]: - """Parse an OTLP ``"k=v,k=v"`` header string into a dict.""" - headers: dict[str, str] = {} - if not raw: - return headers - for pair in raw.split(","): - if "=" in pair: - key, _, value = pair.partition("=") - headers[key.strip()] = value.strip() - return headers diff --git a/litellm/integrations/otel/plumbing/metrics.py b/litellm/integrations/otel/plumbing/metrics.py index 50d0fb75962..6ebaacafdc4 100644 --- a/litellm/integrations/otel/plumbing/metrics.py +++ b/litellm/integrations/otel/plumbing/metrics.py @@ -1,6 +1,6 @@ """GenAI client metrics: the six ``gen_ai.client.*`` histograms plus the recorder that builds attributes, applies the shared cardinality filter, and -records a request's metrics in the success path. +records a request's metrics on both the success and the failure path. The instrument names/units/descriptions and the recording + timing math mirror the v1 :mod:`litellm.integrations.opentelemetry` integration so both engines emit @@ -10,11 +10,12 @@ identical metrics. The attribute cardinality filter is reused from v1 by import from dataclasses import dataclass from datetime import datetime -from typing import Any, FrozenSet, Mapping, Optional +from typing import Any, Final, FrozenSet, Mapping, Optional, TypeAlias from opentelemetry.metrics import Histogram, Meter import litellm +from litellm._logging import verbose_logger from litellm.integrations.opentelemetry import ( METRIC_METADATA_KEYS, TOKEN_TYPE_ATTRIBUTE, @@ -22,11 +23,34 @@ from litellm.integrations.opentelemetry import ( _resolve_metric_attribute_filter, ) from litellm.integrations.otel.model.metadata import time_to_first_chunk_seconds -from litellm.integrations.otel.model.semconv import Metric, resolve_operation +from litellm.integrations.otel.model.semconv import ( + Error, + GenAI, + Metric, + resolve_operation, + resolve_provider, +) from litellm.integrations.otel.model.utils import to_seconds from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +def _provider_attributes(custom_llm_provider: object) -> Mapping[str, str]: + """The provider labels for one call's metrics. + + ``gen_ai.provider.name`` carries the semconv-mapped value; the deprecated + ``gen_ai.system`` spelling is dual-emitted with the raw litellm provider + string it has always carried, so a dashboard already querying it keeps + matching. A call with no provider gets neither label: a placeholder value + would mint a permanent series that no operator can act on. + """ + if not isinstance(custom_llm_provider, str) or not custom_llm_provider: + return {} + return { + GenAI.PROVIDER_NAME: resolve_provider(custom_llm_provider), + GenAI.SYSTEM: custom_llm_provider, + } + + @dataclass(frozen=True) class GenAIMetrics: operation_duration: Histogram @@ -72,8 +96,83 @@ def create_genai_metrics(meter: Meter) -> GenAIMetrics: ) +# A metric datapoint's attributes. Values are the strings the recorder builds, except +# the request model, which is whatever the caller passed and may be absent. +MetricAttributes: TypeAlias = Mapping[str, "str | None"] + +ERROR_TYPE_FALLBACK: Final = "_OTHER" + +# Every attribute a metric datapoint may carry, on either path. A label value that +# is unique per request is a new time series that will never be written to again, so +# this set is what keeps the series count bounded by the deployment's own +# key/team/user/deployment count rather than by its traffic. Each entry is a fixed +# enum or an operator-provisioned identifier. +# +# Deliberately excluded is everything the *client* supplies or that moves per +# request: ``metadata.requester_metadata`` and ``metadata.spend_logs_metadata`` (both +# free-form from the request body), ``metadata.user_api_key_end_user_id`` (the body's +# ``user`` field), and ``metadata.requester_ip_address``. Those stay on the span, +# where cardinality is free and where they already are. +# ``metadata.user_api_key_user_email`` is left out too: it is bounded, but it is PII +# duplicating the user id already here. +# +# This is a CEILING, applied before the operator's own include/exclude filter, so an +# operator can narrow it but never widen it back to an unbounded attribute. +METRIC_ATTRIBUTE_CEILING: Final[frozenset[str]] = frozenset( + ( + "gen_ai.operation.name", + "gen_ai.provider.name", + "gen_ai.system", + "gen_ai.request.model", + "gen_ai.framework", + "metadata.user_api_key_hash", + "metadata.user_api_key_alias", + "metadata.user_api_key_team_id", + "metadata.user_api_key_team_alias", + "metadata.user_api_key_org_id", + "metadata.user_api_key_user_id", + "hidden_params", + ) +) + +# The only ``hidden_params`` field that becomes part of the ``hidden_params`` label. +# The object as a whole is per-request by construction -- ``response_cost``, +# ``litellm_overhead_time_ms``, ``cache_key``, ``usage_object`` and the provider's +# ``additional_headers`` rate-limit counters all move on every call -- so dumping it +# whole made one series per request out of every instrument. +# +# ``model_id`` is the router's own deployment id, so it is bounded by the deployment +# list and is what a per-deployment panel joins on. ``api_base`` is deliberately NOT +# here even though it names the same thing: it is a documented per-call parameter, so +# in SDK use it is chosen by the caller rather than provisioned by the operator, and a +# caller varying it would put the per-request cardinality straight back. +BOUNDED_HIDDEN_PARAM_KEYS: Final[tuple[str, ...]] = ("model_id",) + + +def resolve_error_type(kwargs: Mapping[str, Any]) -> str: + """The ``error.type`` value for a failed request. + + Bounded by construction: the mapped provider exception's class name (the same + ``error_information.error_class`` the failure span stamps), else the provider + status code, else the raw exception's class name, else ``_OTHER`` — the value + the convention reserves for a failure the instrumentation cannot classify. The + exception *message* is unbounded and never becomes a label; it stays on the + span and its exception event, where high cardinality is free. + """ + std_log = kwargs.get("standard_logging_object") + info = getattr(std_log, "error_information", None) or (std_log or {}).get("error_information") or {} + error_class = info.get("error_class") or info.get("error_code") + if error_class: + return str(error_class) + exception = kwargs.get("exception") + if exception is not None: + return type(exception).__name__ + return ERROR_TYPE_FALLBACK + + class GenAIMetricRecorder: - """Records the six GenAI histograms for one successful LLM call. + """Records the six GenAI histograms for one successful LLM call, and the + duration histogram alone for one failed LLM call (see :meth:`record_failure`). The cardinality filter is resolved lazily on the first record: the proxy populates ``callback_settings.otel.attributes`` after the logger is built, so @@ -96,7 +195,7 @@ class GenAIMetricRecorder: start_time: datetime, end_time: datetime, ) -> None: - common_attrs = self._filter_attributes(self._common_attributes(kwargs)) + common_attrs = self._filter_attributes(self._bounded_attributes(kwargs)) duration_s = (end_time - start_time).total_seconds() self._metrics.operation_duration.record(duration_s, attributes=common_attrs) @@ -110,17 +209,48 @@ class GenAIMetricRecorder: self._record_time_per_output_token(kwargs, response_obj, end_time, duration_s, common_attrs) self._record_response_duration(kwargs, end_time, common_attrs) + def record_failure( + self, + kwargs: Mapping[str, Any], + start_time: datetime, + end_time: datetime, + ) -> None: + """Record the one metric a failed request can honestly report: the + operation's duration, tagged with ``error.type``. + + The other five instruments all describe a completed generation and have + nothing to measure here. litellm hands the failure callback no + ``response_obj`` at all, so there is no usage to split into input/output + tokens and no completion-token count to divide generation time by; it also + zeroes ``response_cost`` on failure. Recording them anyway would put a + fabricated zero into series that dashboards average. + + The attribute set is :data:`METRIC_ATTRIBUTE_CEILING`, the same cap the + success path uses. A failure needs no provider spend, so a caller who can put + a unique value into a client-supplied attribute could mint one histogram + series per request for free; the cap is what makes that impossible on either + path. + + ``error.type`` is stamped after both filters, exactly like + ``gen_ai.token.type``, so an operator's include/exclude list cannot strip + the discriminator and silently merge failures back into the success series. + """ + attributes = { + **self._filter_attributes(self._bounded_attributes(kwargs)), + Error.TYPE: resolve_error_type(kwargs), + } + self._metrics.operation_duration.record((end_time - start_time).total_seconds(), attributes=attributes) + # ------------------------------------------------------------------ # # Attribute building + cardinality filter # ------------------------------------------------------------------ # def _common_attributes(self, kwargs: Mapping[str, Any]) -> dict: params = kwargs.get("litellm_params") or {} - provider = params.get("custom_llm_provider", "Unknown") common_attrs: dict = { - "gen_ai.operation.name": resolve_operation(kwargs.get("call_type")).value, - "gen_ai.system": provider, - "gen_ai.request.model": kwargs.get("model"), + GenAI.OPERATION_NAME: resolve_operation(kwargs.get("call_type")).value, + **_provider_attributes(params.get("custom_llm_provider")), + GenAI.REQUEST_MODEL: kwargs.get("model"), "gen_ai.framework": "litellm", } @@ -136,11 +266,25 @@ class GenAIMetricRecorder: common_attrs[f"metadata.{key}"] = str(value) hidden_params = getattr(std_log, "hidden_params", None) or (std_log or {}).get("hidden_params", {}) - if hidden_params: - common_attrs["hidden_params"] = safe_dumps(hidden_params) + bounded_hidden_params = { + key: hidden_params[key] + for key in BOUNDED_HIDDEN_PARAM_KEYS + if isinstance(hidden_params, Mapping) and hidden_params.get(key) is not None + } + if bounded_hidden_params: + common_attrs["hidden_params"] = safe_dumps(bounded_hidden_params) return common_attrs + def _bounded_attributes(self, kwargs: Mapping[str, Any]) -> MetricAttributes: + """The datapoint attributes, capped at :data:`METRIC_ATTRIBUTE_CEILING`. + + The cap runs BEFORE the operator's include/exclude filter so the filter can + only narrow it. An operator who names an excluded attribute in an include + list gets nothing for it rather than reintroducing an unbounded label. + """ + return {k: v for k, v in self._common_attributes(kwargs).items() if k in METRIC_ATTRIBUTE_CEILING} + def _ensure_filter(self) -> None: if self._filter_resolved: return @@ -157,8 +301,29 @@ class GenAIMetricRecorder: # without reconstructing the recorder. self._include, self._exclude = _resolve_metric_attribute_filter(attributes) self._filter_resolved = True + self._warn_about_metric_ineligible_names() - def _filter_attributes(self, attrs: dict) -> dict: + def _warn_about_metric_ineligible_names(self) -> None: + """Say so when the operator's filter names an attribute the ceiling removes. + + The shared validator accepts every span attribute name, so a name that is + legal on a span but metric-ineligible would otherwise be a silent no-op: an + ``include_list`` naming it emits nothing for it and an ``exclude_list`` naming + it looks like it worked. Logged once, when the filter resolves, rather than + per request. + """ + named = (self._include or frozenset()) | (self._exclude or frozenset()) + ineligible = sorted(named - METRIC_ATTRIBUTE_CEILING - {TOKEN_TYPE_ATTRIBUTE}) + if ineligible: + verbose_logger.warning( + "OTel metrics: %s cannot be a metric attribute and is being ignored; it varies " + "per request or is client-supplied, so it would make one time series per request. " + "It is still on the span. Metric attributes are limited to: %s", + ", ".join(ineligible), + ", ".join(sorted(METRIC_ATTRIBUTE_CEILING)), + ) + + def _filter_attributes(self, attrs: MetricAttributes) -> MetricAttributes: self._ensure_filter() if self._include is not None: return {k: v for k, v in attrs.items() if k in self._include} diff --git a/litellm/integrations/otel/plumbing/providers.py b/litellm/integrations/otel/plumbing/providers.py index ced65aa1ec3..ede9acc6d8f 100644 --- a/litellm/integrations/otel/plumbing/providers.py +++ b/litellm/integrations/otel/plumbing/providers.py @@ -29,15 +29,13 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( InMemorySpanExporter, ) from opentelemetry.trace import Span, SpanKind, Tracer +from opentelemetry.util.re import parse_env_headers from litellm._version import version as litellm_version from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config from litellm.integrations.otel.model.semconv import LiteLLM from litellm.integrations.otel.model.spans import LiteLLMSpanKind -# Re-exported so ``providers.parse_headers`` remains a stable entry point. -from litellm.integrations.otel.model.utils import parse_headers as parse_headers - if TYPE_CHECKING: from opentelemetry.metrics import Meter from opentelemetry.sdk.metrics.export import MetricReader @@ -119,6 +117,23 @@ def _otlp_traces_endpoint(endpoint: str | None) -> str | None: return endpoint + "/v1/traces" +def parse_headers(raw: str | None) -> dict[str, str]: + """Parse an OTLP ``"k=v,k=v"`` header string into a dict. + + ``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded per the OTLP spec, so + values are percent-decoded: a vendor that documents + ``Authorization=Basic%20`` (Grafana Cloud does, because a bare space + is not representable there) has to reach the exporter as ``Basic ``, + not with a literal ``%20`` that the backend rejects as malformed. The SDK's + own parser is used so litellm decodes exactly what the OTLP exporters do + when they read the env var themselves; ``liberal`` keeps values that are not + percent-encoded (``Authorization=Bearer ``) working unchanged. + """ + if not raw: + return {} + return dict(parse_env_headers(raw, liberal=True)) + + def _exporter_from_spec(spec: ExporterSpec) -> SpanExporter: kind = (spec.kind or "console").lower() factory = _EXPORTER_FACTORIES.get(kind) @@ -191,6 +206,13 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader": ``console`` (and any unrecognized kind) exports to the console; ``otlp_http`` and ``otlp_grpc`` export over OTLP with the configured endpoint/headers. The reader exports on a 5s period, matching v1. + + Histograms keep the SDK's default cumulative temporality. Prometheus-backed + OTLP receivers (Grafana Cloud / Mimir, and the Prometheus OTLP endpoint) + reject delta histograms outright with ``invalid temporality and type + combination``, which drops the whole metric batch, while backends that + prefer delta still accept cumulative. The enterprise billing exporter + already relies on the same default. """ from opentelemetry.sdk.metrics.export import ( ConsoleMetricExporter, @@ -202,18 +224,12 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader": from opentelemetry.exporter.otlp.proto.http.metric_exporter import ( OTLPMetricExporter as HTTPMetricExporter, ) - from opentelemetry.sdk.metrics import Histogram - from opentelemetry.sdk.metrics.export import AggregationTemporality exporter: Any = HTTPMetricExporter( endpoint=_otlp_metrics_endpoint(config.endpoint), headers=parse_headers(config.headers), - preferred_temporality={Histogram: AggregationTemporality.DELTA}, ) elif kind in ("otlp_grpc", "grpc"): - from opentelemetry.sdk.metrics import Histogram - from opentelemetry.sdk.metrics.export import AggregationTemporality - try: from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import ( OTLPMetricExporter as GRPCMetricExporter, @@ -227,7 +243,6 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader": exporter = GRPCMetricExporter( endpoint=config.endpoint, headers=parse_headers(config.headers), - preferred_temporality={Histogram: AggregationTemporality.DELTA}, ) else: exporter = ConsoleMetricExporter() diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 24597c02ea2..41a4d026fe1 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -69,6 +69,15 @@ else: _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT = 5.0 +# Tiers a caller may name in a request, across the providers that accept the +# parameter: OpenAI ("auto", "default", "flex", "priority", "scale"), Bedrock and +# Groq (subsets of those), Anthropic ("auto", "standard_only") and Vertex, which +# maps "default" to "standard". Used to bound the caller-controlled fallback in +# ``get_service_tier_from_standard_logging_payload``. +KNOWN_REQUEST_SERVICE_TIERS = frozenset( + {"auto", "batch", "default", "flex", "priority", "scale", "standard", "standard_only"} +) + def _get_budget_metrics_per_request_timeout() -> float: raw = os.getenv("PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT") @@ -1245,6 +1254,7 @@ class PrometheusLogger(CustomLogger): client_ip=standard_logging_payload["metadata"].get("requester_ip_address"), user_agent=standard_logging_payload["metadata"].get("user_agent"), stream=(str(standard_logging_payload.get("stream")) if litellm.prometheus_emit_stream_label else None), + service_tier=get_service_tier_from_standard_logging_payload(standard_logging_payload), ) if user_api_key is not None and isinstance(user_api_key, str) and user_api_key.startswith("sk-"): @@ -4098,6 +4108,44 @@ def get_custom_labels_from_metadata(metadata: dict) -> Dict[str, str]: return result +def get_service_tier_from_standard_logging_payload( + standard_logging_payload: StandardLoggingPayload, +) -> str | None: + """ + Resolve the service tier a request ran on, for the ``service_tier`` label. + + The tier the provider actually served wins over the tier the caller asked for, + so latency and spend stay segmentable when the request said ``auto`` and the + provider picked the concrete tier. Providers report the served tier either at + the top level of the response (OpenAI, Bedrock, Groq) or on the usage object + (Anthropic). + + Streaming responses carry no served tier, so the requested tier is the + fallback. That value is caller-controlled and survives param mapping even + where the provider then ignores it (Bedrock and Groq accept the request and + drop an unrecognized tier), so it is only labelled when it names a known + tier; otherwise one caller could mint a Prometheus series per string. Values + the provider itself reports are not caller-controlled and stay unrestricted, + so a tier a provider adds later is still labelled correctly. + """ + response = standard_logging_payload.get("response") + usage_object = standard_logging_payload.get("metadata", {}).get("usage_object") + + served_candidates: tuple[object, ...] = ( + response.get("service_tier") if isinstance(response, dict) else None, + usage_object.get("service_tier") if isinstance(usage_object, dict) else None, + ) + served_tier = next((tier for tier in served_candidates if isinstance(tier, str) and tier), None) + if served_tier is not None: + return served_tier + + model_parameters = standard_logging_payload.get("model_parameters") + requested_tier = model_parameters.get("service_tier") if isinstance(model_parameters, dict) else None + if isinstance(requested_tier, str) and requested_tier in KNOWN_REQUEST_SERVICE_TIERS: + return requested_tier + return None + + def _get_combined_custom_metadata_from_standard_logging_payload( standard_logging_payload: Optional[dict], ) -> Dict[str, Any]: diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 83d6fcc0bee..c2dc7189934 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -4680,6 +4680,7 @@ class StandardLoggingPayloadSetup: applied_guardrails=applied_guardrails, mcp_tool_call_metadata=mcp_tool_call_metadata, vector_store_request_metadata=vector_store_request_metadata, + routing_decision=None, usage_object=usage_object, requester_custom_headers=None, cold_storage_object_key=None, @@ -5519,6 +5520,7 @@ def get_standard_logging_metadata( applied_guardrails=None, mcp_tool_call_metadata=None, vector_store_request_metadata=None, + routing_decision=None, usage_object=None, requester_custom_headers=None, user_api_key_request_route=None, diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 4e3d94e2ab3..d8ce48f05de 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1267,7 +1267,7 @@ def _get_dummy_thought_signature() -> str: def convert_to_gemini_tool_call_invoke( message: ChatCompletionAssistantMessage, model: Optional[str] = None, - custom_llm_provider: Optional[str] = None, + forward_function_call_id: bool = False, ) -> List[VertexPartType]: """ OpenAI tool invokes: @@ -1317,16 +1317,12 @@ def convert_to_gemini_tool_call_invoke( VertexGeminiConfig, ) - forward_tool_call_id = bool( - model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider) - ) - if tool_calls is not None: for idx, tool in enumerate(tool_calls): if "function" in tool: gemini_function_call: Optional[VertexFunctionCall] = _gemini_tool_call_invoke_helper( function_call_params=tool["function"], - tool_call_id=(tool.get("id") if forward_tool_call_id else None), + tool_call_id=(tool.get("id") if forward_function_call_id else None), ) if gemini_function_call is not None: part_dict: VertexPartType = {"function_call": gemini_function_call} @@ -1378,8 +1374,7 @@ def convert_to_gemini_tool_call_invoke( def convert_to_gemini_tool_call_result( message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], last_message_with_tool_calls: Optional[dict], - model: Optional[str] = None, - custom_llm_provider: Optional[str] = None, + forward_function_call_id: bool = False, ) -> Union[VertexPartType, List[VertexPartType]]: """ OpenAI message with a tool result looks like: @@ -1501,14 +1496,8 @@ def convert_to_gemini_tool_call_result( name = tool.get("function", {}).get("name", "") # Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix). - # Only Google AI Studio Gemini 3+ accepts `id` on function_response parts. - # Vertex AI and older Gemini models reject the field with HTTP 400. - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - gemini_call_id: Optional[str] = None - if model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider): + if forward_function_call_id: raw_tool_call_id = message.get("tool_call_id") if raw_tool_call_id and isinstance(raw_tool_call_id, str): stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0] diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index f02333c34c8..853bea636af 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -393,24 +393,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): if compaction_event is not None: return compaction_event - if self.sent_content_block_start is False: - self.sent_content_block_start = True - self.sent_content_block_finish = False - self.chunk_queue.append( - { - "type": "content_block_start", - "index": self.current_content_block_index, - "content_block": {"type": "text", "text": ""}, - } - ) - return self.chunk_queue.popleft() - for chunk in self.completion_stream: if chunk == "None" or chunk is None: raise Exception should_start_new_block = self._should_start_new_content_block(chunk) - if should_start_new_block: + is_opening_first_block = self.sent_content_block_start is False + if is_opening_first_block and self._is_blank_delta(chunk): + continue + if is_opening_first_block: + self.sent_content_block_start = True + self.sent_content_block_finish = False + self.chunk_queue.append( + { + "type": "content_block_start", + "index": self.current_content_block_index, + "content_block": self.current_content_block_start, + } + ) + elif should_start_new_block: self._increment_content_block_index() # applied_edits only needs to flow to the final message_delta @@ -447,7 +448,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): # ``not self.queued_usage_chunk``. continue - if should_start_new_block and not self.sent_content_block_finish: + if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish: # Queue the sequence: content_block_stop -> content_block_start # -> (optionally) the trigger chunk's delta. # @@ -615,25 +616,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): if compaction_event is not None: return compaction_event - if self.sent_content_block_start is False: - self.sent_content_block_start = True - self.sent_content_block_finish = False - self.chunk_queue.append( - { - "type": "content_block_start", - "index": self.current_content_block_index, - "content_block": {"type": "text", "text": ""}, - } - ) - return self.chunk_queue.popleft() - async for chunk in self.completion_stream: if chunk == "None" or chunk is None: raise Exception - # Check if we need to start a new content block should_start_new_block = self._should_start_new_content_block(chunk) - if should_start_new_block: + is_opening_first_block = self.sent_content_block_start is False + if is_opening_first_block and self._is_blank_delta(chunk): + continue + if is_opening_first_block: + self.sent_content_block_start = True + self.sent_content_block_finish = False + self.chunk_queue.append( + { + "type": "content_block_start", + "index": self.current_content_block_index, + "content_block": self.current_content_block_start, + } + ) + elif should_start_new_block: self._increment_content_block_index() # applied_edits only needs to flow to the final message_delta @@ -664,7 +665,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): # Check if this processed chunk has a stop_reason - hold it for next chunk if not self.queued_usage_chunk: - if should_start_new_block and not self.sent_content_block_finish: + if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish: # Queue the sequence: content_block_stop -> content_block_start # -> (optionally) the trigger chunk's delta. # @@ -875,6 +876,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): return False return bool(delta.get(_delta_payload_field(delta_type))) + @staticmethod + def _is_blank_delta(chunk: "ModelResponseStream") -> bool: + choice = chunk.choices[0] + if choice.finish_reason is not None: + return False + delta = choice.delta + if getattr(delta, "tool_calls", None): + return False + if getattr(delta, "content", None): + return False + if getattr(delta, "reasoning_content", None): + return False + if getattr(delta, "thinking_blocks", None): + return False + return True + def _should_start_new_content_block(self, chunk: "ModelResponseStream") -> bool: """ Determine if we should start a new content block based on the processed chunk. diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 4b6617fbeac..86c9c1db481 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -331,6 +331,7 @@ class LiteLLMAnthropicMessagesAdapter: "thinking", "output_format", "output_config", + "stop_sequences", ] def _is_web_search_tool(self, tool: Dict[str, Any]) -> bool: @@ -615,7 +616,7 @@ class LiteLLMAnthropicMessagesAdapter: thinking_type = thinking.get("type", "disabled") if thinking_type == "disabled": - return None + return "none" elif thinking_type == "enabled": return reasoning_effort_from_thinking_budget(thinking.get("budget_tokens", 0)) elif thinking_type == "adaptive": @@ -683,25 +684,37 @@ class LiteLLMAnthropicMessagesAdapter: thinking ) if reasoning_effort: - summary = thinking.get("summary") if isinstance(thinking, dict) else None - auto_summary = is_reasoning_auto_summary_enabled() - if summary: - return { - "reasoning_effort": { - "effort": reasoning_effort, - "summary": summary, - } - } - elif auto_summary: - return { - "reasoning_effort": { - "effort": reasoning_effort, - "summary": "detailed", - } - } - return {"reasoning_effort": reasoning_effort} + return { + "reasoning_effort": LiteLLMAnthropicMessagesAdapter._apply_reasoning_summary_wrapping( + reasoning_effort, thinking + ) + } return {} + @staticmethod + def _apply_reasoning_summary_wrapping( + reasoning_effort: str, + thinking: Dict[str, Any], + ) -> Any: + """ + Apply the reasoning_effort/summary wrapping rules shared by every + thinking->reasoning_effort translation path. + + Disabled thinking always stays a plain string - there's no reasoning + trace to summarize, and non-Claude providers (e.g. Fireworks) expect + reasoning_effort as a plain string, not a summary dict. + """ + thinking_type = thinking.get("type") if isinstance(thinking, dict) else None + if thinking_type == "disabled": + return reasoning_effort + + summary = thinking.get("summary") if isinstance(thinking, dict) else None + if summary: + return {"effort": reasoning_effort, "summary": summary} + if is_reasoning_auto_summary_enabled(): + return {"effort": reasoning_effort, "summary": "detailed"} + return reasoning_effort + def translate_anthropic_tool_choice_to_openai( self, tool_choice: AnthropicMessagesToolChoice ) -> ChatCompletionToolChoiceValues: @@ -919,6 +932,18 @@ class LiteLLMAnthropicMessagesAdapter: tool_choice=cast(AnthropicMessagesToolChoice, tool_choice) ) + def _translate_stop_sequences_to_openai( + self, + anthropic_message_request: AnthropicMessagesRequest, + new_kwargs: ChatCompletionRequest, + ) -> None: + if "stop_sequences" not in anthropic_message_request: + return + stop_sequences = anthropic_message_request["stop_sequences"] + if not stop_sequences: + return + new_kwargs["stop"] = stop_sequences + def _translate_tools_to_openai( self, anthropic_message_request: AnthropicMessagesRequest, @@ -976,32 +1001,17 @@ class LiteLLMAnthropicMessagesAdapter: if not reasoning_effort: return + thinking_type = thinking.get("type") if isinstance(thinking, dict) else None + # For adaptive thinking, override with output_config.effort if available - if isinstance(thinking, dict) and thinking.get("type") == "adaptive": + if thinking_type == "adaptive": output_config = anthropic_message_request.get("output_config") if isinstance(output_config, dict) and output_config.get("effort"): reasoning_effort = output_config["effort"] - summary = thinking.get("summary") if isinstance(thinking, dict) else None - auto_summary = is_reasoning_auto_summary_enabled() - if summary: - new_kwargs["reasoning_effort"] = cast( - Any, - { - "effort": reasoning_effort, - "summary": summary, - }, - ) - elif auto_summary: - new_kwargs["reasoning_effort"] = cast( - Any, - { - "effort": reasoning_effort, - "summary": "detailed", - }, - ) - else: - new_kwargs["reasoning_effort"] = reasoning_effort + new_kwargs["reasoning_effort"] = self._apply_reasoning_summary_wrapping( + reasoning_effort, cast(Dict[str, Any], thinking) + ) def _translate_output_format_to_openai( self, @@ -1098,6 +1108,11 @@ class LiteLLMAnthropicMessagesAdapter: anthropic_message_request=anthropic_message_request, new_kwargs=new_kwargs, ) + ## CONVERT STOP_SEQUENCES + self._translate_stop_sequences_to_openai( + anthropic_message_request=anthropic_message_request, + new_kwargs=new_kwargs, + ) ## CONVERT OUTPUT_FORMAT to RESPONSE_FORMAT self._translate_output_format_to_openai( anthropic_message_request=anthropic_message_request, diff --git a/litellm/llms/bedrock/chat/__init__.py b/litellm/llms/bedrock/chat/__init__.py index c1323b9192a..37dcb270743 100644 --- a/litellm/llms/bedrock/chat/__init__.py +++ b/litellm/llms/bedrock/chat/__init__.py @@ -5,7 +5,6 @@ from .invoke_handler import ( AmazonAnthropicClaudeStreamDecoder, AmazonDeepSeekR1StreamDecoder, AWSEventStreamDecoder, - BedrockLLM, ) diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 4c256be1ab8..c28627d5aec 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -1,19 +1,10 @@ -""" -TODO: DELETE FILE. Bedrock LLM is no longer used. Goto `litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py` -""" - -import copy -import time import types -from functools import partial from typing import ( AsyncIterator, - Callable, Iterator, Optional, Tuple, cast, - get_args, ) import httpx # type: ignore @@ -25,16 +16,6 @@ from litellm.caching.caching import InMemoryCache from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.litellm_core_utils.core_helpers import map_finish_reason from litellm.litellm_core_utils.litellm_logging import Logging -from litellm.litellm_core_utils.logging_utils import track_llm_api_timing -from litellm.litellm_core_utils.prompt_templates.factory import ( - cohere_message_pt, - construct_tool_use_system_prompt, - contains_tag, - custom_prompt, - extract_between_tags, - parse_xml_params, - prompt_factory, -) from litellm.llms.anthropic.chat.handler import ( ModelResponseIterator as AnthropicModelResponseIterator, ) @@ -64,12 +45,9 @@ from litellm.types.utils import ( StreamingChoices, Usage, ) -from litellm.utils import CustomStreamWrapper, get_secret -from ..base_aws_llm import BaseAWSLLM from ..common_utils import ( BedrockError, - ModelResponseIterator, build_bedrock_stream_error, get_bedrock_response_stream_shape, get_bedrock_tool_name, @@ -77,9 +55,6 @@ from ..common_utils import ( bedrock_tool_name_mappings: InMemoryCache = InMemoryCache(max_size_in_memory=50, default_ttl=600) from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig -from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import ( - AmazonBedrockOpenAIConfig, -) converse_config = AmazonConverseConfig() @@ -351,932 +326,6 @@ def make_sync_call( raise BedrockError(status_code=500, message=str(e)) -class BedrockLLM(BaseAWSLLM): - """ - Example call - - ``` - curl --location --request POST 'https://bedrock-runtime.{aws_region_name}.amazonaws.com/model/{bedrock_model_name}/invoke' \ - --header 'Content-Type: application/json' \ - --header 'Accept: application/json' \ - --user "$AWS_ACCESS_KEY_ID":"$AWS_SECRET_ACCESS_KEY" \ - --aws-sigv4 "aws:amz:us-east-1:bedrock" \ - --data-raw '{ - "prompt": "Hi", - "temperature": 0, - "p": 0.9, - "max_tokens": 4096 - }' - ``` - """ - - def __init__(self) -> None: - super().__init__() - - @staticmethod - def is_claude_messages_api_model(model: str) -> bool: - """ - Check if the model uses the Claude Messages API (Claude 3+). - - Handles: - - Regional prefixes: eu.anthropic.claude-*, us.anthropic.claude-* - - Claude 3 models: claude-3-haiku, claude-3-sonnet, claude-3-opus, claude-3-5-*, claude-3-7-* - - Claude 4 models: claude-opus-4, claude-sonnet-4, claude-haiku-4 - """ - # Normalize model string to lowercase for matching - model_lower = model.lower() - - # Claude 3+ indicators (all use Messages API) - messages_api_indicators = [ - "claude-3", # Claude 3.x models - "claude-opus-4", # Claude Opus 4 - "claude-sonnet-4", # Claude Sonnet 4 - "claude-haiku-4", # Claude Haiku 4 - ] - - return any(indicator in model_lower for indicator in messages_api_indicators) - - def convert_messages_to_prompt(self, model, messages, provider, custom_prompt_dict) -> Tuple[str, Optional[list]]: - # handle anthropic prompts and amazon titan prompts - prompt = "" - chat_history: Optional[list] = None - ## CUSTOM PROMPT - if model in custom_prompt_dict: - # check if the model has a registered custom prompt - model_prompt_details = custom_prompt_dict[model] - prompt = custom_prompt( - role_dict=model_prompt_details["roles"], - initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""), - final_prompt_value=model_prompt_details.get("final_prompt_value", ""), - messages=messages, - ) - return prompt, None - ## ELSE - if provider == "anthropic" or provider == "amazon": - prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock") - elif provider == "mistral": - prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock") - elif provider == "meta" or provider == "llama": - prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock") - elif provider == "openai": - # OpenAI uses messages directly, no prompt conversion needed - # Return empty prompt as it won't be used - prompt = "" - elif provider == "cohere": - prompt, chat_history = cohere_message_pt(messages=messages) - else: - prompt = "" - for message in messages: - if "role" in message: - if message["role"] == "user": - prompt += f"{message['content']}" - else: - prompt += f"{message['content']}" - else: - prompt += f"{message['content']}" - return prompt, chat_history # type: ignore - - def process_response( - self, - model: str, - response: httpx.Response, - model_response: ModelResponse, - stream: Optional[bool], - logging_obj: Logging, - optional_params: dict, - api_key: str, - data: Union[dict, str], - messages: List, - print_verbose, - encoding, - ) -> Union[ModelResponse, CustomStreamWrapper]: - provider = self.get_bedrock_invoke_provider(model) - ## LOGGING - logging_obj.post_call( - input=messages, - api_key=api_key, - original_response=response.text, - additional_args={"complete_input_dict": data}, - ) - print_verbose(f"raw model_response: {response.text}") - - ## RESPONSE OBJECT - try: - completion_response = response.json() - except Exception: - raise BedrockError(message=response.text, status_code=422) - - outputText: Optional[str] = None - try: - if provider == "cohere": - if "text" in completion_response: - outputText = completion_response["text"] # type: ignore - elif "generations" in completion_response: - outputText = completion_response["generations"][0]["text"] - model_response.choices[0].finish_reason = map_finish_reason( - completion_response["generations"][0]["finish_reason"] - ) - elif provider == "anthropic": - if self.is_claude_messages_api_model(model): - json_schemas: dict = {} - _is_function_call = False - ## Handle Tool Calling - if "tools" in optional_params: - _is_function_call = True - for tool in optional_params["tools"]: - json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None) - outputText = completion_response.get("content")[0].get("text", None) - if outputText is not None and contains_tag("invoke", outputText): # OUTPUT PARSE FUNCTION CALL - function_name = extract_between_tags("tool_name", outputText)[0] - function_arguments_str = extract_between_tags("invoke", outputText)[0].strip() - function_arguments_str = f"{function_arguments_str}" - function_arguments = parse_xml_params( - function_arguments_str, - json_schema=json_schemas.get( - function_name, None - ), # check if we have a json schema for this function name) - ) - _message = litellm.Message( - tool_calls=[ - { - "id": f"call_{uuid.uuid4()}", - "type": "function", - "function": { - "name": function_name, - "arguments": json.dumps(function_arguments), - }, - } - ], - content=None, - ) - model_response.choices[0].message = _message # type: ignore - model_response._hidden_params["original_response"] = ( - outputText # allow user to access raw anthropic tool calling response - ) - if _is_function_call is True and stream is not None and stream is True: - print_verbose("INSIDE BEDROCK STREAMING TOOL CALLING CONDITION BLOCK") - # return an iterator - streaming_model_response = ModelResponseStream() - streaming_model_response.choices[0].finish_reason = getattr( - model_response.choices[0], "finish_reason", "stop" - ) - # streaming_model_response.choices = [litellm.utils.StreamingChoices()] - streaming_choice = litellm.utils.StreamingChoices() - streaming_choice.index = model_response.choices[0].index - _tool_calls = [] - print_verbose(f"type of model_response.choices[0]: {type(model_response.choices[0])}") - print_verbose(f"type of streaming_choice: {type(streaming_choice)}") - if isinstance(model_response.choices[0], litellm.Choices): - if getattr( - model_response.choices[0].message, "tool_calls", None - ) is not None and isinstance(model_response.choices[0].message.tool_calls, list): - for tool_call in model_response.choices[0].message.tool_calls: - _tool_call = {**tool_call.dict(), "index": 0} - _tool_calls.append(_tool_call) - delta_obj = Delta( - content=getattr(model_response.choices[0].message, "content", None), - role=model_response.choices[0].message.role, - tool_calls=_tool_calls, - ) - streaming_choice.delta = delta_obj - streaming_model_response.choices = [streaming_choice] - completion_stream = ModelResponseIterator(model_response=streaming_model_response) - print_verbose( - "Returns anthropic CustomStreamWrapper with 'cached_response' streaming object" - ) - return litellm.CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="cached_response", - logging_obj=logging_obj, - ) - - model_response.choices[0].finish_reason = map_finish_reason( - completion_response.get("stop_reason", "") - ) - _usage = litellm.Usage( - prompt_tokens=completion_response["usage"]["input_tokens"], - completion_tokens=completion_response["usage"]["output_tokens"], - total_tokens=completion_response["usage"]["input_tokens"] - + completion_response["usage"]["output_tokens"], - ) - setattr(model_response, "usage", _usage) - else: - outputText = completion_response["completion"] - - model_response.choices[0].finish_reason = completion_response["stop_reason"] - elif provider == "ai21": - outputText = completion_response.get("completions")[0].get("data").get("text") - elif provider == "meta" or provider == "llama": - outputText = completion_response["generation"] - elif provider == "openai": - # OpenAI imported models use OpenAI Chat Completions format - if "choices" in completion_response and len(completion_response["choices"]) > 0: - choice = completion_response["choices"][0] - if "message" in choice: - outputText = choice["message"].get("content") - elif "text" in choice: # fallback for completion format - outputText = choice["text"] - - # Set finish reason - if "finish_reason" in choice: - model_response.choices[0].finish_reason = map_finish_reason(choice["finish_reason"]) - - # Set usage if available - if "usage" in completion_response: - usage = completion_response["usage"] - _usage = litellm.Usage( - prompt_tokens=usage.get("prompt_tokens", 0), - completion_tokens=usage.get("completion_tokens", 0), - total_tokens=usage.get("total_tokens", 0), - ) - setattr(model_response, "usage", _usage) - elif provider == "mistral": - outputText = completion_response["outputs"][0]["text"] - model_response.choices[0].finish_reason = completion_response["outputs"][0]["stop_reason"] - else: # amazon titan - outputText = completion_response.get("results")[0].get("outputText") - except Exception as e: - raise BedrockError( - message="Error processing={}, Received error={}".format(response.text, str(e)), - status_code=422, - ) - - try: - if ( - outputText is not None - and len(outputText) > 0 - and hasattr(model_response.choices[0], "message") - and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore - is None - ): - model_response.choices[0].message.content = outputText # type: ignore - elif ( - hasattr(model_response.choices[0], "message") - and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore - is not None - ): - pass - else: - raise Exception() - except Exception as e: - raise BedrockError( - message="Error parsing received text={}.\nError-{}".format(outputText, str(e)), - status_code=response.status_code, - ) - - if stream and provider == "ai21": - streaming_model_response = ModelResponseStream() - streaming_model_response.choices[0].finish_reason = model_response.choices[ # type: ignore - 0 - ].finish_reason - # streaming_model_response.choices = [litellm.utils.StreamingChoices()] - streaming_choice = litellm.utils.StreamingChoices() - streaming_choice.index = model_response.choices[0].index - delta_obj = litellm.utils.Delta( - content=getattr(model_response.choices[0].message, "content", None), # type: ignore - role=model_response.choices[0].message.role, # type: ignore - ) - streaming_choice.delta = delta_obj - streaming_model_response.choices = [streaming_choice] - mri = ModelResponseIterator(model_response=streaming_model_response) - return CustomStreamWrapper( - completion_stream=mri, - model=model, - custom_llm_provider="cached_response", - logging_obj=logging_obj, - ) - - ## CALCULATING USAGE - bedrock returns usage in the headers - # Skip if usage was already set (e.g., from JSON response for OpenAI provider) - if not hasattr(model_response, "usage") or getattr(model_response, "usage", None) is None: - bedrock_input_tokens = response.headers.get("x-amzn-bedrock-input-token-count", None) - bedrock_output_tokens = response.headers.get("x-amzn-bedrock-output-token-count", None) - - prompt_tokens = int(bedrock_input_tokens or litellm.token_counter(messages=messages)) - - completion_tokens = int( - bedrock_output_tokens - or litellm.token_counter( - text=model_response.choices[0].message.content, # type: ignore - count_response_tokens=True, - ) - ) - - model_response.created = int(time.time()) - model_response.model = model - usage = Usage( - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - total_tokens=prompt_tokens + completion_tokens, - ) - setattr(model_response, "usage", usage) - else: - # Ensure created and model are set even if usage was already set - model_response.created = int(time.time()) - model_response.model = model - - return model_response - - def completion( - self, - model: str, - messages: list, - api_base: Optional[str], - custom_prompt_dict: dict, - model_response: ModelResponse, - print_verbose: Callable, - encoding, - logging_obj: Logging, - optional_params: dict, - acompletion: bool, - timeout: Optional[Union[float, httpx.Timeout]], - litellm_params=None, - logger_fn=None, - extra_headers: Optional[dict] = None, - client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None, - ) -> Union[ModelResponse, CustomStreamWrapper]: - try: - from botocore.credentials import Credentials - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - - ## SETUP ## - stream = optional_params.pop("stream", None) - stream_chunk_size = optional_params.pop("stream_chunk_size", None) - - provider = self.get_bedrock_invoke_provider(model) - modelId = self.get_bedrock_model_id( - model=model, - provider=provider, - optional_params=optional_params, - ) - - ## CREDENTIALS ## - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them - aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) - aws_access_key_id = optional_params.pop("aws_access_key_id", None) - aws_session_token = optional_params.pop("aws_session_token", None) - aws_region_name = optional_params.pop("aws_region_name", None) - aws_role_name = optional_params.pop("aws_role_name", None) - aws_session_name = optional_params.pop("aws_session_name", None) - aws_profile_name = optional_params.pop("aws_profile_name", None) - aws_bedrock_runtime_endpoint = optional_params.pop( - "aws_bedrock_runtime_endpoint", None - ) # https://bedrock-runtime.{region_name}.amazonaws.com - aws_web_identity_token = optional_params.pop("aws_web_identity_token", None) - aws_sts_endpoint = optional_params.pop("aws_sts_endpoint", None) - ssl_verify = optional_params.pop("ssl_verify", None) - - ### SET REGION NAME ### - if aws_region_name is None: - # check env # - litellm_aws_region_name = get_secret("AWS_REGION_NAME", None) - - if litellm_aws_region_name is not None and isinstance(litellm_aws_region_name, str): - aws_region_name = litellm_aws_region_name - - standard_aws_region_name = get_secret("AWS_REGION", None) - if standard_aws_region_name is not None and isinstance(standard_aws_region_name, str): - aws_region_name = standard_aws_region_name - - if aws_region_name is None: - aws_region_name = "us-west-2" - - credentials: Credentials = self.get_credentials( - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - aws_region_name=aws_region_name, - aws_session_name=aws_session_name, - aws_profile_name=aws_profile_name, - aws_role_name=aws_role_name, - aws_web_identity_token=aws_web_identity_token, - aws_sts_endpoint=aws_sts_endpoint, - ssl_verify=ssl_verify, - ) - - ### SET RUNTIME ENDPOINT ### - endpoint_url, proxy_endpoint_url = self.get_runtime_endpoint( - api_base=api_base, - aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, - aws_region_name=aws_region_name, - ) - - if (stream is not None and stream is True) and provider != "ai21": - endpoint_url = f"{endpoint_url}/model/{modelId}/invoke-with-response-stream" - proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke-with-response-stream" - else: - endpoint_url = f"{endpoint_url}/model/{modelId}/invoke" - proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke" - - if acompletion and provider == "anthropic" and self.is_claude_messages_api_model(model): - if isinstance(client, HTTPHandler): - client = None - return self._async_anthropic_messages_completion( - model=model, - messages=messages, - endpoint_url=endpoint_url, - proxy_endpoint_url=proxy_endpoint_url, - credentials=credentials, - aws_region_name=aws_region_name, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - litellm_params=litellm_params, - logger_fn=logger_fn, - extra_headers=extra_headers, - timeout=timeout, - client=client, - stream_chunk_size=stream_chunk_size, - ) # type: ignore[return-value] - - prompt, chat_history = self.convert_messages_to_prompt(model, messages, provider, custom_prompt_dict) - inference_params = copy.deepcopy(optional_params) - json_schemas: dict = {} - if provider == "cohere": - if model.startswith("cohere.command-r"): - ## LOAD CONFIG - config = litellm.AmazonCohereChatConfig().get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - _data = {"message": prompt, **inference_params} - if chat_history is not None: - _data["chat_history"] = chat_history - data = json.dumps(_data) - else: - ## LOAD CONFIG - config = litellm.AmazonCohereConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - if stream is True: - inference_params["stream"] = True # cohere requires stream = True in inference params - data = json.dumps({"prompt": prompt, **inference_params}) - elif provider == "anthropic": - if self.is_claude_messages_api_model(model): - # Separate system prompt from rest of message - system_prompt_idx: list[int] = [] - system_messages: list[str] = [] - for idx, message in enumerate(messages): - if message["role"] == "system": - system_messages.append(message["content"]) - system_prompt_idx.append(idx) - if len(system_prompt_idx) > 0: - inference_params["system"] = "\n".join(system_messages) - messages = [i for j, i in enumerate(messages) if j not in system_prompt_idx] - # Format rest of message according to anthropic guidelines - messages = prompt_factory(model=model, messages=messages, custom_llm_provider="anthropic_xml") # type: ignore - ## LOAD CONFIG - config = litellm.AmazonAnthropicClaudeConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - ## Handle Tool Calling - if "tools" in inference_params: - _is_function_call = True - for tool in inference_params["tools"]: - json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None) - tool_calling_system_prompt = construct_tool_use_system_prompt(tools=inference_params["tools"]) - inference_params["system"] = ( - inference_params.get("system", "\n") + tool_calling_system_prompt - ) # add the anthropic tool calling prompt to the system prompt - inference_params.pop("tools") - data = json.dumps({"messages": messages, **inference_params}) - else: - ## LOAD CONFIG - config = litellm.AmazonAnthropicConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - data = json.dumps({"prompt": prompt, **inference_params}) - elif provider == "ai21": - ## LOAD CONFIG - config = litellm.AmazonAI21Config.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - - data = json.dumps({"prompt": prompt, **inference_params}) - elif provider == "mistral": - ## LOAD CONFIG - config = litellm.AmazonMistralConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - - data = json.dumps({"prompt": prompt, **inference_params}) - elif provider == "amazon": # amazon titan - ## LOAD CONFIG - config = litellm.AmazonTitanConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - - data = json.dumps( - { - "inputText": prompt, - "textGenerationConfig": inference_params, - } - ) - elif provider == "meta" or provider == "llama": - ## LOAD CONFIG - config = litellm.AmazonLlamaConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - data = json.dumps({"prompt": prompt, **inference_params}) - elif provider == "openai": - ## OpenAI imported models use OpenAI Chat Completions format (messages-based) - # Use AmazonBedrockOpenAIConfig for proper OpenAI transformation - openai_config = AmazonBedrockOpenAIConfig() - supported_params = openai_config.get_supported_openai_params(model=model) - - # Filter to only supported OpenAI params - filtered_params = {k: v for k, v in inference_params.items() if k in supported_params} - - # OpenAI uses messages format, not prompt - data = json.dumps({"messages": messages, **filtered_params}) - else: - ## LOGGING - logging_obj.pre_call( - input=messages, - api_key="", - additional_args={ - "complete_input_dict": inference_params, - }, - ) - raise BedrockError( - status_code=404, - message="Bedrock Invoke HTTPX: Unknown provider={}, model={}. Try calling via converse route - `bedrock/converse/`.".format( - provider, model - ), - ) - - ## COMPLETION CALL - - headers = {"Content-Type": "application/json"} - if extra_headers is not None: - headers = {"Content-Type": "application/json", **extra_headers} - prepped = self.get_request_headers( - credentials=credentials, - aws_region_name=aws_region_name, - extra_headers=extra_headers, - endpoint_url=endpoint_url, - data=data, - headers=headers, - ) - - ## LOGGING - logging_obj.pre_call( - input=messages, - api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": proxy_endpoint_url, - "headers": prepped.headers, - }, - ) - - ### ROUTING (ASYNC, STREAMING, SYNC) - if acompletion: - if isinstance(client, HTTPHandler): - client = None - if stream is True and provider != "ai21": - return self.async_streaming( - model=model, - messages=messages, - data=data, - api_base=proxy_endpoint_url, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=True, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=prepped.headers, - timeout=timeout, - client=client, - stream_chunk_size=stream_chunk_size, - ) # type: ignore - ### ASYNC COMPLETION - return self.async_completion( - model=model, - messages=messages, - data=data, - api_base=proxy_endpoint_url, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, # type: ignore - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=prepped.headers, - timeout=timeout, - client=client, - ) # type: ignore - - if client is None or isinstance(client, AsyncHTTPHandler): - _params = {} - if timeout is not None: - if isinstance(timeout, float) or isinstance(timeout, int): - timeout = httpx.Timeout(timeout) - _params["timeout"] = timeout - self.client = _get_httpx_client(_params) # type: ignore - else: - self.client = client - if (stream is not None and stream is True) and provider != "ai21": - response = self.client.post( - url=proxy_endpoint_url, - headers=prepped.headers, # type: ignore - data=data, - stream=stream, - logging_obj=logging_obj, - ) - - if response.status_code != 200: - raise BedrockError(status_code=response.status_code, message=str(response.read())) - - decoder = AWSEventStreamDecoder(model=model) - - completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size)) - streaming_response = CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="bedrock", - logging_obj=logging_obj, - ) - - ## LOGGING - logging_obj.post_call( - input=messages, - api_key="", - original_response=streaming_response, - additional_args={"complete_input_dict": data}, - ) - return streaming_response - - try: - response = self.client.post( - url=proxy_endpoint_url, - headers=dict(prepped.headers), - data=data, - logging_obj=logging_obj, - ) - response.raise_for_status() - except httpx.HTTPStatusError as err: - error_code = err.response.status_code - raise BedrockError(status_code=error_code, message=err.response.text) - except httpx.TimeoutException: - raise BedrockError(status_code=408, message="Timeout error occurred.") - - return self.process_response( - model=model, - response=response, - model_response=model_response, - stream=stream, - logging_obj=logging_obj, - optional_params=optional_params, - api_key="", - data=data, - messages=messages, - print_verbose=print_verbose, - encoding=encoding, - ) - - async def _async_anthropic_messages_completion( - self, - model: str, - messages: list, - endpoint_url: str, - proxy_endpoint_url: str, - credentials, - aws_region_name: str, - model_response: ModelResponse, - print_verbose: Callable, - encoding, - logging_obj: Logging, - optional_params: dict, - stream, - litellm_params=None, - logger_fn=None, - extra_headers: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, - stream_chunk_size: Optional[int] = None, - ) -> Union[ModelResponse, CustomStreamWrapper]: - transformed_request = await litellm.AmazonAnthropicClaudeConfig().async_transform_request( - model=model, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params or {}, - headers=extra_headers or {}, - ) - data = json.dumps(transformed_request) - - headers = {"Content-Type": "application/json"} - if extra_headers is not None: - headers = {"Content-Type": "application/json", **extra_headers} - prepped = self.get_request_headers( - credentials=credentials, - aws_region_name=aws_region_name, - extra_headers=extra_headers, - endpoint_url=endpoint_url, - data=data, - headers=headers, - ) - - logging_obj.pre_call( - input=messages, - api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": proxy_endpoint_url, - "headers": prepped.headers, - }, - ) - - if stream is True: - return await self.async_streaming( - model=model, - messages=messages, - data=data, - api_base=proxy_endpoint_url, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=True, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=prepped.headers, - timeout=timeout, - client=client, - stream_chunk_size=stream_chunk_size, - ) - return await self.async_completion( - model=model, - messages=messages, - data=data, - api_base=proxy_endpoint_url, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, # type: ignore - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=prepped.headers, - timeout=timeout, - client=client, - ) - - async def async_completion( - self, - model: str, - messages: list, - api_base: str, - model_response: ModelResponse, - print_verbose: Callable, - data: str, - timeout: Optional[Union[float, httpx.Timeout]], - encoding, - logging_obj: Logging, - stream, - optional_params: dict, - litellm_params=None, - logger_fn=None, - headers={}, - client: Optional[AsyncHTTPHandler] = None, - ) -> Union[ModelResponse, CustomStreamWrapper]: - if client is None: - _params = {} - if timeout is not None: - if isinstance(timeout, float) or isinstance(timeout, int): - timeout = httpx.Timeout(timeout) - _params["timeout"] = timeout - client = get_async_httpx_client(params=_params, llm_provider=litellm.LlmProviders.BEDROCK) # type: ignore - else: - client = client # type: ignore - - try: - response = await client.post( - api_base, - headers=headers, - data=data, - timeout=timeout, - logging_obj=logging_obj, - ) - response.raise_for_status() - except httpx.HTTPStatusError as err: - error_code = err.response.status_code - raise BedrockError(status_code=error_code, message=err.response.text) - except httpx.TimeoutException: - raise BedrockError(status_code=408, message="Timeout error occurred.") - - return self.process_response( - model=model, - response=response, - model_response=model_response, - stream=stream if isinstance(stream, bool) else False, - logging_obj=logging_obj, - api_key="", - data=data, - messages=messages, - print_verbose=print_verbose, - optional_params=optional_params, - encoding=encoding, - ) - - @track_llm_api_timing() # for streaming, we need to instrument the function calling the wrapper - async def async_streaming( - self, - model: str, - messages: list, - api_base: str, - model_response: ModelResponse, - print_verbose: Callable, - data: str, - timeout: Optional[Union[float, httpx.Timeout]], - encoding, - logging_obj: Logging, - stream, - optional_params: dict, - litellm_params=None, - logger_fn=None, - headers={}, - client: Optional[AsyncHTTPHandler] = None, - stream_chunk_size: Optional[int] = None, - ) -> CustomStreamWrapper: - # The call is not made here; instead, we prepare the necessary objects for the stream. - - streaming_response = CustomStreamWrapper( - completion_stream=None, - make_call=partial( - make_call, - client=client, - api_base=api_base, - headers=headers, - data=data, # type: ignore - model=model, - messages=messages, - logging_obj=logging_obj, - fake_stream=True if "ai21" in api_base else False, - stream_chunk_size=stream_chunk_size, - ), - model=model, - custom_llm_provider="bedrock", - logging_obj=logging_obj, - ) - return streaming_response - - @staticmethod - def _get_provider_from_model_path( - model_path: str, - ) -> Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL]: - """ - Helper function to get the provider from a model path with format: provider/model-name - - Args: - model_path (str): The model path (e.g., 'llama/arn:aws:bedrock:us-east-1:086734376398:imported-model/r4c4kewx2s0n' or 'anthropic/model-name') - - Returns: - Optional[str]: The provider name, or None if no valid provider found - """ - parts = model_path.split("/") - if len(parts) >= 1: - provider = parts[0] - if provider in get_args(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL): - return cast(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL, provider) - return None - - class AWSEventStreamDecoder: def __init__(self, model: str, json_mode: Optional[bool] = False) -> None: from botocore.parsers import EventStreamJSONParser diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 5114677ffc0..93998f0610e 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -1109,8 +1109,10 @@ def get_bedrock_chat_config(model: str): Returns: The appropriate Bedrock config class instance """ + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + bedrock_route = BedrockModelInfo.get_bedrock_route(model) - bedrock_invoke_provider = litellm.BedrockLLM.get_bedrock_invoke_provider(model=model) + bedrock_invoke_provider = BaseAWSLLM.get_bedrock_invoke_provider(model=model) base_model = BedrockModelInfo.get_base_model(model) # Handle explicit routes first diff --git a/litellm/llms/custom_httpx/aiohttp_transport.py b/litellm/llms/custom_httpx/aiohttp_transport.py index adac6a1b276..df5b10b3bdc 100644 --- a/litellm/llms/custom_httpx/aiohttp_transport.py +++ b/litellm/llms/custom_httpx/aiohttp_transport.py @@ -143,13 +143,26 @@ class LiteLLMAiohttpTransport(AiohttpTransport): client: Union[ClientSession, Callable[[], ClientSession]], ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None, owns_session: bool = True, + session_factory: Callable[[], ClientSession] | None = None, ): self.client = client self._ssl_verify = ssl_verify # Store for per-request SSL override super().__init__(client=client, owns_session=owns_session) # Store the client factory for recreating sessions when needed - if callable(client): - self._client_factory = client + default_factory: Callable[[], ClientSession] = client if callable(client) else ClientSession + self._client_factory: Callable[[], ClientSession] = session_factory or default_factory + + def _rebuild_session(self) -> ClientSession: + """ + Build a replacement session from the configured factory. + + The replacement is reachable only from this transport, so the transport + owns it from here on even when it was originally handed a session it did + not own (the proxy's shared session). + """ + session = self._client_factory() + self._owns_session = True + return session def _get_valid_client_session(self) -> ClientSession: """ @@ -158,24 +171,16 @@ class LiteLLMAiohttpTransport(AiohttpTransport): This handles the case where the session was created in a different event loop that may have been closed (common in CI/CD environments). """ - from aiohttp.client import ClientSession - # If we don't have a client or it's not a ClientSession, create one if not isinstance(self.client, ClientSession): - if hasattr(self, "_client_factory") and callable(self._client_factory): - self.client = self._client_factory() - else: - self.client = ClientSession() + self.client = self._rebuild_session() # Don't return yet - check if the newly created session is valid # Check if the session itself is closed if self.client.closed: verbose_logger.debug("Session is closed, creating new session") # Create a new session - if hasattr(self, "_client_factory") and callable(self._client_factory): - self.client = self._client_factory() - else: - self.client = ClientSession() + self.client = self._rebuild_session() return self.client # Check if the existing session is still valid for the current event loop @@ -188,7 +193,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport): # Close old session to prevent leaks old_session = self.client try: - if not old_session.closed: + if self._owns_session and not old_session.closed: try: asyncio.create_task(old_session.close()) except RuntimeError: @@ -198,17 +203,11 @@ class LiteLLMAiohttpTransport(AiohttpTransport): verbose_logger.debug(f"Error closing old session: {e}") # Create a new session in the current event loop - if hasattr(self, "_client_factory") and callable(self._client_factory): - self.client = self._client_factory() - else: - self.client = ClientSession() + self.client = self._rebuild_session() except (RuntimeError, AttributeError): # If we can't check the loop or session is invalid, recreate it - if hasattr(self, "_client_factory") and callable(self._client_factory): - self.client = self._client_factory() - else: - self.client = ClientSession() + self.client = self._rebuild_session() return self.client @@ -303,10 +302,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport): if "Session is closed" in str(e): verbose_logger.debug(f"Session closed during request, retrying with new session: {e}") # Force creation of a new session - if hasattr(self, "_client_factory") and callable(self._client_factory): - self.client = self._client_factory() - else: - self.client = ClientSession() + self.client = self._rebuild_session() client_session = self.client # Retry the request with the new session diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 5cec763bb5d..c92f8bcc937 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1013,17 +1013,6 @@ class AsyncHTTPHandler: verbose_logger.debug("Creating AiohttpTransport...") - # Use shared session if provided and valid - if shared_session is not None and not shared_session.closed: - verbose_logger.debug(f"SHARED SESSION: Reusing existing ClientSession (ID: {id(shared_session)})") - return LiteLLMAiohttpTransport( - client=shared_session, - ssl_verify=ssl_for_transport, - owns_session=False, - ) - - # Create new session only if none provided or existing one is invalid - verbose_logger.debug("NEW SESSION: Creating new ClientSession (no shared session provided)") transport_connector_kwargs = { "keepalive_timeout": AIOHTTP_KEEPALIVE_TIMEOUT, "ttl_dns_cache": AIOHTTP_TTL_DNS_CACHE, @@ -1041,11 +1030,26 @@ class AsyncHTTPHandler: if socket_factory is not None: transport_connector_kwargs["socket_factory"] = socket_factory - return LiteLLMAiohttpTransport( - client=lambda: ClientSession( + def session_factory() -> ClientSession: + return ClientSession( connector=TCPConnector(**transport_connector_kwargs), trust_env=trust_env, - ), + ) + + # Use shared session if provided and valid + if shared_session is not None and not shared_session.closed: + verbose_logger.debug(f"SHARED SESSION: Reusing existing ClientSession (ID: {id(shared_session)})") + return LiteLLMAiohttpTransport( + client=shared_session, + ssl_verify=ssl_for_transport, + owns_session=False, + session_factory=session_factory, + ) + + # Create new session only if none provided or existing one is invalid + verbose_logger.debug("NEW SESSION: Creating new ClientSession (no shared session provided)") + return LiteLLMAiohttpTransport( + client=session_factory, ssl_verify=ssl_for_transport, ) diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 6b191144a11..8fc7e6d0ebd 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -1626,7 +1626,7 @@ class OpenAIFilesAPI(BaseLLM): openai_client: AsyncOpenAI, ) -> OpenAIFileObject: response = await openai_client.files.create(**create_file_data) # type: ignore[arg-type] - return OpenAIFileObject(**response.model_dump()) + return OpenAIFileObject.model_validate(response.model_dump()) def create_file( self, @@ -1662,7 +1662,7 @@ class OpenAIFilesAPI(BaseLLM): create_file_data=create_file_data, openai_client=openai_client ) response = cast(OpenAI, openai_client).files.create(**create_file_data) # type: ignore[arg-type] - return OpenAIFileObject(**response.model_dump()) + return OpenAIFileObject.model_validate(response.model_dump()) async def afile_content( self, @@ -1986,7 +1986,7 @@ class OpenAIBatchesAPI(BaseLLM): openai_client: AsyncOpenAI, ) -> LiteLLMBatch: response = await openai_client.batches.create(**create_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) def create_batch( self, @@ -2023,7 +2023,7 @@ class OpenAIBatchesAPI(BaseLLM): ) response = cast(OpenAI, openai_client).batches.create(**create_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) async def aretrieve_batch( self, @@ -2032,7 +2032,7 @@ class OpenAIBatchesAPI(BaseLLM): ) -> LiteLLMBatch: verbose_logger.debug("retrieving batch, args= %s", retrieve_batch_data) response = await openai_client.batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) def retrieve_batch( self, @@ -2068,7 +2068,7 @@ class OpenAIBatchesAPI(BaseLLM): retrieve_batch_data=retrieve_batch_data, openai_client=openai_client ) response = cast(OpenAI, openai_client).batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) async def acancel_batch( self, @@ -2077,7 +2077,7 @@ class OpenAIBatchesAPI(BaseLLM): ) -> LiteLLMBatch: verbose_logger.debug("async cancelling batch, args= %s", cancel_batch_data) response = await openai_client.batches.cancel(**cancel_batch_data) - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) def cancel_batch( self, @@ -2117,7 +2117,7 @@ class OpenAIBatchesAPI(BaseLLM): if not isinstance(openai_client, OpenAI): raise ValueError("OpenAI client is not an instance of OpenAI. Make sure you passed a sync OpenAI client.") response = openai_client.batches.cancel(**cancel_batch_data) - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) async def alist_batches( self, @@ -2477,9 +2477,9 @@ class OpenAIAssistantsAPI(BaseLLM): response_obj: Optional[OpenAIMessage] = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) else: - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) return response_obj # fmt: off @@ -2556,9 +2556,9 @@ class OpenAIAssistantsAPI(BaseLLM): response_obj: Optional[OpenAIMessage] = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) else: - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) return response_obj async def async_get_messages( diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 3c2ae238a0b..dc4e98e6216 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -280,7 +280,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) try: - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) except Exception: verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct") response = ResponsesAPIResponse.model_construct(**raw_response_json) @@ -506,7 +506,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code) raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) response._hidden_params["additional_headers"] = processed_headers response._hidden_params["headers"] = raw_response_headers @@ -588,7 +588,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) response._hidden_params["additional_headers"] = processed_headers response._hidden_params["headers"] = raw_response_headers @@ -647,7 +647,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): processed_headers = process_response_headers(raw_response_headers) try: - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) except Exception: verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct") response = ResponsesAPIResponse.model_construct(**raw_response_json) diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index f0ce3323ef6..36c78974aca 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works """ import re -from typing import List, Optional, Tuple, Literal +from typing import List, Optional, Sequence, Tuple, Literal from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.vertex_ai import CachedContentRequestBody @@ -152,6 +152,20 @@ def separate_cached_messages( return cached_messages, non_cached_messages +def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageValues]) -> bool: + """ + The cachedContents API rejects contents ending on a model turn, which is how it + classifies both assistant messages and tool results, with HTTP 400 + "Requests ending with a model turn are not supported". System messages are + extracted into system_instruction before contents are built, so the terminal + turn is the last non-system message. + """ + non_system_messages = tuple(message for message in cached_messages if message.get("role") != "system") + if not non_system_messages: + return bool(cached_messages) + return non_system_messages[-1].get("role") not in ("assistant", "tool", "function") + + def transform_openai_messages_to_gemini_context_caching( model: str, messages: List[AllMessageValues], diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index 0bf3715f798..f8774e33ca4 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -22,6 +22,7 @@ from litellm.types.llms.vertex_ai import ( from ..common_utils import VertexAIError, get_vertex_base_url from ..vertex_llm_base import VertexBase from .transformation import ( + cached_messages_end_on_supported_turn, separate_cached_messages, transform_openai_messages_to_gemini_context_caching, ) @@ -308,6 +309,14 @@ class ContextCachingEndpoints(VertexBase): if len(cached_messages) == 0: return messages, optional_params, None + if not cached_messages_end_on_supported_turn(cached_messages): + verbose_logger.debug( + "Vertex AI context caching: cached message block ends on a model turn once " + "system messages are extracted, which the cachedContents API rejects. " + "Skipping context caching." + ) + return messages, optional_params, None + # Gemini requires a minimum of 1024 tokens for context caching. # Skip caching if the cached content is too small to avoid API errors. if not is_prompt_caching_valid_prompt( @@ -459,6 +468,14 @@ class ContextCachingEndpoints(VertexBase): if len(cached_messages) == 0: return messages, optional_params, None + if not cached_messages_end_on_supported_turn(cached_messages): + verbose_logger.debug( + "Vertex AI context caching: cached message block ends on a model turn once " + "system messages are extracted, which the cachedContents API rejects. " + "Skipping context caching." + ) + return messages, optional_params, None + # Gemini requires a minimum of 1024 tokens for context caching. # Skip caching if the cached content is too small to avoid API errors. if not is_prompt_caching_valid_prompt( diff --git a/litellm/llms/vertex_ai/files/handler.py b/litellm/llms/vertex_ai/files/handler.py index 3bc09139f8f..4d2a1e18eb5 100644 --- a/litellm/llms/vertex_ai/files/handler.py +++ b/litellm/llms/vertex_ai/files/handler.py @@ -1,7 +1,9 @@ import asyncio +import json +import os import time from urllib.parse import unquote -from typing import Any, Coroutine, Optional, Tuple, Union +from typing import Any, Coroutine, Mapping, Optional, Tuple, Union import httpx @@ -10,6 +12,7 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import ( GCSBucketBase, GCSLoggingConfig, ) +from litellm.types.utils import StandardCallbackDynamicParams from litellm.litellm_core_utils.cloud_storage_security import ( VERTEX_AI_MANAGED_GCS_PREFIX, should_allow_legacy_cloud_file_ids, @@ -39,6 +42,35 @@ class VertexAIFilesHandler(GCSBucketBase): llm_provider=LlmProviders.VERTEX_AI, ) + def _resolve_read_gcs_config( + self, + litellm_params: Mapping[str, object] | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + ) -> tuple[str | None, str | None]: + """ + Resolve the GCS bucket and service-account credentials for the read/content path. + + Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` / + ``bucket_name`` and ``vertex_credentials``), mirroring the write path in + ``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global + ``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch + run entirely at the model-group level, so output written to a per-model bucket is + readable without setting the global env vars. + """ + params: Mapping[str, object] = litellm_params or {} + bucket_candidate = params.get("gcs_bucket_name") or params.get("bucket_name") + configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME") + + credentials = params.get("vertex_credentials") or vertex_credentials + if isinstance(credentials, dict): + path_service_account: str | None = json.dumps(credentials) + elif isinstance(credentials, str): + path_service_account = credentials + else: + path_service_account = os.getenv("GCS_PATH_SERVICE_ACCOUNT") + + return configured_bucket_name, path_service_account + def _extract_bucket_and_object_from_file_id( self, file_id: str, @@ -91,7 +123,17 @@ class VertexAIFilesHandler(GCSBucketBase): if not file_id: raise ValueError("file_id is required in file_content_request") - gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(kwargs={}) + configured_bucket_name, path_service_account = self._resolve_read_gcs_config( + litellm_params=litellm_params, + vertex_credentials=vertex_credentials, + ) + dynamic_params = StandardCallbackDynamicParams( + gcs_bucket_name=configured_bucket_name, + gcs_path_service_account=path_service_account, + ) + gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config( + kwargs={"standard_callback_dynamic_params": dynamic_params} + ) bucket_name, object_path = self._extract_bucket_and_object_from_file_id( file_id=file_id, configured_bucket_name=gcs_logging_config["bucket_name"], diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 0db1118a7b4..cbca57c5e62 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -661,6 +661,10 @@ def _gemini_convert_messages_with_history( vertex_project = litellm_params.get("vertex_project") or litellm_params.get("vertex_ai_project") vertex_credentials = litellm_params.get("vertex_credentials") or litellm_params.get("vertex_ai_credentials") + from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig + + forward_function_call_id = VertexGeminiConfig._forward_gemini_function_call_id(model or "") + try: while msg_i < len(messages): user_content: List[PartType] = [] @@ -910,7 +914,7 @@ def _gemini_convert_messages_with_history( gemini_tool_call_parts = convert_to_gemini_tool_call_invoke( assistant_msg, model=model, - custom_llm_provider=custom_llm_provider, + forward_function_call_id=forward_function_call_id, ) ## check if gemini_tool_call already exists in assistant_content for gemini_tool_call_part in gemini_tool_call_parts: @@ -973,8 +977,7 @@ def _gemini_convert_messages_with_history( _part = convert_to_gemini_tool_call_result( messages[msg_i], # type: ignore last_message_with_tool_calls, # type: ignore - model=model, - custom_llm_provider=custom_llm_provider, + forward_function_call_id=forward_function_call_id, ) msg_i += 1 # Handle both single part and list of parts (for Computer Use with images) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 8dd0dc19b81..126f82436e8 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -289,15 +289,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return False @staticmethod - def _forward_gemini_function_call_id(model: str, custom_llm_provider: Optional[str] = None) -> bool: + def _forward_gemini_function_call_id(model: str) -> bool: """ Whether to include `id` on function_call / function_response parts. - Gemini 3+ on Google AI Studio accepts (and returns) `id` for strict - tool-call matching. Vertex AI rejects the field with HTTP 400. + Gemini 3+ accepts (and returns) `id` for strict tool-call matching, on Vertex AI and + Google AI Studio alike. Older Gemini models reject the field with HTTP 400. """ - if custom_llm_provider != "gemini": - return False return VertexGeminiConfig._is_gemini_3_or_newer(model) def _supports_penalty_parameters(self, model: str) -> bool: diff --git a/litellm/main.py b/litellm/main.py index dc3ec469a1b..acdec7385da 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -207,7 +207,7 @@ from .llms.azure.chat.o_series_handler import AzureOpenAIO1ChatCompletion from .llms.azure.completion.handler import AzureTextCompletion from .llms.azure_ai.anthropic.handler import AzureAnthropicChatCompletion from .llms.azure_ai.embed import AzureAIEmbedding -from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM +from .llms.bedrock.chat import BedrockConverseLLM from .llms.bedrock.embed.embedding import BedrockEmbedding from .llms.bedrock.image_edit.handler import BedrockImageEdit from .llms.bedrock.image_generation.image_handler import BedrockImageGeneration diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 2cef600ea32..cc418c9c428 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3454,22 +3454,16 @@ }, "azure_ai/gpt-5.4-mini": { "cache_read_input_token_cost": 7.5e-08, - "cache_read_input_token_cost_above_272k_tokens": 1.5e-07, "cache_read_input_token_cost_priority": 1.5e-07, - "cache_read_input_token_cost_above_272k_tokens_priority": 3e-07, "input_cost_per_token": 7.5e-07, - "input_cost_per_token_above_272k_tokens": 1.5e-06, "input_cost_per_token_priority": 1.5e-06, - "input_cost_per_token_above_272k_tokens_priority": 3e-06, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.5e-06, - "output_cost_per_token_above_272k_tokens": 6.75e-06, "output_cost_per_token_priority": 9e-06, - "output_cost_per_token_above_272k_tokens_priority": 1.35e-05, "source": "https://ai.azure.com/catalog/models/gpt-5.4-mini", "supported_endpoints": [ "/v1/chat/completions", @@ -3500,22 +3494,16 @@ }, "azure_ai/gpt-5.4-mini-2026-03-17": { "cache_read_input_token_cost": 7.5e-08, - "cache_read_input_token_cost_above_272k_tokens": 1.5e-07, "cache_read_input_token_cost_priority": 1.5e-07, - "cache_read_input_token_cost_above_272k_tokens_priority": 3e-07, "input_cost_per_token": 7.5e-07, - "input_cost_per_token_above_272k_tokens": 1.5e-06, "input_cost_per_token_priority": 1.5e-06, - "input_cost_per_token_above_272k_tokens_priority": 3e-06, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.5e-06, - "output_cost_per_token_above_272k_tokens": 6.75e-06, "output_cost_per_token_priority": 9e-06, - "output_cost_per_token_above_272k_tokens_priority": 1.35e-05, "source": "https://ai.azure.com/catalog/models/gpt-5.4-mini", "supported_endpoints": [ "/v1/chat/completions", @@ -3546,22 +3534,16 @@ }, "azure_ai/gpt-5.4-nano": { "cache_read_input_token_cost": 2e-08, - "cache_read_input_token_cost_above_272k_tokens": 4e-08, "cache_read_input_token_cost_priority": 4e-08, - "cache_read_input_token_cost_above_272k_tokens_priority": 8e-08, "input_cost_per_token": 2e-07, - "input_cost_per_token_above_272k_tokens": 4e-07, "input_cost_per_token_priority": 4e-07, - "input_cost_per_token_above_272k_tokens_priority": 8e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.25e-06, - "output_cost_per_token_above_272k_tokens": 1.875e-06, "output_cost_per_token_priority": 2.5e-06, - "output_cost_per_token_above_272k_tokens_priority": 3.75e-06, "source": "https://ai.azure.com/catalog/models/gpt-5.4-nano", "supported_endpoints": [ "/v1/chat/completions", @@ -3592,22 +3574,16 @@ }, "azure_ai/gpt-5.4-nano-2026-03-17": { "cache_read_input_token_cost": 2e-08, - "cache_read_input_token_cost_above_272k_tokens": 4e-08, "cache_read_input_token_cost_priority": 4e-08, - "cache_read_input_token_cost_above_272k_tokens_priority": 8e-08, "input_cost_per_token": 2e-07, - "input_cost_per_token_above_272k_tokens": 4e-07, "input_cost_per_token_priority": 4e-07, - "input_cost_per_token_above_272k_tokens_priority": 8e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.25e-06, - "output_cost_per_token_above_272k_tokens": 1.875e-06, "output_cost_per_token_priority": 2.5e-06, - "output_cost_per_token_above_272k_tokens_priority": 3.75e-06, "source": "https://ai.azure.com/catalog/models/gpt-5.4-nano", "supported_endpoints": [ "/v1/chat/completions", @@ -7201,7 +7177,7 @@ "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 7.5e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7236,7 +7212,7 @@ "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 7.5e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7271,7 +7247,7 @@ "cache_read_input_token_cost": 2e-08, "input_cost_per_token": 2e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7306,7 +7282,7 @@ "cache_read_input_token_cost": 2e-08, "input_cost_per_token": 2e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -13567,6 +13543,56 @@ } ] }, + "dashscope/qwen3.7-max": { + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token": 2.5e-06, + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 7.5e-06, + "source": "https://www.alibabacloud.com/help/en/model-studio/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "dashscope/qwen3.7-plus": { + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "source": "https://www.alibabacloud.com/help/en/model-studio/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tiered_pricing": [ + { + "cache_read_input_token_cost": 8e-08, + "input_cost_per_token": 4e-07, + "output_cost_per_token": 1.6e-06, + "range": [ + 0, + 256000.0 + ] + }, + { + "cache_read_input_token_cost": 2.4e-07, + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 4.8e-06, + "range": [ + 256000.0, + 1000000.0 + ] + } + ] + }, "dashscope/qwq-plus": { "input_cost_per_token": 8e-07, "litellm_provider": "dashscope", diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index a27d6b92843..423cda5eea2 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -1,6 +1,6 @@ import re from datetime import datetime, timezone -from typing import Dict, List, Optional, Set, Tuple, cast +from typing import TYPE_CHECKING, Dict, List, Optional, Sequence, Set, Tuple, cast from fastapi import HTTPException from starlette.datastructures import Headers @@ -30,6 +30,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credent ) from litellm.proxy._types import ( UI_TEAM_ID, + LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, ProxyException, SpecialHeaders, @@ -43,13 +44,27 @@ from litellm.proxy.auth.user_api_key_auth import ( user_api_key_auth, ) from litellm.proxy.common_utils.http_parsing_utils import _read_request_body -from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl +from litellm.proxy.common_utils.user_api_key_cache import ( + USER_NO_MCP_PERMISSION_SENTINEL, + get_management_object_ttl, + user_object_permission_id_cache_key, +) from litellm.repositories.table_repositories import ( AgentsRepository, MCPServerRepository, ) +from litellm.repositories.user_repository import UserRepository from litellm.types.mcp_server.mcp_server_manager import MCPServer +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + + +def _as_list(values: Sequence[str] | None) -> list[str] | None: # mutable-ok: resolver returns a list + """Widen a read-only allowlist back to the mutable list the resolver's own contract returns, + preserving the ``None`` that means "no restriction".""" + return None if values is None else list(values) + def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: Optional[List[str]] = None) -> Optional[List[str]]: """Resolve the single MCP server name a cold-start passthrough bypass may @@ -1408,6 +1423,15 @@ class MCPRequestHandler: f"Applied agent intersection filter. Final allowed servers: {allowed_mcp_servers}" ) + ######################################################### + # Apply the internal user's own ceiling (the entitlement attached to the human) + ######################################################### + capped, user_restricts = await MCPRequestHandler._apply_user_server_ceiling( + allowed_mcp_servers, user_api_key_auth, keyless_source=keyless_source + ) + allowed_mcp_servers = list(capped) + has_lower_level_mcp_restrictions = has_lower_level_mcp_restrictions or user_restricts + ######################################################### # Apply org-level ceiling if org_id is set ######################################################### @@ -1831,6 +1855,12 @@ class MCPRequestHandler: # No team restrictions → use key restrictions allowed_tools = cast(List[str], key_tools) + allowed_tools = _as_list( + await MCPRequestHandler._apply_user_tool_ceiling( + allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source + ) + ) + return await MCPRequestHandler._apply_agent_and_org_tool_ceilings( allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source ) @@ -2376,6 +2406,203 @@ class MCPRequestHandler: verbose_logger.warning(f"Failed to get allowed MCP servers for end_user: {str(e)}") return [] + @staticmethod + async def _get_user_object_permission( + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> LiteLLM_ObjectPermissionTable | None: + """The internal user's OWN object_permission: the entitlement attached to the HUMAN rather + than to the credential they authenticated with. + + A key's object_permission is the credential's scope and a team's is the group's; this one + answers "which MCP servers and tools is this person entitled to", independent of how many keys + they hold. Caches the ``user_id -> object_permission_id`` mapping (with a sentinel for "no + entitlement") exactly as the agent path does, then reuses the shared ``object_permission_id`` + cache, so a warm request reads no rows. + + ``None`` means the human places NO ceiling: no user row, or a row naming no permission. The + two fault classes are deliberately NOT collapsed into that: a user row we cannot read leaves + us unable to say whether they are entitled at all, which is exactly the state before this + level existed, so it places no ceiling; a row that NAMES a permission we cannot read is a + KNOWN entitlement with unknown contents, so it raises and the caller denies. + """ + from litellm.proxy.auth.auth_checks import get_object_permission + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + if not user_api_key_auth or not user_api_key_auth.user_id: + return None + + if prisma_client is None: + verbose_logger.debug("prisma_client is None") + return None + + user_id = user_api_key_auth.user_id + object_permission_id = await MCPRequestHandler._user_object_permission_id(user_id, prisma_client) + if object_permission_id is None: + return None + + object_permission = await get_object_permission( + object_permission_id=object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=user_api_key_auth.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + if object_permission is None: + raise ValueError( + f"user {user_id!r} names object_permission_id {object_permission_id!r} which could not be loaded" + ) + return object_permission + + @staticmethod + async def _user_object_permission_id(user_id: str, prisma_client: "PrismaClient") -> str | None: + """The permission row this human's user row links to, or None when they link none. + + Caches the link (with a sentinel for "links none") so a human without an entitlement costs no + DB read per MCP request. Anything other than an id string is treated as a cache MISS rather + than carried into the permission lookup, and a read that fails answers None: not knowing + whether someone is entitled is the state that existed before this level, so it places no + ceiling. Only a link we DID resolve can make the caller deny. + """ + from litellm.proxy.proxy_server import user_api_key_cache + + cache_key = user_object_permission_id_cache_key(user_id) + try: + cached: object = await user_api_key_cache.async_get_cache(key=cache_key) + if cached == USER_NO_MCP_PERMISSION_SENTINEL: + return None + if isinstance(cached, str) and cached: + return cached + user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) + linked: object = getattr(user_row, "object_permission_id", None) if user_row is not None else None + object_permission_id = linked if isinstance(linked, str) and linked else None + await user_api_key_cache.async_set_cache( + key=cache_key, + value=object_permission_id or USER_NO_MCP_PERMISSION_SENTINEL, + ttl=get_management_object_ttl(user_api_key_cache), + ) + return object_permission_id + except Exception as e: # noqa: BLE001 # unknown whether entitled at all: no ceiling, as before + verbose_logger.warning(f"MCP user entitlement: link for {user_id!r} unresolved, no ceiling: {str(e)}") + return None + + @staticmethod + async def _get_allowed_mcp_servers_for_user( + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> Sequence[str] | None: + """The MCP servers the internal user is entitled to, as server ids. + + ``[]`` means this human places no restriction (allow-all from this level); ``None`` means the + ceiling is UNRESOLVED, which the caller denies on. Servers named only under + ``mcp_tool_permissions`` count as entitled, exactly as they do for a key or a team, so + granting one tool never requires naming its server twice. + """ + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + try: + object_permissions = await MCPRequestHandler._get_user_object_permission(user_api_key_auth) + if object_permissions is None: + return [] + + direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []) + access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups( + object_permissions.mcp_access_groups or [] + ) + tool_perm_servers = list( + global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys() + ) + return list(set(direct_mcp_servers + access_group_servers + tool_perm_servers)) + except Exception as e: # noqa: BLE001 # any resolution fault is an unresolved ceiling, never "no ceiling" + verbose_logger.warning(f"Failed to get allowed MCP servers for user: {str(e)}") + return None + + @staticmethod + async def _apply_user_server_ceiling( + allowed_mcp_servers: Sequence[str], + user_api_key_auth: UserAPIKeyAuth | None = None, + *, + keyless_source: bool = False, + ) -> tuple[tuple[str, ...], bool]: + """Narrow a resolved server list by the internal user's own entitlement. + + Returns the capped list and whether this human restricted it at all; the caller needs the + second value because an org list may only CAP a lower-level restriction, never replace one, so + a user ceiling has to be visible to the org step. + + RAISES when the entitlement is known but unreadable, which the resolver's own handler turns + into deny-all. That is the point of the level: dropping a ceiling we know exists is exactly the + silent widening it is there to prevent. + """ + if keyless_source: + return tuple(allowed_mcp_servers), False + entitled = await MCPRequestHandler._get_allowed_mcp_servers_for_user(user_api_key_auth) + if entitled is None: + raise ValueError( + f"MCP user ceiling unresolvable for user_id=" + f"{user_api_key_auth.user_id if user_api_key_auth else None!r}" + ) + if not entitled: + return tuple(allowed_mcp_servers), False + capped = tuple(server for server in allowed_mcp_servers if server in set(entitled)) + verbose_logger.debug(f"Applied user ceiling filter. Final allowed servers: {capped}") + return capped, True + + @staticmethod + async def _user_places_mcp_ceiling(user_api_key_auth: UserAPIKeyAuth | None = None) -> bool: + """Whether this human's own entitlement bounds their MCP access at all. + + True when they are entitled to a specific set of servers, and also when that entitlement is + UNRESOLVED — a caller uses this to decide whether it may skip the resolver, and skipping it on + a transient fault would widen access. + """ + entitled_servers = await MCPRequestHandler._get_allowed_mcp_servers_for_user(user_api_key_auth) + return entitled_servers is None or len(entitled_servers) > 0 + + @staticmethod + async def _apply_user_tool_ceiling( + allowed_tools: Sequence[str] | None, + server_id: str, + user_api_key_auth: UserAPIKeyAuth | None = None, + *, + keyless_source: bool = False, + ) -> Sequence[str] | None: + """Narrow a key/team tool allowlist by the internal user's own tool entitlement. + + The human's entitlement can only ever narrow: a user naming tools on ``server_id`` intersects + (and becomes the allowlist when no lower level restricts), while a user naming none places no + restriction. Returns ``[]`` (deny every tool on this server) when the entitlement cannot be + resolved, because the caller's own except-handler treats a raise as allow-all for key auth. + """ + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + if keyless_source: + return allowed_tools + + try: + object_permissions = await MCPRequestHandler._get_user_object_permission(user_api_key_auth) + except Exception as e: # noqa: BLE001 # an unresolved human entitlement must deny, not widen + verbose_logger.warning(f"MCP user tool ceiling unresolvable, denying tools on {server_id!r}: {str(e)}") + return [] + + if object_permissions is None or not object_permissions.mcp_tool_permissions: + return allowed_tools + + user_tools = global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).get( + server_id + ) + if user_tools is None: + return allowed_tools + if allowed_tools is None: + return list(user_tools) + return list(set(allowed_tools) & set(user_tools)) + # Sentinel stored in cache when an agent has no object_permission, so we # don't re-query the DB on every MCP request for that agent. _AGENT_NO_PERMISSION_SENTINEL = "__agent_no_mcp_permission__" diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index caa5c65894c..cdc3ac15b1a 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1778,10 +1778,12 @@ async def token_endpoint( @router.post("/authorize/complete") -async def authorize_complete(request: Request, flow: str = Form(...)): +async def authorize_complete(request: Request, flow: str = Form(...), delivery: str | None = Form(None)): """Finish an aggregate connect flow: mint the gateway authorization code for the - signed-in user and redirect back to the DCR client. POST plus the per-flow HttpOnly - cookie set at /authorize; an anonymous or bad-flow request just 400s.""" + signed-in user and hand it back to the DCR client, by 303 redirect (default) or, for + a loopback client on a different machine, as a copyable callback URL + (``delivery=manual``). POST plus the per-flow HttpOnly cookie set at /authorize; an + anonymous or bad-flow request just 400s.""" from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 # circular import at module load return await complete_connect_flow( @@ -1789,6 +1791,7 @@ async def authorize_complete(request: Request, flow: str = Form(...)): flow_handle=flow, session_user_id=_session_cookie_user_id(request), cache=user_api_key_cache, + delivery=delivery, ) diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index 58233c4c9e5..7177b798c5f 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -39,6 +39,7 @@ from __future__ import annotations import hashlib import hmac +import html import secrets from base64 import urlsafe_b64encode from collections.abc import Mapping @@ -47,7 +48,7 @@ from typing import Awaitable, Callable, Literal, TypeVar from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse from fastapi import HTTPException, Request -from fastapi.responses import JSONResponse, RedirectResponse, Response +from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response from pydantic import BaseModel, ConfigDict, Field, ValidationError from typing_extensions import assert_never @@ -94,6 +95,13 @@ server-side session store, and the sealed value never appears in a URL).""" CONNECT_FLOW_TTL_SECONDS = 600 GATEWAY_AUTH_CODE_TTL_SECONDS = 120 +MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS = 300 +"""Lifetime of a code the user delivers by hand (headless/remote client, LIT-4863 class): +copy-pasting a callback URL from a laptop browser to an SSH session is slower than a +browser redirect, so manual-delivery codes get 5 minutes instead of 2, still well under +the 10-minute ceiling RFC 6749 section 4.1.2 recommends. Single-use and PKCE binding are +unchanged, so the longer window only extends how long the legitimate holder has to paste +it, not what an observer could do with it.""" _CLAIM_TTL_BUFFER_SECONDS = 60 _USED_CODE_CACHE_PREFIX = "mcp_gateway_dcr_code_used:" _USED_FLOW_CACHE_PREFIX = "mcp_gateway_dcr_flow_used:" @@ -390,6 +398,7 @@ async def complete_connect_flow( flow_handle: str, session_user_id: str | None, cache: DualCache, + delivery: str | None = None, ) -> Response: """The deliberate finish step of the connect flow: mint the gateway authorization code and send the browser back to the client. @@ -399,7 +408,24 @@ async def complete_connect_flow( into the flow: a link crafted by another party dies here with ``access_denied`` instead of minting a code for the victim's identity. The flow is single-use (an atomic claim on its ``jti``), so a double-submit cannot mint two codes from one sign-in. + + ``delivery`` chooses how the code reaches the client. Default (absent or + ``"redirect"``) is the 303 to the client's registered redirect URI. ``"manual"`` + renders the callback URL on a page instead, for a client whose redirect URI is a + loopback host but which runs on a DIFFERENT machine than the browser (EC2/SSH box, + container): the 303 would dereference the browser machine's loopback and the code + would never arrive, so the user carries it over by pasting the URL into the client or + fetching it from the client machine's terminal. Manual delivery is honored only for + loopback redirect URIs; a routable redirect URI works from any browser by + construction, so those flows always redirect. The user who sees the page is exactly + the user the 303 would have carried the code to, and the same user already sees the + code today in the dead redirect's address bar, so the page exposes the code to no new + party. Unknown ``delivery`` values are rejected rather than defaulted: a client that + asked for manual delivery and got a dead redirect instead would silently lose its + code. """ + if delivery not in (None, "redirect", "manual"): + return _oauth_error(400, "invalid_request", "delivery must be 'redirect' or 'manual'") sealed_flow = request.cookies.get(_flow_cookie_name(flow_handle)) if sealed_flow is None: return _oauth_error(400, "invalid_request", "unknown or expired connect flow") @@ -417,6 +443,8 @@ async def complete_connect_flow( f"{_USED_FLOW_CACHE_PREFIX}{flow.jti}", CONNECT_FLOW_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS ): return _oauth_error(400, "invalid_request", "this connect flow was already completed; restart the connection") + manual_delivery = delivery == "manual" and is_loopback_redirect_host(urlparse(flow.redirect_uri)) + code_ttl = MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS if manual_delivery else GATEWAY_AUTH_CODE_TTL_SECONDS code = _seal( GATEWAY_AUTH_CODE_PREFIX, _GatewayAuthCode( @@ -426,16 +454,46 @@ async def complete_connect_flow( code_challenge=flow.code_challenge, jti=secrets.token_urlsafe(24), iat=int(now.timestamp()), - exp=int(now.timestamp()) + GATEWAY_AUTH_CODE_TTL_SECONDS, + exp=int(now.timestamp()) + code_ttl, ), ) params = {"code": code, **({"state": flow.state} if flow.state else {})} - response = RedirectResponse(_append_query_params(flow.redirect_uri, params), status_code=303) + callback_url = _append_query_params(flow.redirect_uri, params) + response: Response = ( + _manual_delivery_response(callback_url) if manual_delivery else RedirectResponse(callback_url, status_code=303) + ) path, secure = _cookie_path_and_secure(request) response.delete_cookie(key=_flow_cookie_name(flow_handle), path=path, secure=secure, httponly=True, samesite="lax") return response +def _manual_delivery_response(callback_url: str) -> Response: + """The manual code-delivery page: the callback URL the 303 would have followed, + rendered for the user to carry to the machine the client actually runs on (paste into + the client's prompt, or fetch with curl from that machine's terminal). Served + no-store because the body holds a live single-use code, and the URL is HTML-escaped + because it is client-influenced. The page renders the URL as data only, never as a + ready-to-paste shell command: no single quoting of an attacker-influenced string is + correct across POSIX shells, cmd.exe, and PowerShell (cmd.exe ignores single quotes + and percent-expands inside double quotes), so any command string this page suggested + would be wrong for some shell the user might paste it into.""" + safe_url = html.escape(callback_url, quote=True) + minutes = MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS // 60 + body = ( + "Finish connecting" + "

Almost done

" + "

Your MCP client runs on a different machine, so this browser cannot deliver the" + " authorization code to it. On the machine where the client runs, paste this URL into" + " the client's prompt (Claude Code accepts the pasted callback URL), or pass it as the" + " quoted argument of a curl command from that machine's terminal:

" + f'

' + f"

The code is single-use and expires in {minutes} minutes. You can close this window" + " once the client confirms it is connected.

" + "" + ) + return HTMLResponse(body, headers=TOKEN_NO_CACHE_HEADERS) + + def _pkce_verifier_matches(code_verifier: str, code_challenge: str) -> bool: """RFC 7636 S256 verification, total over hostile input. The comparison is over bytes so a non-ASCII ``code_challenge`` (which reaches here unvalidated from the client's @@ -601,9 +659,11 @@ async def _authorization_code_grant( if failure is not None: return _reload_failure_response(failure) # Atomic single-use claim is the gate: on a concurrent double-redeem exactly one caller - # wins, and a claim that cannot be recorded fails closed. + # wins, and a claim that cannot be recorded fails closed. The marker's TTL derives from + # the code's own remaining lifetime so it outlives whichever lifetime the code was minted with. if not await guard.claim( - f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}", GATEWAY_AUTH_CODE_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS + f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}", + parsed.exp - int(now.timestamp()) + _CLAIM_TTL_BUFFER_SECONDS, ): return _oauth_error(400, "invalid_grant", "the authorization code was already used") return _session_token_pair(SessionPrincipal(user_id=parsed.user_id, client_id=client_id), keys, now) diff --git a/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py b/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py index b668833e638..909925da00a 100644 --- a/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py +++ b/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py @@ -13,11 +13,22 @@ payload (name + arguments) so we just build the tool_call. from typing import TYPE_CHECKING, Any, Dict, Optional +from fastapi import HTTPException from mcp.types import Tool as MCPTool from litellm._logging import verbose_proxy_logger from litellm.experimental_mcp_client.tools import transform_mcp_tool_to_openai_tool from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation +from litellm.proxy._experimental.mcp_server.utils import ( + json_string_leaves, + json_unrewritable_labels, + mcp_content_item_text, + mcp_tool_result_content_list, + mcp_tool_result_structured_content, + set_mcp_tool_result_structured_content, + with_json_string_leaves, + with_mcp_content_item_text, +) from litellm.types.llms.openai import ( ChatCompletionToolParam, ChatCompletionToolParamFunctionChunk, @@ -92,7 +103,93 @@ class MCPGuardrailTranslationHandler(BaseTranslation): user_api_key_dict: Optional[Any] = None, request_data: Optional[dict] = None, ) -> Any: - verbose_proxy_logger.debug( - "MCP Guardrail: Output processing not implemented for MCP tools", + """Scan the text content of an MCP tool result and write masked text back. + + The content list is rewritten in place (only the entries the guardrail + actually changed) rather than returned as a new result: the same object is + already referenced by the logging payload captured before this hook runs, + so a copy would leave the unmasked text in the spend log / span. A + guardrail that rejects the result raises, and the exception propagates to + the caller. + + ``structuredContent`` is scanned and masked too, in the same + ``apply_guardrail`` call: it is serialized to the client alongside + ``content``, so a value living only there would otherwise reach the + client unscanned. + """ + content = mcp_tool_result_content_list(response) + text_blocks = ( + tuple( + (index, text) for index, item in enumerate(content) if (text := mcp_content_item_text(item)) is not None + ) + if content is not None + else () ) + + structured = mcp_tool_result_structured_content(response) + structured_leaves = json_string_leaves(structured) if structured is not None else () + structured_labels = json_unrewritable_labels(structured) if structured is not None else () + if structured_leaves is None or structured_labels is None: + raise HTTPException( + status_code=400, + detail={ + "error": ( + "Content blocked: MCP tool result structuredContent is nested too deeply to be scanned " + "by the configured guardrail" + ) + }, + ) + + if not text_blocks and not structured_leaves and not structured_labels: + verbose_proxy_logger.debug("MCP Guardrail: tool result has no scannable text, nothing to do") + return response + + originals = ( + tuple(text for _, text in text_blocks) + tuple(text for _, text in structured_leaves) + structured_labels + ) + guardrailed_inputs = await guardrail_to_apply.apply_guardrail( + inputs=GenericGuardrailAPIInputs(texts=list(originals)), + request_data=request_data if request_data is not None else {}, + input_type="response", + logging_obj=litellm_logging_obj, + ) + masked_texts = guardrailed_inputs.get("texts") if guardrailed_inputs else None + if masked_texts is None: + return response + if len(masked_texts) != len(originals): + verbose_proxy_logger.warning( + "MCP Guardrail: guardrail returned %d texts for %d tool result texts; leaving the result unmasked", + len(masked_texts), + len(originals), + ) + return response + + split = len(text_blocks) + if content is not None: + for (index, original), masked in zip(text_blocks, masked_texts[:split]): + if masked != original: + content[index] = with_mcp_content_item_text(content[index], masked) + + label_start = split + len(structured_leaves) + if any(masked != original for original, masked in zip(structured_labels, masked_texts[label_start:])): + raise HTTPException( + status_code=400, + detail={ + "error": ( + "Content blocked: MCP tool result matched a masking rule on a non-rewritable field " + "(a structuredContent key or numeric value), which cannot be redacted without changing " + "the payload contract" + ) + }, + ) + + structured_replacements = { + path: masked + for (path, original), masked in zip(structured_leaves, masked_texts[split:label_start]) + if masked != original + } + if structured_replacements: + set_mcp_tool_result_structured_content( + response, with_json_string_leaves(structured, structured_replacements) + ) return response diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 82b820d8cd9..ae1095da336 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -200,6 +200,13 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = ( ) +# OAuth discovery retry cooldown for servers whose endpoints stay unresolved. The base is one +# reload cadence so a transient upstream failure recovers immediately; the cap bounds the request +# amplification and log volume of a permanently broken configuration. +_OAUTH_DISCOVERY_RETRY_BASE_SECONDS = 30.0 +_OAUTH_DISCOVERY_RETRY_MAX_SECONDS = 900.0 + + def _blank_to_none(value: str | None) -> str | None: """Collapse an absent, empty, or whitespace-only string to ``None``. @@ -247,6 +254,7 @@ def _endpoints_yield_to_issuer( authorization_url: str | None, token_url: str | None, registration_url: str | None, + server_ref: str, ) -> tuple[str | None, str | None, str | None]: """The single rule that makes an admin-configured ``issuer`` the sole authoritative endpoint source (RFC 8414 §3.3): when it is set for a discovery auth type, the stored/manual @@ -256,9 +264,29 @@ def _endpoints_yield_to_issuer( i.e. all ``None`` when issuer-anchored, else the inputs unchanged. Called at every resolution site so the invariant holds in one place instead of being re-derived per merge. """ - if issuer is not None and is_discovery_auth_type: - return None, None, None - return authorization_url, token_url, registration_url + if issuer is None or not is_discovery_auth_type: + return authorization_url, token_url, registration_url + discarded = sorted( + label + for label, value in ( + ("authorization_url", authorization_url), + ("token_url", token_url), + ("registration_url", registration_url), + ) + if value + ) + if discarded: + verbose_logger.warning( + "MCP server %s has a pinned Issuer, so its stored %s %s not used: an anchored issuer is the " + "sole endpoint source (RFC 8414 section 3.3) and a failed issuer fetch fails closed rather " + "than falling back to them. To use manually configured endpoints instead, clear the Issuer " + "field and re-enter the endpoint urls (clearing the Issuer also clears endpoints that may " + "have been resolved under it), or clear the Issuer alone to re-discover from the server url.", + server_ref, + ", ".join(discarded), + "is" if len(discarded) == 1 else "are", + ) + return None, None, None def _normalized_authorize_endpoint(url: str) -> str: @@ -280,6 +308,68 @@ def _issuer_matches(claimed_issuer: object, configured_issuer: str) -> bool: return _normalized_authorize_endpoint(claimed_issuer) == _normalized_authorize_endpoint(configured_issuer) +def _flow_endpoints_missing( + auth_type: MCPAuthType | None, + oauth2_flow: str | None, + authorization_url: str | None, + token_url: str | None, + token_exchange_endpoint: str | None = None, +) -> bool: + """Whether a built server is missing an endpoint its flow needs to run at all. + + Used by the reload fast-path exemption: discovery runs at build time only, and the fast path + reuses an unchanged row's registry entry verbatim, so a server whose discovery came back empty + (transient upstream failure, rate limiting) would stay broken until some unrelated config write + bumps ``updated_at``, serving its 400 the whole time. Rebuilding just these entries retries + discovery on the normal reload cadence. It costs no extra fetch for servers that resolved, and + none for those with no discovery source, since the build skips discovery for both. + """ + if auth_type == MCPAuth.oauth2_token_exchange: + # A configured exchange endpoint replaces discovery entirely; only a server that must + # discover its token endpoint and still has none is unresolved. + return token_exchange_endpoint is None and token_url is None + if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: + return False + if oauth2_flow == "client_credentials": + return token_url is None + return authorization_url is None or token_url is None + + +def _oauth_endpoints_unresolved(server: MCPServer) -> bool: + """``_flow_endpoints_missing`` over a built registry entry, for the reload fast-path check. + + The flow comes from ``effective_oauth2_flow``, the one column-first, shape-fallback judge every + flow decision uses, not from the raw column: a legacy row the startup backfill deliberately left + unstamped (the ambiguous M2M shape) serves M2M at request time, and reading the bare column here + would classify it as interactive-missing-endpoints and re-run discovery on every reload. + """ + if ( + server.auth_type == MCPAuth.oauth2_token_exchange + and server.token_exchange_profile == "entra_obo" + and not server.scopes + ): + # entra_obo fails closed at exchange time without a scope (token_exchanger.py), and scopes + # can come from resource discovery, so a server that resolved its endpoints but no scopes is + # still unresolved for its flow. + return True + if server.is_dcr_bridge and not server.client_id and server.registration_url is None: + # A DCR bridge with no admin-configured client can only register callers through the + # upstream's registration endpoint, so a build that resolved the authorize and token + # endpoints but not registration_endpoint (partial metadata) is still unresolved for its + # flow and must keep retrying; without this it silently degrades to the short-circuit arm + # until an unrelated config write. Scopes are deliberately NOT part of completeness: they + # are a request hint the authorization server bounds at consent (RFC 6749 section 3.3), + # and a server without them is fully functional. + return True + return _flow_endpoints_missing( + server.auth_type, + MCPServerManager.effective_oauth2_flow(server), + server.authorization_url, + server.token_url, + server.token_exchange_endpoint, + ) + + def _endpoints_corroborate_authorization_url( source_authorization_url: str | None, trusted_authorization_url: str | None, @@ -311,11 +401,10 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv during re-discovery downgrades a working server (``authorization_url`` set) to a broken one (``None``, /authorize 400s) with no configuration change. Mirrors the ``short_prefix`` carry-forward. Skipped when the server's ``url`` or ``auth_type`` changed, since the previous - endpoints may then belong to a different upstream. ``registration_url`` IS carried even though - ``_persist_discovered_oauth_endpoints`` refuses to write it to the row: carrying only restores - the same in-memory value the previous build already ran with, while persisting it would flip - ``_dcr_bridge_relays_client_registration`` (which keys off the stored column) for dcr_bridge - servers that never had one configured. + endpoints may then belong to a different upstream. Discovery results live only on the in-memory + registry entry; the gateway never writes them to the row, whose OAuth columns carry admin intent + alone, so this carry is the sole last-known-good mechanism and restores exactly the values the + previous build already ran with. Carry-forward is a non-manual endpoint source, so the same trust rule as discovery applies: the previous ``token_url``/``registration_url``/``scopes`` are carried only when the previous @@ -1182,6 +1271,40 @@ class MCPServerManager: # empty result, or failure). Used to throttle re-probes for servers that do # not return instructions, and to apply a short cooldown after failures. self._upstream_initialize_instructions_probed_at: dict[str, float] = {} + # Per-server (consecutive failures, monotonic timestamp) for OAuth discovery retries, so a + # server whose endpoints never resolve backs off instead of re-running the full + # RFC 9728 -> 8414 chain, and re-logging its warning, on every reload forever. + self._oauth_discovery_retry_state: dict[ + str, tuple[int, float] + ] = {} # mutable-ok: retry cooldown cache, keyed per server and pruned on success + + def _oauth_discovery_retry_due(self, server_id: str) -> bool: + """Whether an unresolved server is due for another discovery attempt. + + The reload fast-path exemption is what retries a failed discovery, so without a cooldown a + permanently unresolvable server re-runs the whole RFC 9728 -> RFC 8414 -> origin-fallback + chain and re-emits its unresolved-endpoints warning on every reload, per server, forever. + Delay doubles per consecutive failure from ``_OAUTH_DISCOVERY_RETRY_BASE_SECONDS`` up to + ``_OAUTH_DISCOVERY_RETRY_MAX_SECONDS``, so a transient outage still recovers on the next + reload while a broken configuration settles to one attempt per cap. + """ + state = self._oauth_discovery_retry_state.get(server_id) + if state is None: + return True + failures, attempted_at = state + delay = min( + _OAUTH_DISCOVERY_RETRY_BASE_SECONDS * (2 ** max(failures - 1, 0)), + _OAUTH_DISCOVERY_RETRY_MAX_SECONDS, + ) + return (time.monotonic() - attempted_at) >= delay + + def _record_oauth_discovery_outcome(self, server: MCPServer) -> None: + """Advance or clear a server's retry cooldown after a rebuild resolved it or did not.""" + if not _oauth_endpoints_unresolved(server): + self._oauth_discovery_retry_state.pop(server.server_id, None) + return + failures, _ = self._oauth_discovery_retry_state.get(server.server_id, (0, 0.0)) + self._oauth_discovery_retry_state[server.server_id] = (failures + 1, time.monotonic()) def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None: raw = getattr(client, "_last_initialize_instructions", None) @@ -1357,6 +1480,7 @@ class MCPServerManager: manual_authorization_url, manual_token_url, manual_registration_url, + server_name or server_id, ) should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and ( is_discovery_auth_type or obo_needs_discovery @@ -1834,7 +1958,6 @@ class MCPServerManager: *, credentials_are_encrypted: bool = True, env_vars_are_encrypted: Optional[bool] = None, - persist_discovered_endpoints: bool = True, ) -> MCPServer: _mcp_info: MCPInfo = mcp_server.mcp_info or {} env_dict = _deserialize_json_dict(getattr(mcp_server, "env", None)) @@ -1925,7 +2048,12 @@ class MCPServerManager: or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url), ) manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer( - manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url + manual_issuer, + is_discovery_auth_type, + manual_authorization_url, + manual_token_url, + manual_registration_url, + mcp_server.alias or mcp_server.server_name or mcp_server.server_id, ) gated_oauth_metadata = await self._resolve_table_oauth_metadata( mcp_server=mcp_server, @@ -2033,143 +2161,8 @@ class MCPServerManager: max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None), ) _warn_internal_delegate_pkce_if_applicable(new_server, source="database") - if persist_discovered_endpoints: - await self._persist_discovered_obo_token_url( - server_id=mcp_server.server_id, - auth_type=auth_type, - existing_token_url=manual_token_url, - discovered_token_url=new_server.token_url, - ) - await self._persist_discovered_oauth_endpoints( - server_id=mcp_server.server_id, - auth_type=auth_type, - existing_issuer=manual_issuer, - existing_authorization_url=manual_authorization_url, - existing_token_url=manual_token_url, - existing_scopes=scopes, - metadata=gated_oauth_metadata, - is_issuer_anchored=use_issuer_anchor, - ) return new_server - async def _persist_discovered_obo_token_url( - self, - *, - server_id: str, - auth_type: Optional[MCPAuthType], - existing_token_url: Optional[str], - discovered_token_url: Optional[str], - ) -> None: - """Write a freshly discovered OBO token endpoint back onto the DB row. - - ``build_mcp_server_from_table`` resolves ``token_url`` via RFC 9728 -> RFC 8414 for an - ``oauth2_token_exchange`` server that has none configured, but that resolved value otherwise - lives only on the returned in-memory object; the row keeps ``token_url=None`` so every rebuild - re-runs discovery, and a transient upstream outage during a rebuild leaves the server with no - endpoint until discovery next succeeds. Persisting it makes ``_obo_needs_endpoint_discovery`` - return False on the next build. Fires at most once per server (skipped once the row has a - value), and is best-effort: a write failure just means discovery runs again next time. - """ - if auth_type != MCPAuth.oauth2_token_exchange: - return - if existing_token_url or not discovered_token_url: - return - from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 - - if prisma_client is None: - return - try: - await MCPServerRepository(prisma_client).table.update( - where={"server_id": server_id}, - data={"token_url": discovered_token_url}, - ) - verbose_logger.debug("Persisted discovered OBO token_url for MCP server %s", server_id) - except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build - verbose_logger.warning("Failed to persist discovered OBO token_url for MCP server %s: %s", server_id, exc) - - async def _persist_discovered_oauth_endpoints( - self, - *, - server_id: str, - auth_type: MCPAuthType | None, - existing_issuer: str | None, - existing_authorization_url: str | None, - existing_token_url: str | None, - existing_scopes: list[str] | None, - metadata: MCPOAuthMetadata | None, - is_issuer_anchored: bool = False, - ) -> None: - """Write freshly discovered OAuth endpoints back onto the DB row. - - Same rationale as ``_persist_discovered_obo_token_url`` but for the interactive oauth2 - family: discovered ``authorization_url``/``token_url``/``scopes`` otherwise live only on - the in-memory registry entry, which is rebuilt on every client connect (the DCR reuse path - calls ``update_server``) and on every post-write DB reload, so one failed re-discovery - serves the 400 "authorization url is not configured" from /authorize until a later rebuild succeeds. - Only fills row fields that are currently empty, never persists origin-fallback guesses - (RFC 9728/8414-advertised metadata only), and deliberately skips ``registration_url`` - because ``_dcr_bridge_relays_client_registration`` keys off that column. Best-effort: a - failed write re-discovers on the next build. Scopes go through ``update_mcp_server`` so - they merge into the credentials blob without touching the stored client credentials. - - For an issuer-anchored server (``is_issuer_anchored``) the endpoints are re-derived from the - §3.3-validated issuer document on every build, so they are NOT persisted into the endpoint - columns: persisting them would make the next build see populated endpoints and treat them as - authoritative stored values, defeating the "endpoints come solely from the issuer" invariant. - Only the resource-driven scopes are persisted for such servers. - """ - if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: - return - if metadata is None or metadata.from_origin_fallback: - return - issuer_update = ( - {"issuer": metadata.discovered_issuer} if metadata.discovered_issuer and not existing_issuer else {} - ) - authorization_url_update = ( - {"authorization_url": metadata.authorization_url} - if metadata.authorization_url and not existing_authorization_url and not is_issuer_anchored - else {} - ) - token_url_update = ( - {"token_url": metadata.token_url} - if metadata.token_url and not existing_token_url and not is_issuer_anchored - else {} - ) - scopes_update = {"credentials": {"scopes": metadata.scopes}} if metadata.scopes and not existing_scopes else {} - updates: dict[str, object] = { - **issuer_update, - **authorization_url_update, - **token_url_update, - **scopes_update, - } - if not updates: - return - from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # db.py imports this module at load - update_mcp_server, - ) - from litellm.proxy._types import UpdateMCPServerRequest # noqa: PLC0415 # heavy module; import at call time - from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime value, set after startup - - if prisma_client is None: - return - try: - await update_mcp_server( - prisma_client=prisma_client, - data=UpdateMCPServerRequest.model_validate({"server_id": server_id, **updates}), - touched_by="mcp_oauth_discovery", - ) - verbose_logger.info( - "Persisted discovered OAuth endpoints for MCP server %s: %s", - server_id, - sorted(updates), - ) - except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build - verbose_logger.warning( - "Failed to persist discovered OAuth endpoints for MCP server %s: %s", - server_id, - exc, - ) - async def _maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True): """Register OpenAPI tools if the server has a spec_path configured.""" if server.spec_path: @@ -2418,6 +2411,11 @@ class MCPServerManager: and not is_admitted_subject and _user_has_admin_view(user_api_key_auth) and not has_explicit_object_permission + # An entitlement attached to the HUMAN binds them whatever their role: it is the + # person's scope, not the credential's, so an admin role is not a waiver of it. An + # UNRESOLVED entitlement also skips the shortcut, so the resolver denies rather than + # handing over the whole registry on a transient fault. + and not await MCPRequestHandler._user_places_mcp_ceiling(user_api_key_auth) ): verbose_logger.debug("Admin user without explicit object_permission - returning all servers") return list(self.get_registry().keys()) @@ -5329,7 +5327,7 @@ class MCPServerManager: ] } ) - db_mcp_servers = [LiteLLM_MCPServerTable(**r.model_dump()) for r in raw_rows] + db_mcp_servers = [LiteLLM_MCPServerTable.model_validate(r.model_dump()) for r in raw_rows] verbose_logger.info(f"Found {len(db_mcp_servers)} MCP servers in database") previous_registry = self.registry @@ -5347,6 +5345,10 @@ class MCPServerManager: and existing_server.updated_at is not None and server.updated_at is not None and existing_server.updated_at == server.updated_at + and not ( + _oauth_endpoints_unresolved(existing_server) + and self._oauth_discovery_retry_due(server.server_id) + ) ): # Re-use existing server instance to avoid re-running build_mcp_server_from_table() # which can perform network discovery for OAuth2 servers. @@ -5364,6 +5366,7 @@ class MCPServerManager: # already-decrypted records add_server/update_server are handed. # Decrypt them while building the registry entry. new_server = await self.build_mcp_server_from_table(server, env_vars_are_encrypted=True) + self._record_oauth_discovery_outcome(new_server) # Carry the cached short_prefix from the previous registry entry # (if any) so the prefix is stable across reloads. if existing_server is not None and existing_server.short_prefix: diff --git a/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py b/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py new file mode 100644 index 00000000000..874fcc64772 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py @@ -0,0 +1,148 @@ +"""One-time heal for MCP server rows whose ``issuer`` a released version wrote by itself. + +Until the write was removed, OAuth discovery stamped the issuer it discovered onto the ``issuer`` +column trust-on-first-use. That column means "the admin pinned this trust anchor", so the next +registry build read the gateway's own output back as admin intent: the server turned issuer-anchored +(RFC 8414 section 3.3), its stored authorization/token/registration URLs stopped applying, and a +failed issuer-document fetch left it with no authorize endpoint (GH #34985). + +Deleting the write fixes every row created afterwards but cannot fix a row already stamped, which +still reads as pinned. This heals those rows by clearing the stamp so their configured endpoints +apply again. + +The signal is a heuristic, and deliberately a narrow one. ``updated_by`` records only the most recent +writer, and no audit trail says which field that writer touched, so "discovery wrote this issuer" is +not directly knowable. Two independent clauses bound it, and each rules out a different way of +destroying a pin an admin meant. + +Configured endpoints must be present. A deliberately pinned row very often has none, both because the +Issuer field is documented as overriding them and because ``update_mcp_server`` clears them when an +issuer changes, so "issuer set, endpoints empty" is the canonical shape of a real pin and must never +be cleared on this evidence. Skipping those rows costs little: with nothing configured to restore, the +anchored and resource-rooted paths resolve from the same upstream document, and the row still gets the +unresolved-endpoint retry and the anchored-discard warning. + +The configured endpoints must also share the issuer's origin. A stamped issuer is by construction the +one self-attested by the authorization-server document discovery reached from this very server, so +endpoints typed alongside it address that same authority. An admin who pinned an issuer and typed +endpoints for a different authority is expressing an intent that clearing the issuer would discard, so +that row is warned about and never healed. + +What survives both clauses is a row whose configured endpoints and stamped issuer share an origin, +which is exactly the GH #34985 shape. An admin who pinned that same origin by hand lands here too, and +for them the clear is close to a no-op: their typed endpoints keep serving and still anchor the +RFC 9700 corroboration gate, with only the stricter section 3.3 anchoring lost. Every heal logs the +cleared value so it can be restored, and the clear is recorded under this module's actor so the heal +runs at most once per row. +""" + +from typing import Protocol +from urllib.parse import urlparse + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._experimental.mcp_server.oauth_utils import canonicalize_url_identity +from litellm.proxy.utils import PrismaClient + +# The actor the removed discovery write-back stamped rows with. +_DISCOVERY_ACTOR = "mcp_oauth_discovery" + +# The actor recorded on a healed row, which also makes the heal idempotent: once a row is cleared it +# no longer matches ``updated_by == _DISCOVERY_ACTOR`` and is never reconsidered. +_BACKFILL_ACTOR = "mcp_oauth_issuer_stamp_backfill" + +_AUTH_TYPES_WITH_ISSUER_ANCHORING = ("oauth2", "true_passthrough", "oauth_delegate") + + +def _origin(url: str) -> str | None: + """The scheme-and-authority identity of ``url``, or ``None`` when it has none. + + Built on the shared URL canonicalizer so the lowercase-host and default-port rules match the + RFC 8414 issuer comparison the resolution path uses, instead of being re-derived here. + """ + parsed = urlparse(canonicalize_url_identity(url)) + if not parsed.scheme or not parsed.netloc: + return None + return f"{parsed.scheme}://{parsed.netloc}" + + +class _MCPServerRow(Protocol): + """The MCP server row fields this heal reads, so the untyped DB record is narrowed once here.""" + + server_id: str + alias: str | None + server_name: str | None + auth_type: str | None + issuer: str | None + authorization_url: str | None + token_url: str | None + registration_url: str | None + updated_by: str | None + + +def _is_stamped_issuer_row(row: _MCPServerRow) -> bool: + """Whether this row carries the full signature of a gateway-written issuer stamp. + + The whole rule lives here, including the writer check the query also filters on, so the decision + to clear an admin-visible field is auditable in one place rather than split between a predicate + and a query. + """ + if getattr(row, "updated_by", None) != _DISCOVERY_ACTOR: + return False + if not (getattr(row, "issuer", None) or "").strip(): + return False + if getattr(row, "auth_type", None) not in _AUTH_TYPES_WITH_ISSUER_ANCHORING: + return False + configured = tuple( + value.strip() + for value in (row.authorization_url, row.token_url, row.registration_url) + if value and value.strip() + ) + if not configured: + return False + issuer_origin = _origin(row.issuer or "") + return issuer_origin is not None and all(_origin(endpoint) == issuer_origin for endpoint in configured) + + +async def backfill_discovery_stamped_issuers(prisma_client: PrismaClient) -> int: + """Clear gateway-written issuer stamps, returning the number of rows healed.""" + candidate_rows: list[_MCPServerRow] = await prisma_client.db.litellm_mcpservertable.find_many( + where={ + "updated_by": _DISCOVERY_ACTOR, + "auth_type": {"in": list(_AUTH_TYPES_WITH_ISSUER_ANCHORING)}, + }, + ) + stamped = tuple(row for row in candidate_rows if _is_stamped_issuer_row(row)) + if not stamped: + return 0 + + healed = 0 + for row in stamped: + try: + await prisma_client.db.litellm_mcpservertable.update( + where={"server_id": row.server_id}, + data={"issuer": None, "updated_by": _BACKFILL_ACTOR}, + ) + except Exception as exc: # noqa: BLE001 - per-row best effort; the next boot retries + verbose_proxy_logger.warning( + "MCP issuer stamp backfill: could not heal server_id=%s: %s", row.server_id, exc + ) + continue + healed += 1 + verbose_proxy_logger.warning( + "MCP issuer stamp backfill: cleared issuer %r on server_id=%s (alias=%s). OAuth discovery " + "had written that value onto the Issuer column, which made the server issuer-anchored and " + "fail-closed, and its configured Authorization/Token/Registration URLs were being ignored " + "as a result; those now apply again. If you pinned this issuer deliberately, set it again " + "via the dashboard or PUT /v1/mcp/server to restore RFC 8414 section 3.3 anchoring.", + row.issuer, + row.server_id, + row.alias or row.server_name, + ) + + if healed: + verbose_proxy_logger.warning( + "MCP issuer stamp backfill: healed %d server(s) whose Issuer had been written by OAuth " + "discovery rather than by an admin", + healed, + ) + return healed diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index af3d966c95b..9b51513f4ac 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -12,6 +12,11 @@ import httpx from fastapi import APIRouter, Depends, HTTPException, Query, Request, status from litellm._logging import verbose_logger +from litellm.exceptions import ( + BlockedPiiEntityError, + GuardrailRaisedException, + ModifyResponseException, +) from litellm.proxy._experimental.mcp_server.exceptions import ( MCPServerListError, MCPUpstreamAuthError, @@ -33,6 +38,8 @@ from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.auth.user_api_key_auth import user_api_key_auth if TYPE_CHECKING: + from mcp.types import CallToolResult + from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers from litellm.types.mcp import MCPAuth @@ -51,6 +58,13 @@ router = APIRouter( tags=["mcp"], ) +_MCP_GUARDRAIL_REJECTIONS = ( + BlockedPiiEntityError, + GuardrailRaisedException, + ModifyResponseException, + HTTPException, +) + def _connection_error_message(exc: BaseException) -> str: if isinstance(exc, httpx.LocalProtocolError): @@ -99,9 +113,17 @@ if MCP_AVAILABLE: end_time: datetime, user_api_key_auth: UserAPIKeyAuth | None = None, request_data: Mapping[str, object] | None = None, - ) -> None: + ) -> "CallToolResult": + """Fire post-call logging, returning the tool result to send to the client. + + ``post_mcp_call`` guardrails already ran on ``execute_mcp_tool``'s return + path, so the result arriving here is the guardrailed one. A guardrail + rejection raised by a native ``async_post_mcp_tool_call_hook`` is still + re-raised rather than swallowed as a logging failure, which would return + the unguarded result. + """ if logging_obj is None: - return + return result logging_results = await asyncio.gather( _fire_mcp_tool_call_logging( logging_obj, @@ -113,11 +135,13 @@ if MCP_AVAILABLE: ), return_exceptions=True, ) - logging_error = logging_results[0] - if isinstance(logging_error, asyncio.CancelledError): - raise logging_error - if isinstance(logging_error, BaseException): - verbose_logger.warning("MCP tool call logging failed (continuing): %s", logging_error) + outcome = logging_results[0] + if isinstance(outcome, (asyncio.CancelledError, *_MCP_GUARDRAIL_REJECTIONS)): + raise outcome + if isinstance(outcome, BaseException): + verbose_logger.warning("MCP tool call logging failed (continuing): %s", outcome) + return result + return outcome def _relay_upstream_auth_http_exception(e: MCPUpstreamAuthError, request: Request) -> HTTPException: """Convert a client-forwarded pass-through upstream 401 into an HTTPException that preserves the @@ -196,7 +220,7 @@ if MCP_AVAILABLE: raw_headers=virtual_raw_headers, litellm_logging_obj=virtual_logging_obj, ) - await _safe_fire_mcp_tool_call_logging( + return await _safe_fire_mcp_tool_call_logging( virtual_logging_obj, result, _tool_start_time, @@ -204,7 +228,6 @@ if MCP_AVAILABLE: user_api_key_auth=user_api_key_dict, request_data=data, ) - return result def _get_server_auth_header( server, @@ -998,7 +1021,7 @@ if MCP_AVAILABLE: litellm_logging_obj=data.get("litellm_logging_obj"), requested_server_id=canonical_server_id, ) - await _safe_fire_mcp_tool_call_logging( + return await _safe_fire_mcp_tool_call_logging( logging_obj, result, _tool_start_time, @@ -1006,7 +1029,6 @@ if MCP_AVAILABLE: user_api_key_auth=user_api_key_dict, request_data=data, ) - return result except MCPMissingUserEnvVarsError as e: verbose_logger.info( "MCP tool call missing per-user env vars: server_id=%s missing=%s", diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 14673cf12c1..06a3a5a61e4 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2910,7 +2910,38 @@ if MCP_AVAILABLE: local_content = await _handle_local_mcp_tool(original_tool_name, arguments) response = CallToolResult(content=cast(Any, local_content), isError=False) - return response + return await _run_post_mcp_call_guardrails( + result=response, + litellm_logging_obj=litellm_logging_obj, + user_api_key_auth=user_api_key_auth, + request_data=kwargs, + ) + + async def _run_post_mcp_call_guardrails( + result: CallToolResult, + litellm_logging_obj: LiteLLMLoggingObj | None, + user_api_key_auth: UserAPIKeyAuth | None, + request_data: Mapping[str, object], + ) -> CallToolResult: + """Run ``post_mcp_call`` guardrails over an executed tool result. + + Lives on ``execute_mcp_tool``'s return path rather than inside + ``_fire_mcp_tool_call_logging`` so enforcement never depends on logging + being configured, and so every dispatch route gets it: the MCP protocol + handler, the REST endpoint, and tool search all funnel through here. + A guardrail that rejects the result raises, matching ``pre_mcp_call``. + """ + from litellm.proxy.proxy_server import proxy_logging_obj + + if proxy_logging_obj is None: + return result + return await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data=( + litellm_logging_obj.model_call_details if litellm_logging_obj is not None else dict(request_data) + ), + user_api_key_dict=user_api_key_auth, + ) _MCP_CREDENTIAL_REQUEST_FIELDS = frozenset( { @@ -2929,8 +2960,14 @@ if MCP_AVAILABLE: end_time: datetime, user_api_key_auth: UserAPIKeyAuth | None = None, request_data: Mapping[str, object] | None = None, - ) -> None: - """Fire post-call logging for an executed MCP tool call. + ) -> CallToolResult: + """Fire post-call logging for an executed MCP tool call, returning the result to send. + + The returned result is what the caller must forward to the client: a + ``post_mcp_call`` guardrail may rewrite the tool output (e.g. mask + sensitive values) or reject it, in which case its exception propagates. + Guardrails run before the success/failure logging so the masked text, not + the raw one, is what gets logged. A result with ``isError=True`` is logged as a failure (``status="failure"`` payload, so OTel marks the span ERROR) while the HTTP wire behavior stays @@ -2946,6 +2983,8 @@ if MCP_AVAILABLE: stripped before the dict is handed to ``post_call_failure_hook`` callbacks. """ + from litellm.proxy.proxy_server import proxy_logging_obj + logging_obj.post_call(original_response=result) await logging_obj.async_post_mcp_tool_call_hook( kwargs=logging_obj.model_call_details, @@ -2957,7 +2996,7 @@ if MCP_AVAILABLE: error_message = extract_mcp_tool_result_error_message(result) if error_message is None: await logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time) - return + return result logging_obj.has_run_logging(event_type="sync_success") logging_obj.has_run_logging(event_type="async_success") @@ -2966,8 +3005,7 @@ if MCP_AVAILABLE: await logging_obj.async_failure_handler(tool_error, "", start_time, end_time) if user_api_key_auth is None: - return - from litellm.proxy.proxy_server import proxy_logging_obj + return result if proxy_logging_obj: sanitized_request_data = { @@ -2979,6 +3017,7 @@ if MCP_AVAILABLE: user_api_key_dict=user_api_key_auth, route="/mcp/call_tool", ) + return result @client async def call_mcp_tool( @@ -3062,7 +3101,7 @@ if MCP_AVAILABLE: raise if litellm_logging_obj: - await _fire_mcp_tool_call_logging( + response = await _fire_mcp_tool_call_logging( logging_obj=litellm_logging_obj, result=response, start_time=start_time, diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 80a469b8c1a..afd396adc4c 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -4,6 +4,7 @@ MCP Server Utilities import json import re +from collections.abc import MutableMapping, MutableSequence from typing import ( Any, Dict, @@ -434,6 +435,56 @@ def extract_mcp_tool_result_error_message(result: object) -> Optional[str]: return "MCP tool call returned isError=true" +def mcp_tool_result_content_list(result: object) -> MutableSequence[object] | None: # mutable-ok: see below + """The mutable content list of an MCP tool result, or ``None`` when it has none. + + Deliberately mutable: a guardrail masking the result rewrites entries in place, + because the logging payload captured before the guardrail runs references this + same list, so handing back a copy would leave the unmasked text in the spend log + and the OTel span. + + Accepts both ``mcp.types.CallToolResult`` objects and their dict + equivalents, duck-typed so the ``mcp`` package is not required. + """ + content: object = result.get("content") if isinstance(result, Mapping) else getattr(result, "content", None) + if isinstance(content, MutableSequence): + return content + return None + + +def mcp_content_item_text(item: object) -> str | None: + """The ``text`` of a rewritable MCP content item, or ``None``. + + Only mappings and Pydantic-style models report a text, because those are the + only shapes ``with_mcp_content_item_text`` can rewrite; a caller therefore + never reads text it would be unable to write back (e.g. masked by a + guardrail). Non-text content (images, embedded resources) has no ``text`` + and is reported as ``None``. + """ + text: object + if isinstance(item, Mapping): + text = item.get("text") + elif callable(getattr(item, "model_copy", None)): + text = getattr(item, "text", None) + else: + return None + return text if isinstance(text, str) else None + + +def with_mcp_content_item_text(item: object, text: str) -> object: + """A copy of an MCP content item carrying ``text`` instead of its own. + + Only meaningful for items ``mcp_content_item_text`` returned a text for; any + other item is returned unchanged. + """ + if isinstance(item, Mapping): + return {**item, "text": text} + model_copy = getattr(item, "model_copy", None) + if callable(model_copy): + return model_copy(update={"text": text}) + return item + + TOOL_DISPLAY_NAME_PATTERN = re.compile(r"^[a-zA-Z0-9_-]+$") @@ -618,3 +669,112 @@ def merge_mcp_headers( merged.update({str(k): str(v) for k, v in static_headers.items()}) return merged or None + + +# Local rather than litellm.constants: this module deliberately imports no litellm +# package, so pulling one in for a single integer would drag in litellm/__init__. +MAX_STRUCTURED_CONTENT_SCAN_DEPTH = 100 + + +JSONLeafPath = tuple[str | int, ...] + + +def _flatten_leaf_groups( + groups: Iterable[tuple[tuple[JSONLeafPath, str], ...] | None], +) -> tuple[tuple[JSONLeafPath, str], ...] | None: + """Concatenate child leaf groups, propagating the too-deep sentinel.""" + materialized = tuple(groups) + if any(group is None for group in materialized): + return None + return tuple(leaf for group in materialized if group is not None for leaf in group) + + +def json_string_leaves(value: object, path: JSONLeafPath = ()) -> tuple[tuple[JSONLeafPath, str], ...] | None: + """Depth-first, deterministically ordered string leaves of a JSON value. + + Returns ``None`` when the value is nested past ``MAX_STRUCTURED_CONTENT_SCAN_DEPTH``, + so the caller blocks rather than letting deeper values through unscanned; an + empty tuple means there was simply nothing to scan. A sentinel rather than an + exception because this module is reloaded by tests (see the note above the + environment-backed constants), which would give a custom exception class a new + identity and let it escape a caller's ``except``. + """ + if len(path) > MAX_STRUCTURED_CONTENT_SCAN_DEPTH: + return None + if isinstance(value, str): + return ((path, value),) + if isinstance(value, dict): + return _flatten_leaf_groups(json_string_leaves(item, (*path, key)) for key, item in value.items()) + if isinstance(value, list): + return _flatten_leaf_groups(json_string_leaves(item, (*path, index)) for index, item in enumerate(value)) + return () + + +def with_json_string_leaves( + value: object, + replacements: Mapping[JSONLeafPath, str], + path: JSONLeafPath = (), +) -> object: + """Rebuild a JSON value with the guardrail's rewritten string leaves.""" + if isinstance(value, str): + return replacements.get(path, value) + if isinstance(value, dict): + return {key: with_json_string_leaves(item, replacements, (*path, key)) for key, item in value.items()} + if isinstance(value, list): + return [with_json_string_leaves(item, replacements, (*path, index)) for index, item in enumerate(value)] + return value + + +def json_unrewritable_labels(value: object, path_depth: int = 0) -> tuple[str, ...] | None: + """Strings in a JSON value that carry meaning but cannot be rewritten. + + Dictionary keys and non-string scalars: masking either would change the + payload's contract rather than redact a value, so a caller scans these and + blocks on a match instead of rewriting, matching what the content filter + already does for MCP tool call arguments. ``None`` means the value is nested + past the scan depth, same contract as ``json_string_leaves``. + """ + if path_depth > MAX_STRUCTURED_CONTENT_SCAN_DEPTH: + return None + if isinstance(value, bool) or value is None or isinstance(value, str): + return () + if isinstance(value, (int, float)): + return (str(value),) + if isinstance(value, dict): + own = tuple(key for key in value if isinstance(key, str)) + nested = tuple(json_unrewritable_labels(item, path_depth + 1) for item in value.values()) + if any(group is None for group in nested): + return None + return own + tuple(label for group in nested if group is not None for label in group) + if isinstance(value, list): + nested = tuple(json_unrewritable_labels(item, path_depth + 1) for item in value) + if any(group is None for group in nested): + return None + return tuple(label for group in nested if group is not None for label in group) + return () + + +def mcp_tool_result_structured_content(result: object) -> object: + """The ``structuredContent`` of an MCP tool result, or ``None`` when it has none.""" + if isinstance(result, Mapping): + return result.get("structuredContent") + return getattr(result, "structuredContent", None) + + +def set_mcp_tool_result_structured_content(result: object, value: object) -> bool: + """Replace ``structuredContent`` in place; ``False`` when the shape does not carry it. + + In place for the same reason the content list is: the logging payload captured + before the guardrail ran references this object, so a copy would leave the + unmasked value in the spend log and the OTel span. + """ + if isinstance(result, MutableMapping): + result["structuredContent"] = value + return True + if not hasattr(result, "structuredContent"): + return False + try: + setattr(result, "structuredContent", value) # attribute name is fixed by the MCP result shape + return True + except (AttributeError, TypeError, ValueError): + return False diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index dcd2de07e07..9e98cb46b9a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -53,6 +53,7 @@ from litellm.types.utils import ( StandardLoggingModelInformation, StandardLoggingPayloadErrorInformation, StandardLoggingPayloadStatus, + StandardLoggingRoutingDecision, StandardLoggingVectorStoreRequest, StandardPassThroughResponseObject, TextCompletionResponse, @@ -1169,7 +1170,6 @@ class GenerateKeyResponse(KeyRequestBase): class UpdateKeyRequest(KeyRequestBase): # Note: the defaults of all Params here MUST BE NONE # else they will get overwritten - key: str # type: ignore duration: Optional[str] = None spend: Optional[float] = None metadata: Optional[dict] = None @@ -1186,6 +1186,12 @@ class UpdateKeyRequest(KeyRequestBase): raise ValueError("temp_budget_increase and temp_budget_expiry must be set together") return self + @model_validator(mode="after") + def validate_key_identifier(self) -> "UpdateKeyRequest": + if self.key is None and self.key_alias is None: + raise ValueError("either key or key_alias must be provided") + return self + class RegenerateKeyRequest(GenerateKeyRequest): # This needs to be different from UpdateKeyRequest, because "key" is optional for this @@ -2455,16 +2461,6 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "is active as a reminder that hard enforcement is relaxed." ), ) - skip_user_budget_on_team_key: bool | None = Field( - None, - description=( - "If True, restores the legacy behavior where a user's personal " - "max_budget is NOT enforced when their key belongs to a team; only " - "the team (and team-member) budgets apply. Defaults to False, meaning " - "the user's personal max_budget is always enforced regardless of " - "whether the key belongs to a team (see GitHub issue #12905)." - ), - ) user_url_validation: Optional[bool] = Field( None, description=( @@ -2777,6 +2773,7 @@ class UserInfoV2Response(LiteLLMPydanticObjectBase): updated_at: Optional[datetime] = None sso_user_id: Optional[str] = None teams: List[str] = [] # Just team IDs, not full team objects + object_permission: LiteLLM_ObjectPermissionTable | None = None from litellm.models.config import LiteLLM_Config as LiteLLM_Config # noqa: E402 @@ -3306,6 +3303,7 @@ class SpendLogsMetadata(TypedDict): applied_guardrails: Optional[List[str]] mcp_tool_call_metadata: Optional[StandardLoggingMCPToolCall] vector_store_request_metadata: Optional[List[StandardLoggingVectorStoreRequest]] + routing_decision: StandardLoggingRoutingDecision | None guardrail_information: Optional[List[StandardLoggingGuardrailInformation]] eval_information: Optional[Any] status: StandardLoggingPayloadStatus @@ -4330,7 +4328,13 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): team_id_upsert: bool = False team_ids_jwt_field: Optional[str] = None upsert_sso_user_to_team: bool = False - team_allowed_routes: List[str] = ["openai_routes", "info_routes", "mcp_routes"] + team_allowed_routes: List[str] = [ + "openai_routes", + "info_routes", + "mcp_routes", + "/v1/messages", + "/v1/messages/count_tokens", + ] team_id_default: Optional[str] = Field( default=None, description="If no team_id given, default permissions/spend-tracking to this team.s", diff --git a/litellm/proxy/analytics_endpoints/analytics_endpoints.py b/litellm/proxy/analytics_endpoints/analytics_endpoints.py index 4c1ff31e5a1..cb22c468c1e 100644 --- a/litellm/proxy/analytics_endpoints/analytics_endpoints.py +++ b/litellm/proxy/analytics_endpoints/analytics_endpoints.py @@ -1,105 +1,61 @@ #### Analytics Endpoints ##### from datetime import datetime, timezone -from typing import List, Optional +from typing import Annotated import fastapi from fastapi import APIRouter, Depends, HTTPException, status from litellm.proxy._types import * +from litellm.proxy.analytics_endpoints.cache_activity import CacheActivityResponse, get_cache_activity from litellm.proxy.auth.user_api_key_auth import user_api_key_auth router = APIRouter() +def _parse_date(value: str, param_name: str) -> datetime: + try: + return datetime.strptime(value, "%Y-%m-%d").replace(tzinfo=timezone.utc) + except ValueError: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": f"{param_name} must be a YYYY-MM-DD date, got {value!r}"}, + ) + + @router.get( "/global/activity/cache_hits", tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], - responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, - }, + response_model=CacheActivityResponse, include_in_schema=False, ) async def get_global_activity( - start_date: Optional[str] = fastapi.Query( - default=None, - description="Time from which to start viewing spend", - ), - end_date: Optional[str] = fastapi.Query( - default=None, - description="Time till which to view spend", - ), -): + start_date: Annotated[str, fastapi.Query(description="Time from which to start viewing spend")], + end_date: Annotated[str, fastapi.Query(description="Time till which to view spend")], + key_aliases: Annotated[ + list[str] | None, fastapi.Query(description="Only include spend from these key aliases") + ] = None, + models: Annotated[list[str] | None, fastapi.Query(description="Only include spend for these models")] = None, +) -> CacheActivityResponse: """ - Get number of cache hits, vs misses - - { - "daily_data": [ - const chartdata = [ - { - date: 'Jan 22', - cache_hits: 10, - llm_api_calls: 2000 - }, - { - date: 'Jan 23', - cache_hits: 10, - llm_api_calls: 12 - }, - ], - "sum_cache_hits": 20, - "sum_llm_api_calls": 2012 - } + Cache activity for the Admin UI cache dashboard, aggregated per call_type: + cache hits vs successful LLM API requests vs failed requests, plus totals + for the stat cards and the available key-alias/model filter options. """ - - if start_date is None or end_date is None: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": "Please provide start_date and end_date"}, - ) - - start_date_obj = datetime.strptime(start_date, "%Y-%m-%d").replace(tzinfo=timezone.utc) - end_date_obj = datetime.strptime(end_date, "%Y-%m-%d").replace(tzinfo=timezone.utc) - from litellm.proxy.proxy_server import prisma_client - try: - if prisma_client is None: - raise ValueError( - "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" - ) - - sql_query = """ - SELECT - CASE - WHEN vt."key_alias" IS NOT NULL THEN vt."key_alias" - ELSE 'Unnamed Key' - END AS api_key, - sl."call_type", - sl."model", - COUNT(*) AS total_rows, - SUM(CASE WHEN sl."cache_hit" = 'True' THEN 1 ELSE 0 END) AS cache_hit_true_rows, - SUM(CASE WHEN sl."cache_hit" = 'True' THEN sl."completion_tokens" ELSE 0 END) AS cached_completion_tokens, - SUM(CASE WHEN sl."cache_hit" != 'True' THEN sl."completion_tokens" ELSE 0 END) AS generated_completion_tokens - FROM "LiteLLM_SpendLogs" sl - LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token" - WHERE - sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') - AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC') - GROUP BY - vt."key_alias", - sl."call_type", - sl."model" - """ - db_response = await prisma_client.db.query_raw(sql_query, start_date_obj, end_date_obj) - - if db_response is None: - return [] - - return db_response - - except Exception as e: + if prisma_client is None: raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": str(e)}, + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail={ + "error": "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" + }, ) + + return await get_cache_activity( + prisma_client=prisma_client, + start_date=_parse_date(start_date, "start_date"), + end_date=_parse_date(end_date, "end_date"), + key_aliases=key_aliases or [], + models=models or [], + ) diff --git a/litellm/proxy/analytics_endpoints/cache_activity.py b/litellm/proxy/analytics_endpoints/cache_activity.py new file mode 100644 index 00000000000..6e20382f56e --- /dev/null +++ b/litellm/proxy/analytics_endpoints/cache_activity.py @@ -0,0 +1,137 @@ +import asyncio +import json +from datetime import datetime +from typing import TYPE_CHECKING, Sequence + +from pydantic import BaseModel, TypeAdapter + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +UNKNOWN_CALL_TYPE = "Unknown" + + +class CacheActivityGroup(BaseModel): + call_type: str + api_requests: int + cache_hits: int + failed_requests: int + cached_completion_tokens: int + generated_completion_tokens: int + + +class CacheActivityTotals(BaseModel): + api_requests: int + cache_hits: int + failed_requests: int + cached_completion_tokens: int + cache_hit_ratio: float + + +class CacheActivityFilterOptions(BaseModel): + key_aliases: list[str] + models: list[str] + + +class CacheActivityResponse(BaseModel): + groups: list[CacheActivityGroup] + totals: CacheActivityTotals + filter_options: CacheActivityFilterOptions + + +GROUPS_SQL = """ + SELECT + CASE WHEN sl."call_type" = '' THEN 'Unknown' ELSE sl."call_type" END AS call_type, + (COUNT(*) + - SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN 1 ELSE 0 END) + - SUM(CASE WHEN sl."status" = 'failure' THEN 1 ELSE 0 END))::int AS api_requests, + SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN 1 ELSE 0 END)::int AS cache_hits, + SUM(CASE WHEN sl."status" = 'failure' THEN 1 ELSE 0 END)::int AS failed_requests, + SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN sl."completion_tokens" ELSE 0 END)::int + AS cached_completion_tokens, + SUM(CASE WHEN COALESCE(sl."cache_hit", '') != 'True' THEN sl."completion_tokens" ELSE 0 END)::int + AS generated_completion_tokens + FROM "LiteLLM_SpendLogs" sl + LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token" + WHERE + sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') + AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC') + AND ($3::jsonb = '[]'::jsonb + OR COALESCE(vt."key_alias", 'Unnamed Key') IN (SELECT jsonb_array_elements_text($3::jsonb))) + AND ($4::jsonb = '[]'::jsonb + OR sl."model" IN (SELECT jsonb_array_elements_text($4::jsonb))) + GROUP BY 1 + ORDER BY (COUNT(*)) DESC +""" + +KEY_ALIAS_OPTIONS_SQL = """ + SELECT DISTINCT COALESCE(vt."key_alias", 'Unnamed Key') AS key_alias + FROM "LiteLLM_SpendLogs" sl + LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token" + WHERE + sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') + AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC') + ORDER BY 1 +""" + +MODEL_OPTIONS_SQL = """ + SELECT DISTINCT sl."model" AS model + FROM "LiteLLM_SpendLogs" sl + WHERE + sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') + AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC') + AND sl."model" != '' + ORDER BY 1 +""" + + +class _KeyAliasRow(BaseModel): + key_alias: str + + +class _ModelRow(BaseModel): + model: str + + +_groups_adapter = TypeAdapter(list[CacheActivityGroup]) +_key_alias_rows_adapter = TypeAdapter(list[_KeyAliasRow]) +_model_rows_adapter = TypeAdapter(list[_ModelRow]) + + +def compute_totals(groups: Sequence[CacheActivityGroup]) -> CacheActivityTotals: + api_requests = sum(group.api_requests for group in groups) + cache_hits = sum(group.cache_hits for group in groups) + failed_requests = sum(group.failed_requests for group in groups) + all_requests = api_requests + cache_hits + failed_requests + return CacheActivityTotals( + api_requests=api_requests, + cache_hits=cache_hits, + failed_requests=failed_requests, + cached_completion_tokens=sum(group.cached_completion_tokens for group in groups), + cache_hit_ratio=(cache_hits / all_requests) * 100 if all_requests > 0 else 0.0, + ) + + +async def get_cache_activity( + prisma_client: "PrismaClient", + start_date: datetime, + end_date: datetime, + key_aliases: Sequence[str], + models: Sequence[str], +) -> CacheActivityResponse: + group_rows, key_alias_rows, model_rows = await asyncio.gather( + prisma_client.db.query_raw( + GROUPS_SQL, start_date, end_date, json.dumps(list(key_aliases)), json.dumps(list(models)) + ), + prisma_client.db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date), + prisma_client.db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date), + ) + groups = _groups_adapter.validate_python(group_rows or []) + return CacheActivityResponse( + groups=groups, + totals=compute_totals(groups), + filter_options=CacheActivityFilterOptions( + key_aliases=[row.key_alias for row in _key_alias_rows_adapter.validate_python(key_alias_rows or [])], + models=[row.model for row in _model_rows_adapter.validate_python(model_rows or [])], + ), + ) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 07dfdc4fb43..263fec77d12 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -74,6 +74,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, get_management_object_ttl, + object_permission_cache_key, ) from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.guardrails.tool_name_extraction import ( @@ -485,6 +486,14 @@ MODEL_DISCOVERY_ROUTES = frozenset( } ) +BUDGET_ENFORCED_SIDE_EFFECT_ROUTES = frozenset( + { + "/health", + "/health/services", + "/health/test_connection", + } +) + async def common_checks( request_body: dict, @@ -531,8 +540,10 @@ async def common_checks( request=request, ) - if route in MODEL_DISCOVERY_ROUTES: - skip_budget_checks = True + skip_all_budget_checks = skip_budget_checks or ( + route not in BUDGET_ENFORCED_SIDE_EFFECT_ROUTES + and (route in MODEL_DISCOVERY_ROUTES or not RouteChecks.is_llm_api_route(route=route)) + ) # 1. If team is blocked if team_object is not None and team_object.blocked is True: @@ -606,7 +617,7 @@ async def common_checks( project_object=project_object, _model=_model, llm_router=llm_router, - skip_budget_checks=skip_budget_checks, + skip_budget_checks=skip_all_budget_checks, valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -615,7 +626,7 @@ async def common_checks( _reject_clientside_metadata_tags_check(general_settings, request_body, route) # If this is a free model, skip all budget checks - if not skip_budget_checks: + if not skip_all_budget_checks: # Key metadata.tags are injected into request_body here so the tag budget # check can read them; this mutation must run before the gathered checks. if valid_token is not None: @@ -632,31 +643,28 @@ async def common_checks( ) async def _user_max_budget_check() -> None: - if user_object is None or user_object.max_budget is None: - return - skip_for_team = ( - general_settings.get("skip_user_budget_on_team_key") is True - and team_object is not None - and team_object.team_id is not None - ) - if skip_for_team: - return - from litellm.proxy.proxy_server import get_current_spend + # 4.1 personal budget, if personal key + if ( + (team_object is None or team_object.team_id is None) + and user_object is not None + and user_object.max_budget is not None + ): + from litellm.proxy.proxy_server import get_current_spend - user_budget = user_object.max_budget - user_spend = await get_current_spend( - counter_key=f"spend:user:{user_object.user_id}", - fallback_spend=user_object.spend or 0.0, - max_budget=user_budget, - ) - if math.isfinite(user_budget) and user_spend >= user_budget: - raise litellm.BudgetExceededError( - current_cost=user_spend, + user_budget = user_object.max_budget + user_spend = await get_current_spend( + counter_key=f"spend:user:{user_object.user_id}", + fallback_spend=user_object.spend or 0.0, max_budget=user_budget, - message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}", - entity_type=Litellm_EntityType.USER.value, - entity_id=user_object.user_id, ) + if math.isfinite(user_budget) and user_spend >= user_budget: + raise litellm.BudgetExceededError( + current_cost=user_spend, + max_budget=user_budget, + message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}", + entity_type=Litellm_EntityType.USER.value, + entity_id=user_object.user_id, + ) # Each scope reads a distinct counter key with no cross-scope ordering # dependency, so the per-scope Redis-first reads run concurrently instead @@ -715,7 +723,7 @@ async def common_checks( raise budget_error _enforce_user_param_check(general_settings, request, request_body, route) - _global_proxy_budget_check(global_proxy_spend, skip_budget_checks, route) + _global_proxy_budget_check(global_proxy_spend, skip_all_budget_checks, route) _guardrail_modification_check(request_body, team_object) # 10 [OPTIONAL] Organization RBAC checks @@ -953,7 +961,7 @@ async def get_default_end_user_budget( ) return None - _budget_obj = LiteLLM_BudgetTable(**budget_record.dict()) + _budget_obj = LiteLLM_BudgetTable.model_validate(budget_record.dict()) # Cache the budget for 60 seconds await user_api_key_cache.async_set_cache( key=cache_key, @@ -999,7 +1007,7 @@ async def get_team_member_default_budget( if isinstance(cached_budget, LiteLLM_BudgetTable): return cached_budget if isinstance(cached_budget, dict): - return LiteLLM_BudgetTable(**cached_budget) + return LiteLLM_BudgetTable.model_validate(cached_budget) try: budget_record = await BudgetRepository(prisma_client).table.find_unique(where={"budget_id": budget_id}) @@ -1014,7 +1022,7 @@ async def get_team_member_default_budget( ttl=get_management_object_ttl(user_api_key_cache), ) - return LiteLLM_BudgetTable(**budget_record.dict()) + return LiteLLM_BudgetTable.model_validate(budget_record.dict()) except Exception: verbose_proxy_logger.exception(f"Error fetching team-default member budget {budget_id}") @@ -1168,7 +1176,7 @@ async def get_end_user_object( raise Exception # Convert to LiteLLM_EndUserTable object - _response = LiteLLM_EndUserTable(**response.dict()) + _response = LiteLLM_EndUserTable.model_validate(response.dict()) # Apply default budget if needed _response = await _apply_default_budget_to_end_user( @@ -1360,7 +1368,7 @@ async def get_tag_objects_batch( for db_tag in db_tags: tag_name = db_tag.tag_name cache_key = f"tag:{tag_name}" - _tag_obj = LiteLLM_TagTable(**db_tag.dict()) + _tag_obj = LiteLLM_TagTable.model_validate(db_tag.dict()) await user_api_key_cache.async_set_cache( key=cache_key, value=_tag_obj, @@ -1453,7 +1461,7 @@ async def get_team_membership( if response is None: return None - _response = LiteLLM_TeamMembership(**response.dict()) + _response = LiteLLM_TeamMembership.model_validate(response.dict()) await user_api_key_cache.async_set_cache( key=_key, value=_response, @@ -1719,13 +1727,13 @@ async def get_user_object( if response.organization_memberships is not None and len(response.organization_memberships) > 0: # dump each organization membership to type LiteLLM_OrganizationMembershipTable _dumped_memberships = [ - LiteLLM_OrganizationMembershipTable(**membership.model_dump()) + LiteLLM_OrganizationMembershipTable.model_validate(membership.model_dump()) for membership in response.organization_memberships if membership is not None ] response.organization_memberships = _dumped_memberships - _response = LiteLLM_UserTable(**dict(response)) + _response = LiteLLM_UserTable.model_validate(dict(response)) response_dict = _response.model_dump() # save the user object to cache @@ -1781,9 +1789,22 @@ async def _cache_team_object( ## CACHE REFRESH TIME! team_table.last_refreshed_at = time.time() + key = "team_id:{}".format(team_id) + + if proxy_logging_obj is not None: + try: + await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key) + except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the write + verbose_proxy_logger.warning( + "Failed to invalidate internal usage cache entry %s; " + "a stale team object may be served until its TTL expires: %s", + key, + e, + ) + # team_id is the table primary key — guaranteed unique, safe to write. await _cache_management_object( - key="team_id:{}".format(team_id), + key=key, value=team_table, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -1805,9 +1826,17 @@ async def _cache_team_object( # the cache from a verified single row. if team_table.team_alias: alias_key = "team_alias:{}".format(team_table.team_alias) - user_api_key_cache.delete_cache(key=alias_key) - if proxy_logging_obj is not None: - await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key) + try: + user_api_key_cache.delete_cache(key=alias_key) + if proxy_logging_obj is not None: + await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key) + except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the mutation + verbose_proxy_logger.warning( + "Failed to invalidate cached team alias entry %s; " + "a stale team object may be served until its TTL expires: %s", + alias_key, + e, + ) async def _cache_key_object( @@ -1862,7 +1891,7 @@ async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_ http_request=mock_request, user_api_key_dict=system_admin_user, ) - response = LiteLLM_TeamTable(**created_team_dict) + response = LiteLLM_TeamTable.model_validate(created_team_dict) return response @@ -1894,7 +1923,7 @@ async def _get_team_object_from_user_api_key_cache( if response is None: raise Exception - _response = LiteLLM_TeamTableCachedObj(**response.dict()) + _response = LiteLLM_TeamTableCachedObj.model_validate(response.dict()) # Load object_permission if object_permission_id exists but object_permission is not loaded if _response.object_permission_id and not _response.object_permission: @@ -2085,7 +2114,7 @@ async def get_access_object( detail={"error": f"Access group doesn't exist in db. Access group={access_group_id}."}, ) - _response = LiteLLM_AccessGroupTable(**response.dict()) + _response = LiteLLM_AccessGroupTable.model_validate(response.dict()) # Save to cache await _cache_access_object( @@ -2170,7 +2199,7 @@ async def get_team_object_by_alias( ) team = teams[0] - team_obj = LiteLLM_TeamTableCachedObj(**team.model_dump()) + team_obj = LiteLLM_TeamTableCachedObj.model_validate(team.model_dump()) # Load object_permission if object_permission_id exists but object_permission is not loaded if team_obj.object_permission_id and not team_obj.object_permission: @@ -2272,7 +2301,7 @@ async def get_org_object_by_alias( ) org = orgs[0] - org_obj = LiteLLM_OrganizationTable(**org.model_dump()) + org_obj = LiteLLM_OrganizationTable.model_validate(org.model_dump()) # Cache the result await user_api_key_cache.async_set_cache( @@ -2413,7 +2442,7 @@ class ExperimentalUIJWTToken: if decrypted_token is None: return None try: - return UserAPIKeyAuth(**json.loads(decrypted_token)) + return UserAPIKeyAuth.model_validate(json.loads(decrypted_token)) except Exception as e: raise Exception(f"Invalid hash key. Hash key={hashed_token}. Decrypted token={decrypted_token}. Error: {e}") @@ -2532,7 +2561,7 @@ async def get_key_object( code=status.HTTP_401_UNAUTHORIZED, ) - _response = UserAPIKeyAuth(**_valid_token.model_dump(exclude_none=True)) + _response = UserAPIKeyAuth.model_validate(_valid_token.model_dump(exclude_none=True)) # Load object_permission if object_permission_id exists but object_permission is not loaded if _response.object_permission_id and not _response.object_permission: @@ -2588,7 +2617,7 @@ async def get_object_permission( raise Exception("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys") # check if in cache - key = "object_permission_id:{}".format(object_permission_id) + key = object_permission_cache_key(object_permission_id) deserialized_perm = await user_api_key_cache.async_get_cache( key=key, model_type=LiteLLM_ObjectPermissionTable, @@ -2605,7 +2634,7 @@ async def get_object_permission( if response is None: return None - _perm_obj = LiteLLM_ObjectPermissionTable(**response.dict()) + _perm_obj = LiteLLM_ObjectPermissionTable.model_validate(response.dict()) await user_api_key_cache.async_set_cache( key=key, value=_perm_obj, @@ -2665,7 +2694,7 @@ async def get_managed_vector_store_rows_by_uuids( row_dict = dict(row) if hasattr(row, "__dict__") else {} if not row_dict: continue - cached_obj = LiteLLM_ManagedVectorStoresTable(**row_dict) + cached_obj = LiteLLM_ManagedVectorStoresTable.model_validate(row_dict) key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id) await user_api_key_cache.async_set_cache( key=key, @@ -2746,7 +2775,7 @@ async def get_org_object( f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call." ) - _org_obj = LiteLLM_OrganizationTable(**response.model_dump()) + _org_obj = LiteLLM_OrganizationTable.model_validate(response.model_dump()) # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, @@ -4221,7 +4250,7 @@ async def get_project_object( if project_row is None: return None - project_obj = LiteLLM_ProjectTableCachedObj(**project_row.model_dump()) + project_obj = LiteLLM_ProjectTableCachedObj.model_validate(project_row.model_dump()) # Cache with TTL following _cache_management_object pattern project_obj.last_refreshed_at = time.time() diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index ecb37e67c14..644253ceac7 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1432,7 +1432,7 @@ def _extract_models_from_managed_resource_id( ) _append_model_candidates( candidates=candidates, - value=get_model_id_from_unified_batch_id(unified_file_id), + value=_resolve_model_id_with_router(get_model_id_from_unified_batch_id(unified_file_id), llm_router), ) except Exception as e: verbose_proxy_logger.debug("Unable to extract model from managed file/batch ID: %s", str(e)) @@ -1442,7 +1442,10 @@ def _extract_models_from_managed_resource_id( parsed_id = parse_unified_id(resource_id) if parsed_id: - _append_model_candidates(candidates=candidates, value=parsed_id.get("model_id")) + _append_model_candidates( + candidates=candidates, + value=_resolve_model_id_with_router(parsed_id.get("model_id"), llm_router), + ) _append_model_candidates(candidates=candidates, value=parsed_id.get("target_model_names")) except Exception as e: verbose_proxy_logger.debug("Unable to extract model from unified managed resource ID: %s", str(e)) diff --git a/litellm/proxy/auth/oauth2_proxy_hook.py b/litellm/proxy/auth/oauth2_proxy_hook.py index 2b0593d3618..ca6a7ee4b1d 100644 --- a/litellm/proxy/auth/oauth2_proxy_hook.py +++ b/litellm/proxy/auth/oauth2_proxy_hook.py @@ -1,4 +1,5 @@ -from typing import Any, Dict, FrozenSet +from collections.abc import Mapping +from typing import Dict, FrozenSet, List, Union from fastapi import Request @@ -83,21 +84,17 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth: "(signature-validated) instead of header-trust." ) - auth_data: Dict[str, Any] = {} - for key, header in oauth2_config_mappings.items(): - value = request.headers.get(header) - if not value: - continue - if key == "models": - auth_data[key] = [model.strip() for model in value.split(",")] - else: - auth_data[key] = value + auth_data: Mapping[str, Union[str, List[str]]] = { + key: [model.strip() for model in value.split(",")] if key == "models" else value + for key, header in oauth2_config_mappings.items() + if (value := request.headers.get(header)) + } verbose_proxy_logger.debug( "Auth data before creating UserAPIKeyAuth object: keys=%s", list(auth_data.keys()), ) - user_api_key_auth = UserAPIKeyAuth(**auth_data) + user_api_key_auth = UserAPIKeyAuth.model_validate(auth_data) verbose_proxy_logger.debug( "UserAPIKeyAuth object created with keys: %s", list(user_api_key_auth.__fields_set__), diff --git a/litellm/proxy/auth/resolvers/store.py b/litellm/proxy/auth/resolvers/store.py index ad9fd234163..7c2bd324064 100644 --- a/litellm/proxy/auth/resolvers/store.py +++ b/litellm/proxy/auth/resolvers/store.py @@ -118,7 +118,7 @@ class IdentityStore: if from_db is None: raise KeyNotFoundError(hashed_token) - key = UserAPIKeyAuth(**from_db.model_dump(exclude_none=True)) + key = UserAPIKeyAuth.model_validate(from_db.model_dump(exclude_none=True)) if key.object_permission_id and not key.object_permission: try: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e7c12448435..f4d07c1a674 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1960,17 +1960,6 @@ async def _user_api_key_auth_builder( else: valid_token.team_object_permission = None - # Cache under the canonical "team_id:{id}" key so get_team_object and - # _update_team_cache serve this write from the L2 cache. The guard keeps a - # non-team (personal) key, whose team_id is None, from reaching the cache - # layer, which Redis rejects with a NoneType key error. - if valid_token.team_id is not None and _team_obj is not None: - await user_api_key_cache.async_set_cache( - key=f"team_id:{valid_token.team_id}", - value=_team_obj, - model_type=LiteLLM_TeamTableCachedObj, - ) - # Fetch project object if key belongs to a project _project_obj = None if valid_token.project_id is not None: @@ -2461,7 +2450,6 @@ async def _reserve_budget_after_common_checks( proxy_logging_obj=proxy_logging_obj, end_user_id=end_user_id, end_user_object=end_user_object, - skip_user_budget_on_team_key=general_settings.get("skip_user_budget_on_team_key") is True, fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True, ) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index a2dc5e1caf5..a91b29002e3 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -23,6 +23,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( ) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, + apply_team_provider_credentials, decode_model_from_file_id, encode_batch_response_ids, encode_file_id_with_model, @@ -295,6 +296,12 @@ async def create_batch( verbose_proxy_logger.debug(f"Created batch using model: {model_param}") else: # SCENARIO 3: Fallback to custom_llm_provider (uses env variables) + apply_team_provider_credentials( + data=cast(dict, _create_batch_data), # cast-ok: TypedDict is a dict at runtime + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) response = await litellm.acreate_batch( custom_llm_provider=custom_llm_provider, **_create_batch_data, # type: ignore @@ -525,6 +532,12 @@ async def retrieve_batch( or get_custom_llm_provider_from_request_query(request=request) or "openai" ) + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) response = await litellm.aretrieve_batch( custom_llm_provider=custom_llm_provider, **data, # type: ignore @@ -718,6 +731,12 @@ async def list_batches( or get_custom_llm_provider_from_request_query(request=request) or "openai" ) + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) response = await litellm.alist_batches( custom_llm_provider=custom_llm_provider, # type: ignore after=after, @@ -908,6 +927,12 @@ async def cancel_batch( # Extract batch_id from data to avoid "multiple values for keyword argument" error # data was cast from CancelBatchRequest which already contains batch_id data.pop("batch_id", None) + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) _cancel_batch_data = CancelBatchRequest(batch_id=batch_id, **data) response = await litellm.acancel_batch( custom_llm_provider=custom_llm_provider, # type: ignore diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index 2ad8a08b8c3..de9d38963c1 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -10,11 +10,32 @@ uv tool install 'litellm[proxy]' ## Configuration -The CLI can be configured using environment variables or command-line options: +The CLI can be configured using environment variables, command-line options, or a persistent config file: - `LITELLM_PROXY_URL`: Base URL of the LiteLLM proxy server (default: http://localhost:4000) - `LITELLM_PROXY_API_KEY`: API key for authentication +To stop exporting `LITELLM_PROXY_URL` in every shell session, store the proxy URL once in `~/.litellm/config.json`: + +```bash +lite config set base_url https://your-proxy.example.com +``` + +Manage the stored config with: + +```bash +lite config get base_url # print the stored value +lite config get # print all stored config +lite config unset base_url # remove the stored value +``` + +The base URL is resolved in this order of precedence: + +1. `--base-url` command-line option +2. `LITELLM_PROXY_URL` environment variable +3. `base_url` from `~/.litellm/config.json` +4. `http://localhost:4000` + ## Global Options - `--version`, `-v`: Print the LiteLLM Proxy client and server version and exit. @@ -581,6 +602,8 @@ The CLI respects the following environment variables: - `LITELLM_PROXY_URL`: Base URL of the proxy server - `LITELLM_PROXY_API_KEY`: API key for authentication +`LITELLM_PROXY_URL` takes precedence over a `base_url` stored via `lite config set`, and the `--base-url` option overrides both. See the Configuration section for the full precedence order. + ## Examples 1. List all models in table format: diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 61495403407..970d801dc6d 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -15,6 +15,8 @@ from rich.table import Table from litellm.constants import CLI_JWT_EXPIRATION_HOURS from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh +from .private_json import write_private_json + # Token storage utilities def get_token_file_path() -> str: @@ -27,11 +29,7 @@ def get_token_file_path() -> str: def save_token(token_data: Dict[str, Any]) -> None: """Save token data to file""" - token_file = get_token_file_path() - with open(token_file, "w") as f: - json.dump(token_data, f, indent=2) - # Set file permissions to be readable only by owner - os.chmod(token_file, 0o600) + write_private_json(get_token_file_path(), token_data) def load_token() -> Optional[Dict[str, Any]]: diff --git a/litellm/proxy/client/cli/commands/config.py b/litellm/proxy/client/cli/commands/config.py new file mode 100644 index 00000000000..851a6c11529 --- /dev/null +++ b/litellm/proxy/client/cli/commands/config.py @@ -0,0 +1,108 @@ +import json +import os +import sys +from collections.abc import Mapping +from pathlib import Path +from urllib.parse import urlparse + +import click +from pydantic import TypeAdapter + +from .private_json import write_private_json + +ALLOWED_CONFIG_KEYS: tuple[str, ...] = ("base_url",) + +_config_adapter: TypeAdapter[Mapping[str, str]] = TypeAdapter(Mapping[str, str]) + + +def get_config_file_path() -> str: + """Get the path to the persistent CLI config file""" + home_dir = Path.home() + config_dir = home_dir / ".litellm" + return str(config_dir / "config.json") + + +def load_config() -> Mapping[str, str]: + """Load CLI config from file; returns {} if missing or unreadable""" + try: + config_file = get_config_file_path() + except RuntimeError: + return {} + if not os.path.exists(config_file): + return {} + try: + with open(config_file, "r") as f: + return _config_adapter.validate_python(json.load(f)) + except (OSError, ValueError) as e: + click.echo(f"Warning: ignoring invalid config file {config_file}: {e}", err=True) + return {} + + +def save_config(config: Mapping[str, str]) -> None: + """Save CLI config to file""" + write_private_json(get_config_file_path(), config) + + +def get_config_value(key: str) -> str | None: + """Get a single value from the persistent CLI config""" + return load_config().get(key) + + +@click.group(name="config") +def config_commands() -> None: + """Manage persistent CLI configuration (~/.litellm/config.json)""" + + +@config_commands.command(name="set") +@click.argument("key") +@click.argument("value") +def set_config(key: str, value: str) -> None: + """Set a config KEY to VALUE (e.g. `lite config set base_url https://your-proxy.example.com`)""" + if key not in ALLOWED_CONFIG_KEYS: + raise click.UsageError(f"Unknown config key '{key}'. Allowed keys: {', '.join(ALLOWED_CONFIG_KEYS)}") + + if key == "base_url": + parsed = urlparse(value) + if parsed.scheme not in ("http", "https") or not parsed.netloc: + raise click.UsageError("base_url must be a full http:// or https:// URL including a host") + if "?" in value or "#" in value: + raise click.UsageError("base_url must not include a query string or fragment") + + normalized_value = value.rstrip("/") + save_config({**load_config(), key: normalized_value}) + click.echo(f"Set {key} = {normalized_value} in {get_config_file_path()}") + + +@config_commands.command(name="get") +@click.argument("key", required=False) +def get_config(key: str | None) -> None: + """Print the value of KEY, or all stored config when KEY is omitted""" + config = load_config() + + if key is not None: + value = config.get(key) + if value is None: + click.echo(f"{key} is not set", err=True) + sys.exit(1) + click.echo(value) + return + + if not config: + click.echo("(no config set)") + return + + for entry_key, entry_value in config.items(): + click.echo(f"{entry_key} = {entry_value}") + + +@config_commands.command(name="unset") +@click.argument("key") +def unset_config(key: str) -> None: + """Remove KEY from the config file""" + config = load_config() + if key not in config: + click.echo(f"{key} was not set") + return + + save_config({k: v for k, v in config.items() if k != key}) + click.echo(f"Removed {key} from {get_config_file_path()}") diff --git a/litellm/proxy/client/cli/commands/private_json.py b/litellm/proxy/client/cli/commands/private_json.py new file mode 100644 index 00000000000..70aac0c6de0 --- /dev/null +++ b/litellm/proxy/client/cli/commands/private_json.py @@ -0,0 +1,20 @@ +import json +import os +import tempfile +from collections.abc import Mapping +from pathlib import Path + + +def write_private_json(path: str, data: Mapping[str, object]) -> None: + """Atomically write JSON to path with owner-only permissions (0600)""" + parent = Path(path).parent + parent.mkdir(parents=True, exist_ok=True) + fd, tmp_path = tempfile.mkstemp(dir=str(parent), prefix=".tmp-", suffix=".json") + try: + with os.fdopen(fd, "w") as f: + json.dump(data, f, indent=2) + f.flush() + os.fsync(f.fileno()) + os.replace(tmp_path, path) + finally: + Path(tmp_path).unlink(missing_ok=True) diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index e641956b2c5..24e5cdf747b 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -11,6 +11,7 @@ from .commands.agents import agent_commands from .commands.auth import auth_group, get_stored_api_key, login, logout, whoami from .commands.autoroute.commands import autoroute_group from .commands.chat import chat +from .commands.config import config_commands, get_config_value from .commands.credentials import credentials from .commands.encryption import encryption from .commands.http import http @@ -45,27 +46,16 @@ def print_version(base_url: str, api_key: Optional[str]): @click.option( "--version", "-v", + "show_version", is_flag=True, - is_eager=True, - expose_value=False, help="Show the LiteLLM Proxy CLI and server version and exit.", - callback=lambda ctx, param, value: ( - ( - print_version( - ctx.params.get("base_url") or "http://localhost:4000", - ctx.params.get("api_key"), - ) - or ctx.exit() - ) - if value and not ctx.resilient_parsing - else None - ), ) @click.option( "--base-url", envvar="LITELLM_PROXY_URL", show_envvar=True, - default="http://localhost:4000", + default=None, + show_default="base_url from `lite config`, else http://localhost:4000", help="Base URL of the LiteLLM proxy server", ) @click.option( @@ -75,13 +65,16 @@ def print_version(base_url: str, api_key: Optional[str]): help="API key for authentication", ) @click.pass_context -def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None: +def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: Optional[str]) -> None: """LiteLLM Proxy CLI - Manage your LiteLLM proxy server""" ctx.ensure_object(dict) + stored_base_url = get_config_value("base_url") + base_url_provided = base_url is not None + # Normalize once here so every downstream command (login, agents, http, ...) can safely # do f"{base_url}/some/path" without producing a double slash. - base_url = base_url.rstrip("/") + base_url = ((stored_base_url or "http://localhost:4000") if base_url is None else base_url).rstrip("/") # If no API key provided via flag or environment variable, try to load from saved token. # Pass base_url so we only use the stored key when it was issued for this server. @@ -94,8 +87,13 @@ def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None: # apiKeyHelper is invoked bare (no flags) -- commands that must work # unattended (print-token) need to tell "user didn't say" apart from # "user said localhost:4000 on purpose" so they can fall back to - # whatever server the stored token was actually issued for. - ctx.obj["base_url_explicit"] = ctx.get_parameter_source("base_url") != click.core.ParameterSource.DEFAULT + # whatever server the stored token was actually issued for. A base_url + # saved via `lite config set` counts as the user saying it. + ctx.obj["base_url_explicit"] = base_url_provided or bool(stored_base_url) + + if show_version: + print_version(base_url, api_key) + ctx.exit() # If no subcommand was invoked, start interactive mode if ctx.invoked_subcommand is None: @@ -141,6 +139,7 @@ cli.add_command(down) cli.add_command(model_groups) # Add the autoroute command group (QA auto-routing against your real proxy) cli.add_command(autoroute_group, name="autoroute") +cli.add_command(config_commands) if __name__ == "__main__": diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index 09921a3ac1d..dbb5b2c24d0 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -150,6 +150,28 @@ class UserApiKeyCache(DualCache): return await super().async_set_cache_pipeline(cache_list=normalized, local_only=local_only, **kwargs) +#: Value cached under ``user_object_permission_id_cache_key`` when the user links no permission row, +#: so a human without an entitlement costs no DB read per request. Lives beside the key builder +#: because it is part of the same cache protocol: a reader that knows the key must know this value. +USER_NO_MCP_PERMISSION_SENTINEL = "__user_no_mcp_permission__" + + +def user_object_permission_id_cache_key(user_id: str) -> str: + """Cache key for the ``user_id -> object_permission_id`` link. + + Lives here rather than next to either user because two modules own the two halves: the MCP auth + resolver writes it on read, and ``/user/update`` deletes it after changing the link. A key format + duplicated across those two drifts silently, and the failure is an entitlement change that never + takes effect. + """ + return f"user_object_permission_id:{user_id}" + + +def object_permission_cache_key(object_permission_id: str) -> str: + """Cache key ``get_object_permission`` stores a permission row under.""" + return f"object_permission_id:{object_permission_id}" + + def get_management_object_ttl(cache: DualCache) -> float: """ In-memory TTL for management-object cache writes (keys, teams, users, budgets, ...). diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index cc608d6e82c..1e0d8f5e010 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -834,6 +834,21 @@ class PrismaManager: dname = os.path.dirname(os.path.dirname(abspath)) return dname + @staticmethod + def _apply_replica_identity_full_if_requested() -> None: + """ + `prisma db push` bypasses litellm-proxy-extras, so the opt-in + REPLICA IDENTITY FULL step has to be driven from here too. + + litellm-proxy-extras is an optional install, so this is a no-op when it + is absent. + """ + try: + from litellm_proxy_extras.utils import ProxyExtrasDBManager + except ImportError: + return + ProxyExtrasDBManager.apply_replica_identity_full_if_requested() + @staticmethod def setup_database(use_migrate: bool = False, use_v2_resolver: bool = False) -> bool: """ @@ -880,6 +895,7 @@ class PrismaManager: timeout=60, check=True, ) + PrismaManager._apply_replica_identity_full_if_requested() return True except subprocess.TimeoutExpired: verbose_proxy_logger.warning(f"Attempt {attempt + 1} timed out") diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index a0c822964a0..60e1947b752 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -75,6 +75,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): GuardrailEventHooks.post_call, GuardrailEventHooks.logging_only, GuardrailEventHooks.pre_mcp_call, + GuardrailEventHooks.post_mcp_call, ] # Class variables or attributes diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index f72492881b3..719496785dd 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -28,6 +28,7 @@ from litellm import DualCache from litellm._logging import verbose_proxy_logger from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_str_from_messages, ) @@ -2447,13 +2448,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): data["litellm_proxy_rate_limit_response"] = response # Mirror into metadata so streaming success logging can find # it via ``kwargs["litellm_params"]["metadata"]``. - self._stash_value_in_metadata_channels( + self._stash_value_in_internal_metadata( data=data, key=RATE_LIMIT_RESPONSE_KEY, value=response, ) if parallel_slot_id is not None: - self._stash_value_in_metadata_channels( + self._stash_value_in_internal_metadata( data=data, key=MAX_PARALLEL_SLOT_ACQUIRED_KEY, value={ @@ -2533,7 +2534,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): requested_model=requested_model, ) else: - self._stash_value_in_metadata_channels( + self._stash_value_in_internal_metadata( data=data, key=RATE_LIMIT_DESCRIPTORS_KEY, value=descriptors, @@ -2566,7 +2567,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): data["litellm_proxy_rate_limit_response"] = tpm_response # Keep the metadata stash in sync when this is the # first snapshot written. - self._stash_value_in_metadata_channels( + self._stash_value_in_internal_metadata( data=data, key=RATE_LIMIT_RESPONSE_KEY, value=tpm_response, @@ -2803,19 +2804,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return merged @staticmethod - def _stash_value_in_metadata_channels( + def _stash_value_in_internal_metadata( data: Dict[str, Any], key: str, value: Any, ) -> None: - for channel in ("metadata", "litellm_metadata"): - existing = data.get(channel) - if isinstance(existing, dict): - existing[key] = value - elif channel == "metadata": - # ``litellm_metadata`` is owned by the router; don't conjure - # it here. - data[channel] = {key: value} + # Writes only the proxy-internal bucket. Routes that own + # ``litellm_metadata`` (Responses, /v1/messages, batches, files) expose + # ``metadata`` as a provider request parameter, so creating or adding to + # it here would forward internal state upstream. + _, metadata_bucket = get_or_create_metadata_bucket(data) + metadata_bucket[key] = value @classmethod def _stash_reservation_in_data( @@ -2831,11 +2830,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ scopes_payload: Optional[List[List[str]]] = [[k, v] for k, v in reserved_scopes] if reserved_scopes else None - cls._stash_value_in_metadata_channels(data=data, key=TPM_RESERVED_TOKENS_KEY, value=estimated_tokens) + cls._stash_value_in_internal_metadata(data=data, key=TPM_RESERVED_TOKENS_KEY, value=estimated_tokens) if reserved_model: - cls._stash_value_in_metadata_channels(data=data, key=TPM_RESERVED_MODEL_KEY, value=reserved_model) + cls._stash_value_in_internal_metadata(data=data, key=TPM_RESERVED_MODEL_KEY, value=reserved_model) if scopes_payload is not None: - cls._stash_value_in_metadata_channels(data=data, key=TPM_RESERVED_SCOPES_KEY, value=scopes_payload) + cls._stash_value_in_internal_metadata(data=data, key=TPM_RESERVED_SCOPES_KEY, value=scopes_payload) @staticmethod def _lookup_stashed_value( @@ -2858,9 +2857,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return candidate litellm_params = kwargs.get("litellm_params") if isinstance(litellm_params, dict): - lp_metadata = litellm_params.get("metadata") - if isinstance(lp_metadata, dict): - candidate = lp_metadata.get(key) + for channel in ("litellm_metadata", "metadata"): + lp_metadata = litellm_params.get(channel) + if isinstance(lp_metadata, dict) and lp_metadata.get(key) is not None: + return lp_metadata[key] if candidate is None and isinstance(standard_logging_metadata, dict): candidate = standard_logging_metadata.get(key) return candidate diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index d94fed0ee5b..673e73f72fb 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -150,6 +150,7 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS = ( "applied_guardrails", "applied_policies", "policy_sources", + "routing_decision", "pillar_response_headers", "_guardrail_pipelines", "_pipeline_managed_guardrails", @@ -197,6 +198,7 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS = ( "applied_guardrails", "applied_policies", "policy_sources", + "routing_decision", "standard_logging_object", "proxy_server_request", "secret_fields", @@ -1829,6 +1831,15 @@ async def add_litellm_data_to_request( return data +def _warn_stale_team_alias_once(warning_key: str, message: str, *args: str) -> None: + if warning_key in _STALE_TEAM_ALIAS_WARNING_KEYS: + return + _STALE_TEAM_ALIAS_WARNING_KEYS[warning_key] = None + while len(_STALE_TEAM_ALIAS_WARNING_KEYS) > _MAX_STALE_ALIAS_WARNING_KEYS: + _STALE_TEAM_ALIAS_WARNING_KEYS.popitem(last=False) + verbose_proxy_logger.warning(message, *args) + + def _update_model_if_team_alias_exists( data: dict, user_api_key_dict: UserAPIKeyAuth, @@ -1848,49 +1859,63 @@ def _update_model_if_team_alias_exists( Note: model_aliases for team models are deprecated. This function only applies to legacy non-team-scoped aliases. Team-scoped deployments use team_public_model_name and are resolved via map_team_model in route_llm_request. + + An alias that targets a team-scoped internal name (``model_name_{team_id}_{uuid}``) + with no live deployment behind it is never applied: the deployment was deleted, so + the rewrite could only fail with an error naming a model the caller never sent. + Keeping the requested model name lets it resolve against the deployments that still + exist (e.g. a gateway-level model group shared with the team). """ _model = data.get("model") - if _model and user_api_key_dict.team_model_aliases and _model in user_api_key_dict.team_model_aliases: - from litellm.proxy.proxy_server import llm_router + if not _model or not user_api_key_dict.team_model_aliases or _model not in user_api_key_dict.team_model_aliases: + return - # Skip alias rewrite if this model resolves to team-specific deployments - # (team models use team_public_model_name, not model_aliases) - aliased_target = user_api_key_dict.team_model_aliases[_model] + from litellm.proxy.proxy_server import llm_router - # Optional bypass for stale aliases from pre-PR deployments: - # only enabled via feature flag to preserve backwards compatibility. - # Cached at module level to avoid hot-path secret lookups on every request. - global _ENABLE_TEAM_STALE_ALIAS_BYPASS - if _ENABLE_TEAM_STALE_ALIAS_BYPASS is None: - _ENABLE_TEAM_STALE_ALIAS_BYPASS = get_secret_bool("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", False) - enable_stale_alias_bypass = _ENABLE_TEAM_STALE_ALIAS_BYPASS - # Check if the alias points to a team-scoped UUID name - # (format: "model_name_{team_id}_{uuid}") - is_stale_team_alias = aliased_target.startswith(f"model_name_{user_api_key_dict.team_id}_") - if is_stale_team_alias and llm_router: - # This is a stale alias from pre-PR deployments. - # Check if current team deployments exist for the public name. - key = (user_api_key_dict.team_id, _model) - if key in llm_router.team_model_to_deployment_indices: - if enable_stale_alias_bypass: - # Team deployments exist; skip stale alias - return - warning_key = f"{user_api_key_dict.team_id}:{_model}:{aliased_target}" - if warning_key not in _STALE_TEAM_ALIAS_WARNING_KEYS: - _STALE_TEAM_ALIAS_WARNING_KEYS[warning_key] = None - while len(_STALE_TEAM_ALIAS_WARNING_KEYS) > _MAX_STALE_ALIAS_WARNING_KEYS: - _STALE_TEAM_ALIAS_WARNING_KEYS.popitem(last=False) - verbose_proxy_logger.warning( - "Stale team model alias detected for model='%s', team_id='%s'. " - "New sibling deployments may be unreachable. " - "Set LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS=true to enable " - "team-scoped sibling routing.", - _sanitize_for_log(_model), - user_api_key_dict.team_id, - ) + # Skip alias rewrite if this model resolves to team-specific deployments + # (team models use team_public_model_name, not model_aliases) + aliased_target = user_api_key_dict.team_model_aliases[_model] - data["model"] = aliased_target - return + # Optional bypass for stale aliases from pre-PR deployments: + # only enabled via feature flag to preserve backwards compatibility. + # Cached at module level to avoid hot-path secret lookups on every request. + global _ENABLE_TEAM_STALE_ALIAS_BYPASS + if _ENABLE_TEAM_STALE_ALIAS_BYPASS is None: + _ENABLE_TEAM_STALE_ALIAS_BYPASS = get_secret_bool("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", False) + enable_stale_alias_bypass = _ENABLE_TEAM_STALE_ALIAS_BYPASS + # Check if the alias points to a team-scoped UUID name + # (format: "model_name_{team_id}_{uuid}") + is_stale_team_alias = aliased_target.startswith(f"model_name_{user_api_key_dict.team_id}_") + if is_stale_team_alias and llm_router: + if aliased_target not in llm_router.model_name_to_deployment_indices: + _warn_stale_team_alias_once( + f"deleted:{user_api_key_dict.team_id}:{_model}:{aliased_target}", + "Team model alias for model='%s', team_id='%s' targets '%s', which has no live " + "deployment. Routing with the requested model name instead; remove the stale " + "entry from the team's model_aliases to silence this warning.", + _sanitize_for_log(_model), + _sanitize_for_log(user_api_key_dict.team_id), + _sanitize_for_log(aliased_target), + ) + return + # This is a stale alias from pre-PR deployments. + # Check if current team deployments exist for the public name. + key = (user_api_key_dict.team_id, _model) + if key in llm_router.team_model_to_deployment_indices: + if enable_stale_alias_bypass: + # Team deployments exist; skip stale alias + return + _warn_stale_team_alias_once( + f"{user_api_key_dict.team_id}:{_model}:{aliased_target}", + "Stale team model alias detected for model='%s', team_id='%s'. " + "New sibling deployments may be unreachable. " + "Set LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS=true to enable " + "team-scoped sibling routing.", + _sanitize_for_log(_model), + _sanitize_for_log(user_api_key_dict.team_id), + ) + + data["model"] = aliased_target def _update_model_if_key_alias_exists( diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 9b756d14815..8a5a31710cf 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -233,27 +233,28 @@ def update_breakdown_metrics( ) # Update model group breakdown - if record.model_group and record.model_group not in breakdown.model_groups: - breakdown.model_groups[record.model_group] = MetricWithMetadata( + model_group_key = record.model_group or record.model + if model_group_key and model_group_key not in breakdown.model_groups: + breakdown.model_groups[model_group_key] = MetricWithMetadata( metrics=SpendMetrics(), - metadata=model_metadata.get(record.model_group, {}), + metadata=model_metadata.get(model_group_key, {}), ) - if record.model_group: - breakdown.model_groups[record.model_group].metrics = update_metrics( - breakdown.model_groups[record.model_group].metrics, record + if model_group_key: + breakdown.model_groups[model_group_key].metrics = update_metrics( + breakdown.model_groups[model_group_key].metrics, record ) # Update API key breakdown for this model - if record.api_key not in breakdown.model_groups[record.model_group].api_key_breakdown: - breakdown.model_groups[record.model_group].api_key_breakdown[record.api_key] = KeyMetricWithMetadata( + if record.api_key not in breakdown.model_groups[model_group_key].api_key_breakdown: + breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key] = KeyMetricWithMetadata( metrics=SpendMetrics(), metadata=KeyMetadata( key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None), team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None), ), ) - breakdown.model_groups[record.model_group].api_key_breakdown[record.api_key].metrics = update_metrics( - breakdown.model_groups[record.model_group].api_key_breakdown[record.api_key].metrics, + breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics = update_metrics( + breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics, record, ) @@ -574,11 +575,11 @@ def _build_aggregated_sql_query( date, api_key, model, - model_group, + COALESCE(NULLIF(model_group, ''), model) AS model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint, - GROUPING(date, api_key, model, model_group, + GROUPING(date, api_key, model, COALESCE(NULLIF(model_group, ''), model), custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level, SUM(spend)::float AS spend, @@ -599,8 +600,8 @@ def _build_aggregated_sql_query( (date, api_key), (date, model), (date, model, api_key), - (date, model_group), - (date, model_group, api_key), + (date, COALESCE(NULLIF(model_group, ''), model)), + (date, COALESCE(NULLIF(model_group, ''), model), api_key), (date, custom_llm_provider), (date, custom_llm_provider, api_key), (date, mcp_namespaced_tool_name), diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 8162babef40..877130c2066 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -216,7 +216,7 @@ async def _user_has_admin_privileges( teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_obj.teams}}) for team in teams: - team_obj = LiteLLM_TeamTable(**team.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): return True @@ -288,7 +288,7 @@ async def _team_admin_can_invite_user( for team in teams if _is_user_team_admin( user_api_key_dict=user_api_key_dict, - team_obj=LiteLLM_TeamTable(**team.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(team.model_dump()), ) ] if not admin_team_ids: diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index a2c16e88839..80d9ee21a44 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -45,6 +45,14 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( generate_key_helper_fn, prepare_metadata_fields, ) +from litellm.proxy.common_utils.user_api_key_cache import ( + object_permission_cache_key, + user_object_permission_id_cache_key, +) +from litellm.proxy.management_helpers.object_permission_utils import ( + _set_object_permission, + handle_update_object_permission_common, +) from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.proxy.utils import handle_exception_on_proxy, hash_password from litellm.repositories.organization_repository import OrganizationRepository @@ -401,7 +409,7 @@ async def new_user( - duration: Optional[str] - Duration for the key auto-created on `/user/new`. Default is None. - key_alias: Optional[str] - Alias for the key auto-created on `/user/new`. Default is None. - sso_user_id: Optional[str] - The id of the user in the SSO provider. - - object_permission: Optional[LiteLLM_ObjectPermissionBase] - internal user-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission. + - object_permission: Optional[LiteLLM_ObjectPermissionBase] - internal user-specific object permission. Example - {"vector_stores": ["vector_store_1"], "mcp_servers": ["github"], "mcp_tool_permissions": {"github": ["list_issues"]}}. The MCP grants act as a ceiling on every key this user holds. IF null or {} then no object permission. - prompts: Optional[List[str]] - List of allowed prompts for the user. If specified, the user will only be able to use these specific prompts. - organizations: List[str] - List of organization id's the user is a member of - budget_limits: Optional[list] - List of concurrent budget windows for the user. Each window specifies a budget_limit, time_period, and optional budget_duration. Example - [{"budget_limit": 10.0, "time_period": "1d"}, {"budget_limit": 50.0, "time_period": "7d"}]. @@ -466,6 +474,10 @@ async def new_user( data_json = data.json() # type: ignore data_json = _update_internal_new_user_params(data_json, data) + # Persist the requested grants as their own row and link it, mirroring key/team creation. + # generate_key_helper_fn only forwards object_permission_id, so without this the entitlement + # the caller sent would be dropped on the floor. + data_json = await _set_object_permission(data_json=data_json, prisma_client=prisma_client) _hash_password_in_dict(data_json) teams = data.teams if teams is None: @@ -513,7 +525,7 @@ async def new_user( response_dict["key"] = response.get("token", "") - new_user_response = NewUserResponse(**response_dict) + new_user_response = NewUserResponse.model_validate(response_dict) ######################################################### ########## USER CREATED HOOK ################ @@ -852,9 +864,12 @@ async def _check_user_info_v2_access( if prisma_client is None: return None - # Helper: fetch the target user row (reused across branches) + # Helper: fetch the target user row (reused across branches). object_permission is included so + # callers can read the user's MCP/vector-store entitlements without a second round trip. async def _fetch_target_user(): - return await UserRepository(prisma_client).table.find_unique(where={"user_id": target_user_id}) + return await UserRepository(prisma_client).table.find_unique( + where={"user_id": target_user_id}, include={"object_permission": True} + ) # Rule 1: Proxy admins — fetch and return the target row directly if _user_has_admin_view(user_api_key_dict): @@ -879,7 +894,7 @@ async def _check_user_info_v2_access( # Get all teams the caller belongs to teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": caller_user.teams}}) for team in teams: - team_obj = LiteLLM_TeamTable(**team.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): # Check if target user is in this team if team.team_id in (target_user.teams or []): @@ -972,6 +987,7 @@ async def user_info_v2( updated_at=user_data.get("updated_at"), sso_user_id=user_data.get("sso_user_id"), teams=user_data.get("teams") or [], + object_permission=user_data.get("object_permission"), ) except Exception as e: verbose_proxy_logger.exception( @@ -1013,11 +1029,11 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth): for key in _keys_in_db: if key.get("models") is None: key["models"] = [] - keys_in_db.append(LiteLLM_VerificationToken(**key)) + keys_in_db.append(LiteLLM_VerificationToken.model_validate(key)) # cast all teams to LiteLLM_TeamTable _teams_in_db: list = results[0]["teams"] or [] - _teams_in_db = [LiteLLM_TeamTable(**team) for team in _teams_in_db] + _teams_in_db = [LiteLLM_TeamTable.model_validate(team) for team in _teams_in_db] _teams_in_db.sort(key=lambda x: getattr(x, "team_alias", "") or "") returned_keys = _process_keys_for_user_info(keys=keys_in_db, all_teams=_teams_in_db) @@ -1146,7 +1162,7 @@ async def _schedule_user_update_audit_log( try: updated_user_row = await UserRepository(prisma_client).table.find_first(where={"user_id": response["user_id"]}) if updated_user_row: - user_row_typed = LiteLLM_UserTable(**updated_user_row.model_dump(exclude_none=True)) + user_row_typed = LiteLLM_UserTable.model_validate(updated_user_row.model_dump(exclude_none=True)) asyncio.create_task( UserManagementEventHooks.create_internal_user_audit_log( user_id=user_row_typed.user_id, @@ -1172,7 +1188,7 @@ def _check_user_update_authz( raise HTTPException(status_code=403, detail="Only proxy admins can modify user roles.") if existing_user_row is not None: - typed_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True)) + typed_row = LiteLLM_UserTable.model_validate(existing_user_row.model_dump(exclude_none=True)) if not can_user_call_user_update(user_api_key_dict=user_api_key_dict, user_info=typed_row): raise HTTPException( status_code=403, @@ -1207,6 +1223,48 @@ async def _invalidate_user_spend_counter_if_changed( await _invalidate_spend_counter(counter_key=f"spend:user:{non_default_values['user_id']}") +def _clears_object_permission(user_request: UpdateUserRequest) -> bool: + """Whether the caller explicitly asked to remove this user's object_permission. + + Distinguishes "sent nothing" from "sent an empty grant set". Only the latter clears; an omitted + field must leave an existing entitlement alone. + """ + if "object_permission" not in (user_request.fields_set() if hasattr(user_request, "fields_set") else set()): + return False + sent = user_request.object_permission + return sent is None or not sent.model_dump(exclude_unset=True, exclude_none=True) + + +async def _invalidate_cached_user_entitlement(user_id: str | None, object_permission_ids: tuple[str, ...]) -> None: + """Drop the cache entries an entitlement change makes stale. + + All three kinds are needed: a permission row is cached under its own id (so re-reading the same + link still yields the OLD grants), the ``user_id -> object_permission_id`` link is cached + separately (so a user who previously had NO entitlement keeps its "none" sentinel), and the user + row itself is cached whole. Leaving any behind means an admin revoking a tool keeps serving it + until the management-object TTL expires. + + Both the outgoing and incoming permission ids are passed, because a clear leaves no incoming id + at all and an upsert may mint a new row; invalidating only one of the two leaves the other's + grants live. + + Each deletion is isolated: one that fails must not skip the others, or a single unreachable key + would silently leave the rest of a revocation in place. Best-effort overall, exactly as the caches + are everywhere else, since one we cannot clear still expires on its own. + """ + from litellm.proxy.proxy_server import user_api_key_cache + + keys = ( + *(object_permission_cache_key(permission_id) for permission_id in dict.fromkeys(object_permission_ids)), + *((user_object_permission_id_cache_key(user_id), user_id) if user_id is not None else ()), + ) + for key in keys: + try: + await user_api_key_cache.async_delete_cache(key=key) + except Exception as e: # noqa: BLE001 # a cache we cannot clear still expires; never fail the write + verbose_proxy_logger.warning(f"Failed to invalidate cached entitlement key {key!r}: {str(e)}") + + async def _update_single_user_helper( user_request: UpdateUserRequest, user_api_key_dict: UserAPIKeyAuth, @@ -1248,7 +1306,7 @@ async def _update_single_user_helper( _check_user_update_authz(user_request, user_api_key_dict, existing_user_row) if existing_user_row is not None: - existing_user_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True)) + existing_user_row = LiteLLM_UserTable.model_validate(existing_user_row.model_dump(exclude_none=True)) # Prevent budget self-escalation (GHSA-wvg4-6222-3q4r): non-admin callers # must not be able to raise their own budget/spend fields. @@ -1259,9 +1317,15 @@ async def _update_single_user_helper( ) _is_self_update = _target_user_id is not None and user_api_key_dict.user_id == _target_user_id if _is_self_update and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: - _protected_fields = ("max_budget", "soft_budget", "spend") + # object_permission is a CEILING on what this human may reach, so a self-write is an + # escalation path: sending an empty grant list means "no restriction" and would lift a + # restriction an admin placed on them. Checked against the fields the caller actually SENT, + # because `_update_internal_user_params` drops empty values, and `object_permission: {}` is + # precisely the clear-my-own-ceiling case this must refuse. + _sent_fields = user_request.fields_set() if hasattr(user_request, "fields_set") else set() + _protected_fields = ("max_budget", "soft_budget", "spend", "object_permission") for _field in _protected_fields: - if _field in non_default_values: + if _field in non_default_values or _field in _sent_fields: raise HTTPException( status_code=403, detail={ @@ -1282,6 +1346,22 @@ async def _update_single_user_helper( # Reject NaN/±inf spend before it can reach the DB / spend counter. validate_finite_spend(non_default_values.get("spend")) + # Upsert the grants into their own row and link it, mirroring /key/update and /team/update. + # This also removes object_permission from the payload, which is not a column on the user table. + if "object_permission" in non_default_values: + object_permission_id = await handle_update_object_permission_common( + data_json=non_default_values, + existing_object_permission_id=getattr(existing_user_row, "object_permission_id", None), + prisma_client=prisma_client, + ) + if object_permission_id is not None: + non_default_values["object_permission_id"] = object_permission_id + elif _clears_object_permission(user_request): + # An explicit `{}` or null means "no object permission", which the merge-based upsert cannot + # express: merging an empty grant set over the existing row leaves every grant in place. So + # the link is dropped instead, which is what makes the documented clear actually clear. + non_default_values["object_permission_id"] = None + # Perform the update response: dict[str, Any] | None = None @@ -1326,6 +1406,19 @@ async def _update_single_user_helper( await _invalidate_user_spend_counter_if_changed(non_default_values) + if "object_permission_id" in non_default_values: + await _invalidate_cached_user_entitlement( + user_id=non_default_values.get("user_id"), + object_permission_ids=tuple( + permission_id + for permission_id in ( + getattr(existing_user_row, "object_permission_id", None), + non_default_values.get("object_permission_id"), + ) + if isinstance(permission_id, str) + ), + ) + if response is None: raise HTTPException( status_code=400, @@ -1407,7 +1500,7 @@ async def user_update( - team_id: Optional[str] - [DEPRECATED PARAM] The team id of the user. Default is None. - duration: Optional[str] - [NOT IMPLEMENTED]. - key_alias: Optional[str] - [NOT IMPLEMENTED]. - - object_permission: Optional[LiteLLM_ObjectPermissionBase] - internal user-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission. + - object_permission: Optional[LiteLLM_ObjectPermissionBase] - internal user-specific object permission. Example - {"vector_stores": ["vector_store_1"], "mcp_servers": ["github"], "mcp_tool_permissions": {"github": ["list_issues"]}}. The MCP grants act as a ceiling on every key this user holds. IF null or {} then no object permission. - prompts: Optional[List[str]] - List of allowed prompts for the user. If specified, the user will only be able to use these specific prompts. - budget_limits: Optional[list] - List of concurrent budget windows for the user. Each window specifies a budget_limit, time_period, and optional budget_duration. Example - [{"budget_limit": 10.0, "time_period": "1d"}, {"budget_limit": 50.0, "time_period": "7d"}]. @@ -1998,7 +2091,11 @@ async def get_users( for user in users: user_dump = user.model_dump() user_dump["metadata"] = _redact_scim_enterprise_metadata(user_dump.get("metadata")) - user_list.append(LiteLLM_UserTableWithKeyCount(**user_dump, key_count=user_key_counts.get(user.user_id, 0))) + user_list.append( + LiteLLM_UserTableWithKeyCount.model_validate( + {**user_dump, "key_count": user_key_counts.get(user.user_id, 0)} + ) + ) else: user_list = [] @@ -2157,7 +2254,7 @@ async def delete_user( teams_to_update = [] for team in fetch_all_teams: is_member_in_team, new_team_members = _cleanup_members_with_roles( - existing_team_row=LiteLLM_TeamTable(**team.model_dump()), + existing_team_row=LiteLLM_TeamTable.model_validate(team.model_dump()), data=TeamMemberDeleteRequest( team_id=team.team_id, user_id=user_row.user_id, @@ -2438,7 +2535,7 @@ async def ui_view_users( if not users: return [] - return [LiteLLM_UserTableFiltered(**user.model_dump()) for user in users] + return [LiteLLM_UserTableFiltered.model_validate(user.model_dump()) for user in users] except HTTPException: raise diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index ac6a2a4a7db..a94a75fdfa3 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1078,7 +1078,7 @@ async def _common_key_generation_helper( response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response - response = GenerateKeyResponse(**response) + response = GenerateKeyResponse.model_validate(response) response.token = response.token_id # remap token to use the hash, and leave the key in the `key` field [TODO]: clean up generate_key_helper_fn to do this @@ -2023,7 +2023,7 @@ def _validate_max_budget(max_budget: Optional[float]) -> None: async def _get_and_validate_existing_key( - token: str, prisma_client: Optional[PrismaClient] + token: str | None, prisma_client: Optional[PrismaClient], key_alias: str | None = None ) -> LiteLLM_VerificationToken: """ Get existing key from database and validate it exists. @@ -2031,12 +2031,13 @@ async def _get_and_validate_existing_key( Args: token: The key token to look up prisma_client: Prisma client instance + key_alias: Alias to look the key up by when token is not provided Returns: LiteLLM_VerificationToken: The existing key row Raises: - ProxyException: 404 if key is not found + ProxyException: 404 if key is not found, 400 if the alias matches multiple keys """ if prisma_client is None: raise HTTPException( @@ -2044,19 +2045,65 @@ async def _get_and_validate_existing_key( detail={"error": "Database not connected"}, ) - hashed_token = _hash_token_if_needed(token=token) + if token is not None: + hashed_token = _hash_token_if_needed(token=token) - existing_key_row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token}) + existing_key_row: LiteLLM_VerificationToken | None = await VerificationTokenRepository( + prisma_client + ).table.find_unique(where={"token": hashed_token}) - if existing_key_row is None: + if existing_key_row is None: + raise ProxyException( + message="Key not found.", + type=ProxyErrorTypes.not_found_error, + param="key", + code=status.HTTP_404_NOT_FOUND, + ) + + return existing_key_row + + if key_alias is None: + raise ProxyException( + message="either key or key_alias must be provided", + type=ProxyErrorTypes.bad_request_error, + param="key", + code=status.HTTP_400_BAD_REQUEST, + ) + + rows: list[LiteLLM_VerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many( + where={"key_alias": key_alias}, take=2 + ) + + if len(rows) == 0: + raise ProxyException( + message=f"Key not found. No key with key_alias='{key_alias}'.", + type=ProxyErrorTypes.not_found_error, + param="key_alias", + code=status.HTTP_404_NOT_FOUND, + ) + + if len(rows) > 1: + raise ProxyException( + message=f"Multiple keys share key_alias='{key_alias}', so it cannot be used as an identifier.", + type=ProxyErrorTypes.bad_request_error, + param="key_alias", + code=status.HTTP_400_BAD_REQUEST, + ) + + return rows[0] + + +def _resolve_token_to_update(data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken) -> str: + if data.key is not None: + return data.key + if existing_key_row.token is None: raise ProxyException( message="Key not found.", type=ProxyErrorTypes.not_found_error, param="key", code=status.HTTP_404_NOT_FOUND, ) - - return existing_key_row + return existing_key_row.token async def _process_single_key_update( @@ -2508,8 +2555,8 @@ async def update_key_fn( Update an existing API key's parameters. Parameters: - - key: str - The key to update - - key_alias: Optional[str] - User-friendly key alias + - key: Optional[str] - The key to update. Either key or key_alias must be provided. + - key_alias: Optional[str] - User-friendly key alias. If key is omitted, also identifies the key to update (must match exactly one key, same as /key/delete's key_aliases) - user_id: Optional[str] - User ID associated with key - team_id: Optional[str] - Team ID associated with key - agent_id: Optional[str] - The agent id associated with the key. @@ -2592,14 +2639,14 @@ async def update_key_fn( detail={"error": f"max_budget must be a non-negative finite number. Received: {data.max_budget}"}, ) - data_json: dict = data.model_dump(exclude_unset=True) - key = data_json.pop("key") - # get the row from db existing_key_row = await _get_and_validate_existing_key( token=data.key, prisma_client=prisma_client, + key_alias=data.key_alias, ) + key = _resolve_token_to_update(data=data, existing_key_row=existing_key_row) + data.key = key await _validate_update_key_data( data=data, @@ -3047,10 +3094,12 @@ async def bulk_update_team_keys( ) # team_id from validated scope, never user payload — drives _check_team_key_limits. - update_key_request = UpdateKeyRequest( - key=token, - team_id=data.team_id, - **update_field_dict, + update_key_request = UpdateKeyRequest.model_validate( + { + "key": token, + "team_id": data.team_id, + **update_field_dict, + } ) updated_key_info = await _process_single_key_update( update_key_request=update_key_request, @@ -4048,12 +4097,14 @@ def _transform_verification_tokens_to_deleted_records( records = [] for key in keys: key_payload = key.model_dump() - deleted_record = LiteLLM_DeletedVerificationToken( - **key_payload, - deleted_at=deleted_at, - deleted_by=user_api_key_dict.user_id, - deleted_by_api_key=user_api_key_dict.api_key, - litellm_changed_by=litellm_changed_by, + deleted_record = LiteLLM_DeletedVerificationToken.model_validate( + { + **key_payload, + "deleted_at": deleted_at, + "deleted_by": user_api_key_dict.user_id, + "deleted_by_api_key": user_api_key_dict.api_key, + "litellm_changed_by": litellm_changed_by, + } ) record = deleted_record.model_dump() @@ -4535,7 +4586,7 @@ async def _execute_virtual_key_regeneration( proxy_logging_obj=proxy_logging_obj, ) - response = GenerateKeyResponse(**updated_token_dict) + response = GenerateKeyResponse.model_validate(updated_token_dict) asyncio.create_task( KeyManagementEventHooks.async_key_rotated_hook( data=data, @@ -4853,7 +4904,7 @@ async def _check_proxy_or_team_admin_for_key( ) -def _validate_reset_spend_value(reset_to: Any, key_in_db: LiteLLM_VerificationToken) -> float: +def _validate_reset_spend_value(reset_to: object, key_in_db: LiteLLM_VerificationToken) -> float: if not isinstance(reset_to, (int, float)): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -5029,7 +5080,7 @@ async def validate_key_list_check( code=status.HTTP_403_FORBIDDEN, ) - complete_user_info = LiteLLM_UserTable(**complete_user_info_db_obj.model_dump()) + complete_user_info = LiteLLM_UserTable.model_validate(complete_user_info_db_obj.model_dump()) # internal user can only see their own keys if user_id: @@ -5102,7 +5153,7 @@ async def _fetch_user_team_objects( if teams is None: return [] - return [LiteLLM_TeamTable(**team.model_dump()) for team in teams] + return [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in teams] def _get_admin_team_ids_from_objects( @@ -5851,7 +5902,7 @@ async def _list_key_helper( if return_full_object is True or (expand and "user" in expand): if use_deleted_table: # Use deleted key type to preserve deleted_at, deleted_by, etc. - key_list.append(LiteLLM_DeletedVerificationToken(**key_dict)) + key_list.append(LiteLLM_DeletedVerificationToken.model_validate(key_dict)) else: key_list.append(UserAPIKeyAuth(**key_dict)) # Return full key object else: diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 282184d6495..64cc13a5543 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -459,7 +459,7 @@ if MCP_AVAILABLE: payload_dict: dict[str, Any] = loaded try: - return MCPServer(**payload_dict) + return MCPServer.model_validate(payload_dict) except Exception as e: verbose_proxy_logger.debug(f"Invalid temporary MCP server payload in Redis cache: {str(e)}") return None @@ -704,7 +704,7 @@ if MCP_AVAILABLE: except AttributeError: payload_dict = payload.dict() # type: ignore[attr-defined] payload_dict["credentials"] = inherited_credentials - return NewMCPServerRequest(**payload_dict) + return NewMCPServerRequest.model_validate(payload_dict) def _build_temporary_mcp_server_record( payload: NewMCPServerRequest, @@ -1526,7 +1526,6 @@ if MCP_AVAILABLE: temporary_server = await global_mcp_server_manager.build_mcp_server_from_table( temp_record, credentials_are_encrypted=False, - persist_discovered_endpoints=False, ) _cache_temporary_mcp_server( temporary_server, diff --git a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py index 1ca29eee89d..5e3ff8eb7f8 100644 --- a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py @@ -6,6 +6,7 @@ Endpoints here: """ import json +from collections.abc import Mapping, Sequence from typing import Any, Dict, List, Tuple from fastapi import APIRouter, Depends, HTTPException @@ -16,6 +17,9 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth # Clear cache and reload models to pick up the access group changes from litellm.proxy.management_endpoints.model_management_endpoints import ( + live_model_ids_snapshot, + model_info_as_mapping, + reload_serving_verdict, clear_cache, ) from litellm.proxy.utils import PrismaClient @@ -72,11 +76,92 @@ def add_access_group_to_deployment(model_info: Dict[str, Any], access_group: str return model_info, True +def _raise_http_if_reload_degraded_serving( + before: frozenset[str], + written_models: Sequence[tuple[str, object]], + access_group: str, +) -> None: + """Same verdict as the model-write endpoints, expressed through this file's + HTTPException error convention, with the metadata-only obligation: these writes + change group membership, not the models themselves, so a row that was already not + serving before the reload is never blamed here; only a model this reload stopped + serving is reported.""" + missing, collateral = reload_serving_verdict(before=before, written_models=written_models, written_must_serve=False) + gone = tuple(dict.fromkeys((*missing, *collateral))) + if not gone: + return + raise HTTPException( + status_code=500, + detail={ + "error": ( + f"Access group '{access_group}' was saved to the database, but model id(s) {list(gone)} that " + "this pod was serving are no longer live after the reload it triggered. Other pods reload on " + "their own interval. Check server logs for 'Error upserting deployment' for the cause." + ) + }, + ) + + +async def _tag_deployment_with_access_group( + model_id: str, + model_info: object, + access_group: str, + prisma_client: PrismaClient, +) -> tuple[str, Mapping[str, object]] | None: + """Write `access_group` into one deployment's model_info; returns the + (model_id, updated model_info) pair when a write happened, None when the + deployment already carried the group.""" + updated_model_info, was_modified = add_access_group_to_deployment( + model_info=dict(_readable_model_info_or_raise(model_id=model_id, model_info=model_info)), + access_group=access_group, + ) + if not was_modified: + return None + await ModelRepository(prisma_client).table.update( + where={"model_id": model_id}, + data={"model_info": json.dumps(updated_model_info)}, + ) + verbose_proxy_logger.debug(f"Updated deployment {model_id} with access group: {access_group}") + return (model_id, updated_model_info) + + +def _readable_model_info_or_raise(model_id: str, model_info: object) -> Mapping[str, object]: + """These helpers rewrite the model_info column wholesale, so a present-but-unreadable + value must refuse loudly rather than be silently replaced with a fresh object; an + absent value stays a legitimate empty start.""" + parsed = model_info_as_mapping(model_info) + if parsed is None and model_info is not None: + raise ValueError(f"model_info for deployment {model_id} is not a readable JSON object; refusing to rewrite it") + return parsed or {} + + +async def _strip_access_group_from_deployment( + model_id: str, + model_info: object, + access_group: str, + prisma_client: PrismaClient, +) -> tuple[str, Mapping[str, object]] | None: + """Remove `access_group` from one deployment's model_info; returns the + (model_id, updated model_info) pair when a write happened, None when the + deployment did not carry the group.""" + updated_model_info, was_modified = remove_access_group_from_deployment( + model_info=dict(_readable_model_info_or_raise(model_id=model_id, model_info=model_info)), + access_group=access_group, + ) + if not was_modified: + return None + await ModelRepository(prisma_client).table.update( + where={"model_id": model_id}, + data={"model_info": json.dumps(updated_model_info)}, + ) + return (model_id, updated_model_info) + + async def update_deployments_with_access_group( model_names: List[str], access_group: str, prisma_client: PrismaClient, -) -> int: +) -> tuple[tuple[str, Mapping[str, object]], ...]: """ Update all deployments for the given model names to include the access group. @@ -86,20 +171,15 @@ async def update_deployments_with_access_group( prisma_client: Database client Returns: - int: Number of deployments updated + The (model_id, updated model_info) pair of every deployment actually written, + so callers can verify each one survived the post-write reload """ - models_updated = 0 + deployments = await ModelRepository(prisma_client).table.find_many(where={"model_name": {"in": model_names}}) + verbose_proxy_logger.debug(f"Found {len(deployments)} deployments for model_names: {model_names}") + found_names = {deployment.model_name for deployment in deployments} for model_name in model_names: - verbose_proxy_logger.debug(f"Updating deployments for model_name: {model_name}") - - # Get all deployments with this model_name - deployments = await ModelRepository(prisma_client).table.find_many(where={"model_name": model_name}) - - verbose_proxy_logger.debug(f"Found {len(deployments)} deployments for model_name: {model_name}") - - # If no deployments found, this is a config model (not in DB) - if len(deployments) == 0: + if model_name not in found_names: raise HTTPException( status_code=400, detail={ @@ -107,65 +187,52 @@ async def update_deployments_with_access_group( }, ) - # Update each deployment - for deployment in deployments: - model_info = deployment.model_info or {} - - # Add access group using helper - updated_model_info, was_modified = add_access_group_to_deployment( - model_info=model_info, - access_group=access_group, - ) - - # Only update in DB if modified - if was_modified: - await ModelRepository(prisma_client).table.update( - where={"model_id": deployment.model_id}, - data={"model_info": json.dumps(updated_model_info)}, - ) - - models_updated += 1 - verbose_proxy_logger.debug( - f"Updated deployment {deployment.model_id} with access group: {access_group}" - ) - - return models_updated + tagged = [ + await _tag_deployment_with_access_group( + model_id=deployment.model_id, + model_info=deployment.model_info, + access_group=access_group, + prisma_client=prisma_client, + ) + for deployment in deployments + ] + return tuple(pair for pair in tagged if pair is not None) async def update_specific_deployments_with_access_group( model_ids: List[str], access_group: str, prisma_client: PrismaClient, -) -> int: +) -> tuple[tuple[str, Mapping[str, object]], ...]: """ Update specific deployments (by model_id) to include the access group. Unlike update_deployments_with_access_group which tags ALL deployments sharing a model_name, this function only tags the specific deployments identified by - their unique model_id. + their unique model_id. Returns the (model_id, updated model_info) pair of every + deployment actually written. """ - models_updated = 0 - for model_id in model_ids: - verbose_proxy_logger.debug(f"Updating specific deployment model_id: {model_id}") - deployment = await ModelRepository(prisma_client).table.find_unique(where={"model_id": model_id}) - if deployment is None: - raise HTTPException( - status_code=400, - detail={"error": f"Deployment with model_id '{model_id}' not found in Database."}, - ) - model_info = deployment.model_info or {} - updated_model_info, was_modified = add_access_group_to_deployment( - model_info=model_info, + verbose_proxy_logger.debug(f"Updating specific deployment model_ids: {model_ids}") + tagged = [ + await _tag_deployment_with_access_group( + model_id=model_id, + model_info=(await _find_deployment_or_400(model_id=model_id, prisma_client=prisma_client)), access_group=access_group, + prisma_client=prisma_client, ) - if was_modified: - await ModelRepository(prisma_client).table.update( - where={"model_id": model_id}, - data={"model_info": json.dumps(updated_model_info)}, - ) - models_updated += 1 - verbose_proxy_logger.debug(f"Updated deployment {model_id} with access group: {access_group}") - return models_updated + for model_id in model_ids + ] + return tuple(pair for pair in tagged if pair is not None) + + +async def _find_deployment_or_400(model_id: str, prisma_client: PrismaClient) -> Mapping[str, object] | None: + deployment = await ModelRepository(prisma_client).table.find_unique(where={"model_id": model_id}) + if deployment is None: + raise HTTPException( + status_code=400, + detail={"error": f"Deployment with model_id '{model_id}' not found in Database."}, + ) + return deployment.model_info def remove_access_group_from_deployment(model_info: Dict[str, Any], access_group: str) -> Tuple[Dict[str, Any], bool]: @@ -335,20 +402,28 @@ async def create_model_group( # Update deployments using the appropriate method if use_model_ids: assert data.model_ids is not None - models_updated = await update_specific_deployments_with_access_group( + updated_pairs = await update_specific_deployments_with_access_group( model_ids=data.model_ids, access_group=data.access_group, prisma_client=prisma_client, ) else: assert data.model_names is not None - models_updated = await update_deployments_with_access_group( + updated_pairs = await update_deployments_with_access_group( model_names=data.model_names, access_group=data.access_group, prisma_client=prisma_client, ) + models_updated = len(updated_pairs) + + live_before_reload = live_model_ids_snapshot() await clear_cache() + _raise_http_if_reload_degraded_serving( + before=live_before_reload, + written_models=updated_pairs, + access_group=data.access_group, + ) verbose_proxy_logger.info( f"Successfully created access group '{data.access_group}' with {models_updated} models updated" @@ -573,38 +648,42 @@ async def update_access_group( # Step 1: Remove access group from ALL DB deployments (skip config models) all_deployments = await ModelRepository(prisma_client).table.find_many() - for deployment in all_deployments: - model_info = deployment.model_info or {} - - updated_model_info, was_modified = remove_access_group_from_deployment( - model_info=model_info, + stripped = [ + await _strip_access_group_from_deployment( + model_id=deployment.model_id, + model_info=deployment.model_info, access_group=access_group, + prisma_client=prisma_client, ) - - if was_modified: - await ModelRepository(prisma_client).table.update( - where={"model_id": deployment.model_id}, - data={"model_info": json.dumps(updated_model_info)}, - ) + for deployment in all_deployments + ] + stripped_pairs = tuple(pair for pair in stripped if pair is not None) # Step 2: Add access group using the appropriate method if use_model_ids: assert data.model_ids is not None - models_updated = await update_specific_deployments_with_access_group( + updated_pairs = await update_specific_deployments_with_access_group( model_ids=data.model_ids, access_group=access_group, prisma_client=prisma_client, ) else: assert data.model_names is not None - models_updated = await update_deployments_with_access_group( + updated_pairs = await update_deployments_with_access_group( model_names=data.model_names, access_group=access_group, prisma_client=prisma_client, ) + models_updated = len(updated_pairs) # Clear cache and reload models to pick up the access group changes + live_before_reload = live_model_ids_snapshot() await clear_cache() + _raise_http_if_reload_degraded_serving( + before=live_before_reload, + written_models=list({**dict(stripped_pairs), **dict(updated_pairs)}.items()), + access_group=access_group, + ) verbose_proxy_logger.info( f"Successfully updated access group '{access_group}' with {models_updated} models updated" @@ -686,25 +765,27 @@ async def delete_access_group( try: # Remove access group from all DB deployments (skip config models) all_deployments = await ModelRepository(prisma_client).table.find_many() - models_updated = 0 - for deployment in all_deployments: - model_info = deployment.model_info or {} - - updated_model_info, was_modified = remove_access_group_from_deployment( - model_info=model_info, + removed = [ + await _strip_access_group_from_deployment( + model_id=deployment.model_id, + model_info=deployment.model_info, access_group=access_group, + prisma_client=prisma_client, ) - - if was_modified: - await ModelRepository(prisma_client).table.update( - where={"model_id": deployment.model_id}, - data={"model_info": json.dumps(updated_model_info)}, - ) - models_updated += 1 + for deployment in all_deployments + ] + removed_pairs = tuple(pair for pair in removed if pair is not None) + models_updated = len(removed_pairs) # Clear cache and reload models to pick up the access group changes + live_before_reload = live_model_ids_snapshot() await clear_cache() + _raise_http_if_reload_degraded_serving( + before=live_before_reload, + written_models=removed_pairs, + access_group=access_group, + ) verbose_proxy_logger.info( f"Successfully deleted access group '{access_group}' from {models_updated} deployments" diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 91f0b9b2790..b6422d7f5ae 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -13,6 +13,7 @@ model/{model_id}/update - PATCH endpoint for model update. import asyncio import datetime import json +from collections.abc import Mapping, Sequence from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast from fastapi import APIRouter, Depends, HTTPException, Header, Request, status @@ -52,13 +53,19 @@ from litellm.proxy.utils import PrismaClient from litellm.repositories.model_repository import ModelRepository from litellm.repositories.table_repositories import ModelTableRepository from litellm.repositories.team_repository import TeamRepository +from litellm.router import Router from litellm.types.proxy.management_endpoints.model_management_endpoints import ( UpdateUsefulLinksRequest, ) +from litellm.router_utils.auto_router_model_naming import ( + STRATEGY_ROUTER_PARAM_FIELDS, + validate_strategy_router_model_write, +) from litellm.types.router import ( SPECIAL_MODEL_INFO_PARAMS, Deployment, DeploymentTypedDict, + GenericLiteLLMParams, LiteLLMParamsTypedDict, updateDeployment, ) @@ -96,6 +103,45 @@ async def get_db_model(model_id: str, prisma_client: PrismaClient) -> Optional[D return deployment_pydantic_obj +def _strategy_router_write_violation( + incoming_params: GenericLiteLLMParams | None, + existing_params: GenericLiteLLMParams | None, +) -> str | None: + """Reject writes that would corrupt a strategy router's pseudo-model. + + An auto-router deployment's ``litellm_params.model`` (``auto_router/...``) is + the discriminator the router loads it by; a write that mangles it makes the + router drop the deployment silently under ``ignore_invalid_deployments``. + Only writes that supply ``litellm_params.model`` are judged, against the + merged (stored + incoming) params, so partial patches and restores of an + already-corrupted row stay legal. Returns the violation, or None. + """ + if incoming_params is None or incoming_params.model is None: + return None + present_fields = frozenset( + field + for field in STRATEGY_ROUTER_PARAM_FIELDS + for source in (incoming_params, existing_params) + if source is not None and getattr(source, field, None) is not None + ) + return validate_strategy_router_model_write(model=incoming_params.model, present_fields=present_fields) + + +def _raise_on_strategy_router_write_violation( + incoming_params: GenericLiteLLMParams | None, + existing_params: GenericLiteLLMParams | None, +) -> None: + violation = _strategy_router_write_violation(incoming_params=incoming_params, existing_params=existing_params) + if violation is None: + return + raise ProxyException( + message=violation, + type=ProxyErrorTypes.validation_error.value, + code=status.HTTP_400_BAD_REQUEST, + param="litellm_params.model", + ) + + def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> PrismaCompatibleUpdateDBModel: merged_deployment_dict = DeploymentTypedDict( model_name=db_model.model_name, @@ -253,6 +299,11 @@ async def patch_model( param="blocked", ) + _raise_on_strategy_router_write_violation( + incoming_params=patch_data.litellm_params, + existing_params=db_model.litellm_params, + ) + # Handle team model updates with proper alias management update_data = await _update_team_model_in_db( db_model=db_model, @@ -272,6 +323,7 @@ async def patch_model( ) # Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates) + live_before_reload = live_model_ids_snapshot() await clear_cache() ## CREATE AUDIT LOG ## @@ -288,6 +340,12 @@ async def patch_model( ) ) + raise_if_reload_degraded_serving( + before=live_before_reload, + written_models=[(model_id, getattr(updated_model, "model_info", None))], + action="update", + ) + return updated_model except Exception as e: @@ -370,6 +428,7 @@ async def _set_model_blocked_status( }, ) + live_before_reload = live_model_ids_snapshot() await clear_cache() asyncio.create_task( @@ -387,6 +446,12 @@ async def _set_model_blocked_status( ) ) + raise_if_reload_degraded_serving( + before=live_before_reload, + written_models=[(data.model_id, getattr(updated_model, "model_info", None))], + action=action, + ) + return updated_model except Exception as e: @@ -713,13 +778,8 @@ async def _get_team_deployments( # Confirm team_id in model_info (defensive check) result = [] for row in response: - model_info = row.model_info - if isinstance(model_info, str): - try: - model_info = json.loads(model_info) - except (TypeError, ValueError): - continue - if isinstance(model_info, dict) and model_info.get("team_id") == team_id: + model_info = model_info_as_mapping(row.model_info) + if model_info is not None and model_info.get("team_id") == team_id: result.append(row) return result @@ -770,13 +830,8 @@ async def _get_team_public_model_names( deployments = await _get_team_deployments(team_id, prisma_client) public_names: Set[str] = set() for row in deployments: - model_info = row.model_info - if isinstance(model_info, str): - try: - model_info = json.loads(model_info) - except (TypeError, ValueError): - continue - if isinstance(model_info, dict): + model_info = model_info_as_mapping(row.model_info) + if model_info is not None: public_name = model_info.get("team_public_model_name") if public_name: public_names.add(public_name) @@ -788,6 +843,7 @@ async def _remove_unbacked_team_models( prisma_client: PrismaClient, user_api_key_cache: Any, proxy_logging_obj: Any, + llm_router: Router | None = None, ) -> None: """ Strip a deleted team model's public name(s) from team.models and refresh the cache. @@ -795,26 +851,50 @@ async def _remove_unbacked_team_models( Must be called after the deployment row is deleted: a public name is removed only when no remaining team deployment still backs it, so a load-balanced replica isn't revoked while siblings serve it, and concurrent deletes can't leave a ghost. + + Legacy team models (created before team_public_model_name existed) store a + ``{public_name: "model_name_{team_id}_{uuid}"}`` entry in the team's model_aliases, + so the alias scan runs for every team model; skipping it for internal-shaped names + left stale aliases that rewrote requests to deployments that no longer exist. + Aliases are scrubbed only when the deleted deployment's name no longer resolves in + the router, so deleting one replica of a load-balanced group never breaks aliases + that still route to the surviving replicas (in any team). + + A public name that still resolves to a live router deployment (e.g. a gateway-level + model group shared with the team) is kept in team.models, so deleting a per-team + duplicate does not revoke the team's access to the shared deployment. """ team_id = model_params.model_info.team_id if team_id is None: return - # BYOK models carry an internal `model_name_{team_id}_{uuid}` name that can never - # be a team alias value, so skip the full litellm_modeltable scan for them. - removed_model_aliases: List[Tuple[str, str]] = [] - if not model_params.model_name.startswith(f"model_name_{team_id}_"): - removed_model_aliases = await delete_team_model_alias( + deleted_name_still_served = ( + llm_router is not None and model_params.model_name in llm_router.model_name_to_deployment_indices + ) + removed_model_aliases: List[Tuple[str, str]] = ( + [] + if deleted_name_still_served + else await delete_team_model_alias( public_model_name=model_params.model_name, prisma_client=prisma_client, ) - names_to_remove = {alias for alias_team_id, alias in removed_model_aliases if alias_team_id == team_id} - if model_params.model_info.team_public_model_name is not None: - names_to_remove.add(model_params.model_info.team_public_model_name) - - if names_to_remove: - names_to_remove -= await _get_team_public_model_names(team_id=team_id, prisma_client=prisma_client) + ) + removed_alias_names = {alias for alias_team_id, alias in removed_model_aliases if alias_team_id == team_id} + candidate_names = ( + removed_alias_names | {model_params.model_info.team_public_model_name} + if model_params.model_info.team_public_model_name is not None + else removed_alias_names + ) + if not candidate_names: + return + team_backed_names = await _get_team_public_model_names(team_id=team_id, prisma_client=prisma_client) + router_served_names = ( + frozenset(name for name in candidate_names if name in llm_router.model_name_to_deployment_indices) + if llm_router is not None + else frozenset() + ) + names_to_remove = candidate_names - team_backed_names - router_served_names if not names_to_remove: return @@ -853,18 +933,11 @@ async def _update_existing_team_model_assignment( def _get_team_public_model_name( model_info: Optional[Union[dict, str]], ) -> Optional[str]: - if isinstance(model_info, dict): - value = model_info.get("team_public_model_name") - return value if isinstance(value, str) else None - if isinstance(model_info, str): - try: - parsed = json.loads(model_info) - except (TypeError, ValueError): - return None - if isinstance(parsed, dict): - value = parsed.get("team_public_model_name") - return value if isinstance(value, str) else None - return None + parsed = model_info_as_mapping(model_info) + if parsed is None: + return None + value = parsed.get("team_public_model_name") + return value if isinstance(value, str) else None old_public_name = db_model.model_info.team_public_model_name if db_model.model_info else None @@ -978,7 +1051,7 @@ class ModelManagementAuthChecks: status_code=400, detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)}, ) - existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) + existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) ModelManagementAuthChecks.can_user_make_team_model_call( team_id=model_params.model_info.team_id, @@ -1016,7 +1089,7 @@ class ModelManagementAuthChecks: status_code=400, detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)}, ) - team_obj = LiteLLM_TeamTable(**team_obj_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_obj_row.model_dump()) return ModelManagementAuthChecks.can_user_make_team_model_call( team_id=model_params.model_info.team_id, @@ -1120,6 +1193,7 @@ async def delete_model( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, ) ## CREATE AUDIT LOG ## @@ -1267,6 +1341,11 @@ async def add_new_model( premium_user=premium_user, ) + _raise_on_strategy_router_write_violation( + incoming_params=model_params.litellm_params, + existing_params=None, + ) + model_response: Optional[LiteLLM_ProxyModelTable] = None # update DB if store_model_in_db is True: @@ -1275,6 +1354,7 @@ async def add_new_model( - store keys separately """ + live_before_reload = live_model_ids_snapshot() try: _original_litellm_model_name = model_params.model_name if model_params.model_info.team_id is None: @@ -1330,6 +1410,12 @@ async def add_new_model( ) ) + raise_if_reload_degraded_serving( + before=live_before_reload, + written_models=[(model_response.model_id, getattr(model_response, "model_info", None))], + action="create", + ) + return model_response except Exception as e: @@ -1414,6 +1500,11 @@ async def update_model( premium_user=premium_user, ) + _raise_on_strategy_router_write_violation( + incoming_params=model_params.litellm_params, + existing_params=deployment.litellm_params, + ) + # update DB if store_model_in_db is True: _existing_litellm_params_dict = dict(_existing_litellm_params.litellm_params) @@ -1450,8 +1541,8 @@ async def update_model( ) # Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates) + live_before_reload = live_model_ids_snapshot() await clear_cache() - ## CREATE AUDIT LOG ## asyncio.create_task( create_object_audit_log( @@ -1474,6 +1565,12 @@ async def update_model( ) ) + raise_if_reload_degraded_serving( + before=live_before_reload, + written_models=[(_model_id, getattr(model_response, "model_info", None))], + action="update", + ) + return model_response except Exception as e: verbose_proxy_logger.exception( @@ -1677,6 +1774,114 @@ def _deduplicate_litellm_router_models(models: List[Dict]) -> List[Dict]: return unique_models +def model_info_as_mapping(model_info: object) -> Mapping[str, object] | None: + """A DB row's model_info column arrives as a dict or as its JSON string depending on + the query path, and every consumer needs the mapping. Single owner of that parse: + returns None when no usable mapping exists (None, an unparseable string, or JSON + that is not an object), and callers choose what None means for them.""" + if isinstance(model_info, Mapping): + return model_info + if not isinstance(model_info, str): + return None + try: + parsed = json.loads(model_info) + except (TypeError, ValueError): + return None + return parsed if isinstance(parsed, Mapping) else None + + +def _expects_liveness_on_this_pod(model_info: object) -> bool: + from litellm.router import model_info_is_active_for_environment + + try: + return model_info_is_active_for_environment(model_info=model_info_as_mapping(model_info)) + except ValueError: + return True + + +def live_model_ids_snapshot() -> frozenset[str]: + """The ids this pod's router is currently serving, read fresh from the module global + because a reload can rebind it. The empirical ground truth every verdict below is + computed from; an absent router serves nothing.""" + from litellm.proxy.proxy_server import llm_router + + if llm_router is None: + return frozenset() + return frozenset(llm_router.get_model_ids()) + + +def reload_serving_verdict( + before: frozenset[str], + written_models: Sequence[tuple[str, object]], + written_must_serve: bool, +) -> tuple[tuple[str, ...], tuple[str, ...]]: + """Judge a write-triggered reload by diffing the router's serving state instead of + trusting any layer of the reload stack to report its own failure. + + The full cell matrix, per id: + - written, must-serve (the write's purpose is this model's serving state): live now + is fine; not live is reported unless the row is deliberately inactive for this + pod's LITELLM_ENVIRONMENT; a row whose model_info cannot be read counts as + expecting to serve, so its drop is still reported + - written, metadata-only (must_not_degrade): live before and gone now is reported; + a row that was already not serving stays silent, because its deadness predates + this write and blaming it would block unrelated metadata fixes + - not written but live before and gone now: collateral degradation of this pod + caused by the reload this request triggered (a wholesale re-add failure, or a + newly introduced conflict), always reported + + Returns (written ids violating their obligation, collateral ids no longer served). + Best effort under concurrent admin writes: the snapshot spans only this request. + """ + now = live_model_ids_snapshot() + written_ids = frozenset(model_id for model_id, _ in written_models) + if written_must_serve: + missing = tuple( + model_id + for model_id, model_info in written_models + if model_id not in now and _expects_liveness_on_this_pod(model_info) + ) + else: + missing = tuple(model_id for model_id, _ in written_models if model_id in before and model_id not in now) + collateral = tuple(sorted(before - now - written_ids)) + return (missing, collateral) + + +def raise_if_reload_degraded_serving( + before: frozenset[str], + written_models: Sequence[tuple[str, object]], + action: str, +) -> None: + """The caller-visible error this pod's model-write endpoints owe their caller when + the model they wrote is not being served after the reload they triggered. The DB + write is durable either way and every other pod reloads on its own interval; this + speaks only for the handling pod.""" + missing, collateral = reload_serving_verdict(before=before, written_models=written_models, written_must_serve=True) + if not missing and not collateral: + return + missing_clause = ( + f"the model id(s) {list(missing)} are not live in this pod's router after the reload and are not " + "being served by this pod." + if missing + else "the reload it triggered degraded this pod's serving state." + ) + collateral_clause = ( + f" Previously served model id(s) {list(collateral)} are also no longer being served by this pod." + if collateral + else "" + ) + raise ProxyException( + message=( + f"Model {action} was saved to the database, but {missing_clause}{collateral_clause} " + "Other pods reload on their own interval. Check server logs for 'Error upserting deployment' or " + "'Error creating deployment' for the cause." + ), + type=ProxyErrorTypes.internal_server_error, + code=status.HTTP_500_INTERNAL_SERVER_ERROR, + param=None, + ) + + async def clear_cache(): """ Clear router caches and reload models. diff --git a/litellm/proxy/management_endpoints/scim/scim_transformations.py b/litellm/proxy/management_endpoints/scim/scim_transformations.py index 65651752944..d572255fd32 100644 --- a/litellm/proxy/management_endpoints/scim/scim_transformations.py +++ b/litellm/proxy/management_endpoints/scim/scim_transformations.py @@ -175,6 +175,7 @@ class ScimTransformations: SCIMMember( value=ScimTransformations._get_scim_member_value(member), display=ScimTransformations._get_scim_member_display(member), + type="User", ) ) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 90eae5bbb21..e44299d018d 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -5,7 +5,9 @@ This is an enterprise feature and requires a premium license. """ import re -from typing import Any, Dict, Iterable, List, Optional, Set, Tuple +from collections.abc import Mapping, Sequence +from itertools import chain +from typing import Any, Dict, Iterable, List, NamedTuple, Optional, Set, Tuple from fastapi import ( APIRouter, @@ -17,8 +19,8 @@ from fastapi import ( Request, Response, ) -from pydantic import BaseModel, ValidationError -from typing_extensions import TypedDict +from pydantic import BaseModel, TypeAdapter, ValidationError +from typing_extensions import TypedDict, assert_never import litellm from litellm._logging import verbose_proxy_logger @@ -50,7 +52,11 @@ from litellm.proxy.management_endpoints.team_endpoints import ( team_member_add, team_member_delete, ) -from litellm.proxy.utils import _premium_user_check, handle_exception_on_proxy +from litellm.proxy.utils import ( + PrismaClient, + _premium_user_check, + handle_exception_on_proxy, +) from litellm.repositories.table_repositories import ( InvitationLinkRepository, OrganizationMembershipRepository, @@ -143,7 +149,11 @@ class ScimUserData(TypedDict): class GroupMemberExtractionResult(BaseModel): - """Result of extracting and processing group members.""" + """Result of extracting and processing group members. + + ``all_member_ids`` is deduped order-preserving; ``existing_member_ids`` is not, + so a repeated resolved id appears once in the former and twice in the latter. + """ existing_member_ids: List[str] created_users: List[NewUserResponse] @@ -371,6 +381,216 @@ async def _recompute_scim_member_roles(prisma_client: Any, user_ids: Iterable[st ) +class _ResolvedUserMember(NamedTuple): + user_id: str + + +class _SkippedGroupMember(NamedTuple): + value: str + reason: Literal["nested_group", "non_user_type", "existing_team"] + + +class _UnknownMember(NamedTuple): + value: str + + +_ClassifiedGroupMember = Union[_ResolvedUserMember, _SkippedGroupMember, _UnknownMember] + + +class _PartitionedMembers(NamedTuple): + resolved_ids: tuple[str, ...] + skipped: tuple[_SkippedGroupMember, ...] + unknown_ids: tuple[str, ...] + + +def _member_value(member: SCIMMember) -> str: + """A member id is opaque to us but has to be there; an empty one is a client error.""" + if not member.value or not member.value.strip(): + raise HTTPException( + status_code=400, + detail={"error": "Invalid member: user ID cannot be empty."}, + ) + return member.value + + +def _normalized_member_type(member: SCIMMember) -> str | None: + """The canonical ``type`` a member declares, lowercased; blank or absent means none.""" + normalized = (member.type or "").strip().lower() + return normalized or None + + +_JSON_OBJECT_ADAPTER = TypeAdapter(Dict[str, object]) + + +def _json_object_fields(raw: object) -> Mapping[str, object] | None: + """A typed, read-only view of a JSON object, or None when it is not one.""" + try: + return _JSON_OBJECT_ADAPTER.validate_python(raw) + except ValidationError: + return None + + +def _team_metadata_has_scim_provenance(team_metadata: object) -> bool: + """Whether a group write from the identity provider left its mark on this team. + + ``SCIM_TEAM_DATA_METADATA_KEY`` counts because PUT has been writing it since + long before the explicit marker, so a team the identity provider already + syncs is recognized without waiting to be written again. + """ + fields = _json_object_fields(team_metadata) + if fields is None: + return False + return bool(fields.get(SCIM_MANAGED_TEAM_METADATA_KEY)) or fields.get(SCIM_TEAM_DATA_METADATA_KEY) is not None + + +async def _classify_group_member(member: SCIMMember, prisma_client: PrismaClient) -> _ClassifiedGroupMember: + """ + Decide what a single SCIM group member refers to. + + A LiteLLM team only holds users, so a member is dropped when it declares a type + other than ``User`` or when its id names an existing team. Both of those checks + are placed around the user lookup rather than before it, because the id of a + real user is the one thing that outranks them: + + - ``"type": "Group"`` (what Entra sends for a nested group) is dropped without + a lookup. This bug provisioned nested group GUIDs as users, so those rows + exist in the wild and would otherwise resolve as members all over again. + - any other unrecognized type is dropped only after the user lookup misses. + Clients do send non-canonical types on real members (RFC 7643 defines + ``direct`` for ``User.groups``), and dropping a live user over one would + revoke that user's team access on the next full sync. + - an id that names an existing team is dropped only when the member arrives + untyped, which is how Okta sends nested groups, and only when that team is + one the identity provider writes. An id the IdP called a User is a user + even if some team happens to share the id, and a team created here rather + than through SCIM is not evidence of anything about the member. + """ + value = _member_value(member) + member_type = _normalized_member_type(member) + + if member_type == "group": + return _SkippedGroupMember(value=value, reason="nested_group") + + user = await UserRepository(prisma_client).table.find_unique(where={"user_id": value}) + if user is not None: + return _ResolvedUserMember(user_id=value) + + if member_type is not None and member_type != "user": + return _SkippedGroupMember(value=value, reason="non_user_type") + + if member_type is None: + team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": value}) + if team is not None and _team_metadata_has_scim_provenance(team.metadata): + return _SkippedGroupMember(value=value, reason="existing_team") + + return _UnknownMember(value=value) + + +def _bucketed_member(entry: _ClassifiedGroupMember) -> _PartitionedMembers: + """The single-member partition one classified entry contributes.""" + match entry: + case _ResolvedUserMember(user_id=user_id): + return _PartitionedMembers(resolved_ids=(user_id,), skipped=(), unknown_ids=()) + case _SkippedGroupMember(): + return _PartitionedMembers(resolved_ids=(), skipped=(entry,), unknown_ids=()) + case _UnknownMember(value=value): + return _PartitionedMembers(resolved_ids=(), skipped=(), unknown_ids=(value,)) + case _: + assert_never(entry) + + +def _partition_classified_members(classified: Iterable[_ClassifiedGroupMember]) -> _PartitionedMembers: + """Split classified members into the buckets the resolver acts on, keeping request order.""" + bucketed = tuple(_bucketed_member(entry) for entry in classified) + return _PartitionedMembers( + resolved_ids=tuple(chain.from_iterable(bucket.resolved_ids for bucket in bucketed)), + skipped=tuple(chain.from_iterable(bucket.skipped for bucket in bucketed)), + unknown_ids=tuple(chain.from_iterable(bucket.unknown_ids for bucket in bucketed)), + ) + + +def _admitted_member_id(entry: _ClassifiedGroupMember, created_ids: frozenset[str]) -> str | None: + match entry: + case _ResolvedUserMember(user_id=user_id): + return user_id + case _UnknownMember(value=value): + return value if value in created_ids else None + case _SkippedGroupMember(): + return None + case _: + assert_never(entry) + + +def _admitted_member_ids(classified: Iterable[_ClassifiedGroupMember], created_ids: frozenset[str]) -> tuple[str, ...]: + """Member ids that survive resolution, in the order the request listed them. + + An id the request repeats is one member: the roster these ids are written to + holds one row per member, and a second creation attempt for the same id fails + against the real unique constraint even though the first one succeeded. + """ + return tuple( + dict.fromkeys( + member_id for entry in classified if (member_id := _admitted_member_id(entry, created_ids)) is not None + ) + ) + + +async def _resolve_group_member_ids( + members: Sequence[SCIMMember], + created_via: str, + prisma_client: PrismaClient, +) -> GroupMemberExtractionResult: + """ + Resolve SCIM group members to LiteLLM user ids, dropping members that are not users. + + Only the operations that put ids onto a roster resolve their members: an id + that resolves to nothing is created when litellm_settings.scim_upsert_user is + True (default) and rejected per SCIM 2.0 otherwise. Removals do not come + through here; dropping an id is idempotent, so it needs neither a lookup nor a + user to drop. + + Raises: + HTTPException: 400 when a member id is empty, or when scim_upsert_user is + False and a member id is neither an existing user, an existing team, nor a + member declared to be something other than a user. + """ + classified = tuple([await _classify_group_member(member, prisma_client) for member in members]) + partition = _partition_classified_members(classified) + + for skipped in partition.skipped: + verbose_proxy_logger.info( + "SCIM: ignoring non-user group member '%s' (%s); LiteLLM teams contain users only", + skipped.value, + skipped.reason, + ) + + if partition.unknown_ids and not await _get_scim_upsert_user_setting(): + raise HTTPException( + status_code=400, + detail={ + "error": f"User with ID '{partition.unknown_ids[0]}' does not exist. " + "Please create the user first via POST /Users before adding to group." + }, + ) + + creations = tuple( + [ + (user_id, await _create_user_if_not_exists(user_id=user_id, created_via=created_via)) + for user_id in partition.unknown_ids + ] + ) + created_users = tuple(created for _, created in creations if created is not None) + + return GroupMemberExtractionResult( + existing_member_ids=partition.resolved_ids, + created_users=created_users, + all_member_ids=_admitted_member_ids( + classified, + frozenset(user_id for user_id, created in creations if created is not None), + ), + ) + + async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionResult: """ Extract member IDs from SCIMGroup, validating that all users exist. @@ -386,56 +606,10 @@ async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionRe HTTPException: If scim_upsert_user is False and any member user does not exist (400 Bad Request) """ prisma_client = await _get_prisma_client_or_raise_exception() - existing_member_ids = [] - created_users = [] - all_member_ids = [] - - # Check the feature flag - scim_upsert_user = await _get_scim_upsert_user_setting() - - if group.members: - for member in group.members: - user_id = member.value - - # Validate user_id is not empty - if not user_id or not user_id.strip(): - raise HTTPException( - status_code=400, - detail={"error": "Invalid member: user ID cannot be empty."}, - ) - - # Check if user exists - user = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) - - if user: - existing_member_ids.append(user_id) - all_member_ids.append(user_id) - else: - if scim_upsert_user: - # Create the user if they don't exist (backward compatible behavior) - created_user = await _create_user_if_not_exists( - user_id=user_id, created_via="scim_group_membership" - ) - if created_user: - created_users.append(created_user) - all_member_ids.append(user_id) - # If creation failed, user is skipped (logged in helper) - else: - # User doesn't exist - reject per SCIM 2.0 protocol - # This prevents security issues where users not assigned to app - # get provisioned via group membership - raise HTTPException( - status_code=400, - detail={ - "error": f"User with ID '{user_id}' does not exist. " - "Please create the user first via POST /Users before adding to group." - }, - ) - - return GroupMemberExtractionResult( - existing_member_ids=existing_member_ids, - created_users=created_users, - all_member_ids=all_member_ids, + return await _resolve_group_member_ids( + members=group.members or [], + created_via="scim_group_membership", + prisma_client=prisma_client, ) @@ -448,7 +622,7 @@ async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]: user = await UserRepository(prisma_client).table.find_unique(where={"user_id": member_id}) if user: display_name = user.user_email or user.user_id - members.append(SCIMMember(value=user.user_id, display=display_name)) + members.append(SCIMMember(value=user.user_id, display=display_name, type="User")) return members @@ -863,6 +1037,14 @@ def _get_schemas() -> list: type="string", description="Member display name.", ), + SCIMSchemaAttribute( + name="type", + type="string", + description=( + 'The type of member; canonical values are "User" and "Group". ' + "Only members of type User are honored, LiteLLM teams contain users only." + ), + ), ], ), ], @@ -1317,7 +1499,7 @@ async def delete_user( where={"team_id": team.team_id}, data={"members": new_members} ) - team_row = LiteLLM_TeamTable(**team.model_dump()) + team_row = LiteLLM_TeamTable.model_validate(team.model_dump()) if any(member.user_id == user_id for member in team_row.members_with_roles or []): await team_member_delete( data=TeamMemberDeleteRequest(team_id=team_row.team_id, user_id=user_id), @@ -1336,21 +1518,42 @@ async def delete_user( raise handle_exception_on_proxy(e) -def _extract_group_values(value: Any) -> List[str]: +def _parse_member_entry(entry: object) -> SCIMMember | None: + """Parse one entry of a SCIM patch value, or None when it carries no id.""" + if isinstance(entry, str): + return SCIMMember(value=entry) + + fields = _json_object_fields(entry) + if fields is None: + return None + + entry_value = fields.get("value") + if not entry_value: + return None + + entry_display = fields.get("display") + entry_type = fields.get("type") + return SCIMMember( + value=str(entry_value), + display=str(entry_display) if entry_display is not None else None, + type=entry_type if isinstance(entry_type, str) else None, + ) + + +def _parse_member_entries(value: object) -> tuple[SCIMMember, ...]: + """Parse a SCIM patch value into members, keeping each entry's ``type``. + + PATCH bodies bypass SCIMGroup parsing (SCIMPatchOperation.value is untyped), + so member objects arrive as raw dicts and the ``type`` that marks a nested + group would otherwise be lost. + """ + entries: tuple[object, ...] = tuple(value) if isinstance(value, list) else (value,) + return tuple(member for member in (_parse_member_entry(entry) for entry in entries) if member is not None) + + +def _extract_group_values(value: object) -> List[str]: """Return group ids from a SCIM patch value.""" - group_values: List[str] = [] - if isinstance(value, list): - for v in value: - if isinstance(v, dict) and v.get("value"): - group_values.append(str(v.get("value"))) - elif isinstance(v, str): - group_values.append(v) - elif isinstance(value, dict): - if value.get("value"): - group_values.append(str(value.get("value"))) - elif isinstance(value, str): - group_values.append(value) - return group_values + return [member.value for member in _parse_member_entries(value)] def _extract_ids_from_path_filter(path: str | None, attribute: str) -> List[str]: @@ -1833,6 +2036,7 @@ async def create_group( team_id=team_id, team_alias=group.displayName, members_with_roles=members_with_roles, + metadata={SCIM_MANAGED_TEAM_METADATA_KEY: True}, ), http_request=Request(scope={"type": "http", "path": "/scim/v2/Groups"}), user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), @@ -1875,7 +2079,11 @@ async def update_group( # Prepare update data existing_metadata = existing_team.metadata if existing_team.metadata else {} - updated_metadata = {**existing_metadata, "scim_data": group.model_dump()} + updated_metadata = { + **existing_metadata, + SCIM_TEAM_DATA_METADATA_KEY: group.model_dump(), + SCIM_MANAGED_TEAM_METADATA_KEY: True, + } update_data = { "team_alias": group.displayName, @@ -1968,12 +2176,17 @@ async def _process_group_patch_operations( is absolute: it declares the roster is exactly this set, so the caller must reconcile against it as a set-to-target rather than rebasing it onto a concurrently-mutated roster. + + A ``remove`` drops the ids it names without resolving them first. Removal is + idempotent and cannot put anything on a roster, while resolving would make it + conditional on what the id turns out to be and leave members we should never + have admitted - the phantom users this endpoint used to create for nested + groups - impossible to clean up. """ update_data: Dict[str, Any] = {} # Create a fresh copy of existing metadata to avoid Prisma issues - existing_metadata = existing_team.metadata or {} - metadata = dict(existing_metadata) if existing_metadata else {} + metadata = {**(existing_team.metadata or {}), SCIM_MANAGED_TEAM_METADATA_KEY: True} # Track member changes. members_with_roles is the source of truth for team # membership; the legacy `members` column is not populated by team creation @@ -2001,50 +2214,26 @@ async def _process_group_patch_operations( metadata["externalId"] = str(value) elif path.startswith("members"): # Handle member operations - member_values = _extract_group_values(value) - if not member_values and value is None: - member_values = _extract_ids_from_path_filter(op.path, "members") - # Check the feature flag - scim_upsert_user = await _get_scim_upsert_user_setting() - # Validate all users exist or create them based on feature flag - valid_members = [] - for member_id in member_values: - # Validate member_id is not empty - if not member_id or not member_id.strip(): - raise HTTPException( - status_code=400, - detail={"error": "Invalid member: user ID cannot be empty."}, - ) + patched_members = ( + _parse_member_entries(value) + if value is not None + else tuple( + SCIMMember(value=member_id) for member_id in _extract_ids_from_path_filter(op.path, "members") + ) + ) - user = await UserRepository(prisma_client).table.find_unique(where={"user_id": member_id}) - if user: - valid_members.append(member_id) - else: - if scim_upsert_user: - # Create the user if they don't exist (backward compatible behavior) - created_user = await _create_user_if_not_exists( - user_id=member_id, created_via="scim_group_patch" - ) - if created_user: - valid_members.append(member_id) - # If creation failed, user is skipped (logged in helper) - else: - # User doesn't exist - reject per SCIM 2.0 protocol - raise HTTPException( - status_code=400, - detail={ - "error": f"User with ID '{member_id}' does not exist. " - "Please create the user first via POST /Users before adding to group." - }, - ) - - if op_type == "replace": - final_members = set(valid_members) - elif op_type == "add": - final_members.update(valid_members) - elif op_type == "remove": - for member_id in valid_members: - final_members.discard(member_id) + if op_type == "remove": + final_members = final_members - {_member_value(member) for member in patched_members} + else: + member_result = await _resolve_group_member_ids( + members=patched_members, + created_via="scim_group_patch", + prisma_client=prisma_client, + ) + if op_type == "replace": + final_members = set(member_result.all_member_ids) + elif op_type == "add": + final_members = final_members | set(member_result.all_member_ids) else: # Handle other generic metadata if op_type == "remove": @@ -2052,9 +2241,7 @@ async def _process_group_patch_operations( else: metadata[path] = value - # Include metadata in update data if it exists - if metadata: - update_data["metadata"] = metadata + update_data["metadata"] = metadata member_replace_present = any( op.op == "replace" and (op.path or "").lower().startswith("members") for op in patch_ops.Operations @@ -2145,7 +2332,9 @@ async def patch_group( refreshed_team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": group_id}) refreshed_current = ( - set(await _get_team_member_user_ids_from_team(LiteLLM_TeamTable(**refreshed_team.model_dump()))) + set( + await _get_team_member_user_ids_from_team(LiteLLM_TeamTable.model_validate(refreshed_team.model_dump())) + ) if refreshed_team else snapshot_members ) @@ -2173,7 +2362,7 @@ async def patch_group( # Convert to SCIM format and return scim_group = await ScimTransformations.transform_litellm_team_to_scim_group( - LiteLLM_TeamTable(**updated_team.model_dump()) + LiteLLM_TeamTable.model_validate(updated_team.model_dump()) ) return scim_group diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 59b0cbc4ae7..c35c17aa359 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -141,7 +141,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import ( router = APIRouter() -def _sanitize_for_log(value: Any) -> str: +def _sanitize_for_log(value: object) -> str: """Strip CR/LF from user-controlled values to prevent log injection.""" try: text = str(value) @@ -171,7 +171,7 @@ async def _refresh_cached_team( """ await _cache_team_object( team_id=team_row.team_id, - team_table=LiteLLM_TeamTableCachedObj(**team_row.model_dump()), + team_table=LiteLLM_TeamTableCachedObj.model_validate(team_row.model_dump()), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -510,7 +510,7 @@ async def get_all_team_memberships( returned_tm: List[LiteLLM_TeamMembership] = [] for tm in team_memberships: - returned_tm.append(LiteLLM_TeamMembership(**tm.model_dump())) + returned_tm.append(LiteLLM_TeamMembership.model_validate(tm.model_dump())) return returned_tm @@ -772,7 +772,7 @@ async def _check_org_team_limits( # Convert teams to LiteLLM_TeamTable objects team_objs: List[LiteLLM_TeamTable] = [] for team in teams: - team_objs.append(LiteLLM_TeamTable(**team.model_dump())) + team_objs.append(LiteLLM_TeamTable.model_validate(team.model_dump())) check_org_team_model_specific_limits( teams=team_objs, @@ -1467,9 +1467,9 @@ async def fetch_and_validate_organization( ) is_proxy_admin = user_api_key_dict is not None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN - organization = LiteLLM_OrganizationTableWithMembers(**organization_row.model_dump()) + organization = LiteLLM_OrganizationTableWithMembers.model_validate(organization_row.model_dump()) validate_team_org_change( - team=LiteLLM_TeamTable(**existing_team_row.model_dump()), + team=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()), organization=organization, llm_router=llm_router, is_proxy_admin=is_proxy_admin, @@ -1477,7 +1477,7 @@ async def fetch_and_validate_organization( if is_proxy_admin: await _auto_add_team_members_to_organization( - team=LiteLLM_TeamTable(**existing_team_row.model_dump()), + team=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()), organization=organization, prisma_client=prisma_client, ) @@ -1714,7 +1714,7 @@ async def update_team( # Verify caller has access to manage this team await _verify_team_access( - team_obj=LiteLLM_TeamTable(**existing_team_row.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()), user_api_key_dict=user_api_key_dict, ) @@ -2013,7 +2013,7 @@ async def patch_team( existing_metadata = existing_team_row.metadata if isinstance(existing_team_row.metadata, dict) else {} patch_fields["metadata"] = apply_json_merge_patch(existing_metadata, patch_fields["metadata"]) - update_request = UpdateTeamRequest(team_id=team_id, **patch_fields) + update_request = UpdateTeamRequest.model_validate({"team_id": team_id, **patch_fields}) result = await update_team( data=update_request, @@ -2591,7 +2591,7 @@ async def team_member_add( detail={"error": f"Team not found for team_id={getattr(data, 'team_id', None)}"}, ) - complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()) team_member_add_duplication_check( data=data, @@ -2636,10 +2636,12 @@ async def team_member_add( _emit_team_members_metric(complete_team_data) - return TeamAddMemberResponse( - **updated_team.model_dump(), - updated_users=updated_users, - updated_team_memberships=updated_team_memberships, + return TeamAddMemberResponse.model_validate( + { + **updated_team.model_dump(), + "updated_users": updated_users, + "updated_team_memberships": updated_team_memberships, + } ) @@ -2711,7 +2713,7 @@ async def team_member_delete( status_code=400, detail={"error": "Team id={} does not exist in db".format(data.team_id)}, ) - existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) + existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) ## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN @@ -2915,7 +2917,7 @@ async def team_member_update( status_code=400, detail={"error": "Team id={} does not exist in db".format(data.team_id)}, ) - existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) + existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) ## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN @@ -3261,7 +3263,7 @@ async def delete_team( status_code=404, detail={"error": f"Team not found, passed team_id={team_id}"}, ) - team_row_pydantic = LiteLLM_TeamTable(**team_row_base.model_dump()) + team_row_pydantic = LiteLLM_TeamTable.model_validate(team_row_base.model_dump()) # Verify caller has access to manage this team await _verify_team_access( @@ -3385,12 +3387,14 @@ def _transform_teams_to_deleted_records( records = [] for team in teams: team_payload = team.model_dump() - deleted_record = LiteLLM_DeletedTeamTable( - **team_payload, - deleted_at=deleted_at, - deleted_by=user_api_key_dict.user_id, - deleted_by_api_key=user_api_key_dict.api_key, - litellm_changed_by=litellm_changed_by, + deleted_record = LiteLLM_DeletedTeamTable.model_validate( + { + **team_payload, + "deleted_at": deleted_at, + "deleted_by": user_api_key_dict.user_id, + "deleted_by_api_key": user_api_key_dict.api_key, + "litellm_changed_by": litellm_changed_by, + } ) record = deleted_record.model_dump() @@ -3580,7 +3584,7 @@ async def team_info( ) await validate_membership( user_api_key_dict=user_api_key_dict, - team_table=LiteLLM_TeamTable(**team_info.model_dump()), + team_table=LiteLLM_TeamTable.model_validate(team_info.model_dump()), ) ## GET ALL KEYS ## @@ -3615,9 +3619,9 @@ async def team_info( returned_tm = await get_all_team_memberships(prisma_client, [team_id], user_id=None) if isinstance(team_info, dict): - _team_info = TeamInfoResponseObjectTeamTable(**team_info) + _team_info = TeamInfoResponseObjectTeamTable.model_validate(team_info) elif isinstance(team_info, BaseModel): - _team_info = TeamInfoResponseObjectTeamTable(**team_info.model_dump()) + _team_info = TeamInfoResponseObjectTeamTable.model_validate(team_info.model_dump()) else: _team_info = TeamInfoResponseObjectTeamTable() @@ -3823,7 +3827,7 @@ async def block_team( # Verify caller has access to manage this team await _verify_team_access( - team_obj=LiteLLM_TeamTable(**existing_team.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(existing_team.model_dump()), user_api_key_dict=user_api_key_dict, ) @@ -3872,7 +3876,7 @@ async def unblock_team( # Verify caller has access to manage this team await _verify_team_access( - team_obj=LiteLLM_TeamTable(**existing_team.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(existing_team.model_dump()), user_api_key_dict=user_api_key_dict, ) @@ -3916,13 +3920,13 @@ async def list_available_teams( status_code=404, detail={"error": "User not found"}, ) - user_info_correct_type = LiteLLM_UserTable(**user_info.model_dump()) + user_info_correct_type = LiteLLM_UserTable.model_validate(user_info.model_dump()) available_teams = [team for team in available_teams if team not in user_info_correct_type.teams] available_teams_db = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": available_teams}}) - available_teams_correct_type = [LiteLLM_TeamTable(**team.model_dump()) for team in available_teams_db] + available_teams_correct_type = [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in available_teams_db] return available_teams_correct_type @@ -4090,7 +4094,7 @@ def _convert_teams_to_response_models( team_dict = team.dict() if use_deleted_table: - team_list.append(LiteLLM_DeletedTeamTable(**team_dict)) + team_list.append(LiteLLM_DeletedTeamTable.model_validate(team_dict)) else: members_with_roles = team_dict.get("members_with_roles") if not isinstance(members_with_roles, list): @@ -4705,7 +4709,7 @@ async def team_model_add( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) # Authorization check - only proxy admin, team admin, or org admin can add models if ( @@ -4805,7 +4809,7 @@ async def team_model_delete( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) # Authorization check - only proxy admin, team admin, or org admin can remove models if ( @@ -4873,7 +4877,7 @@ async def team_member_permissions( check_db_only=True, ) - complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()) # Admin Viewer follows the read-parity rule: see team permissions like # a Proxy Admin would. Team / org admins keep their existing scope. @@ -4940,7 +4944,7 @@ async def update_team_member_permissions( check_db_only=True, ) - complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()) # Available-team self-join must NOT grant write access to team-wide # permission policies; only proxy/team/org admins can update them. @@ -5201,7 +5205,7 @@ async def get_team_daily_activity( if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases: has_full_team_view = True for team_alias in team_aliases: - team_obj = LiteLLM_TeamTable(**team_alias.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_alias.model_dump()) is_admin = _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj) has_perm = _team_member_has_permission( user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index 7c52b04c4eb..14ba1dbfde1 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -308,7 +308,7 @@ async def add_new_member( ) await _append_team_id_if_absent(prisma_client, new_member.user_id, team_id) if _returned_user is not None: - returned_user = LiteLLM_UserTable(**_returned_user.model_dump()) + returned_user = LiteLLM_UserTable.model_validate(_returned_user.model_dump()) elif new_member.user_email is not None: new_user_defaults = get_new_internal_user_defaults(user_id=str(uuid.uuid4()), user_email=new_member.user_email) ## user email is not unique acc. to prisma schema -> future improvement @@ -323,11 +323,11 @@ async def add_new_member( _returned_user = await prisma_client.insert_data(data=new_user_defaults, table_name="user") # type: ignore if _returned_user is not None: - returned_user = LiteLLM_UserTable(**_returned_user.model_dump()) + returned_user = LiteLLM_UserTable.model_validate(_returned_user.model_dump()) elif len(existing_user_row) == 1: user_info = existing_user_row[0] await _append_team_id_if_absent(prisma_client, user_info.user_id, team_id) - returned_user = LiteLLM_UserTable(**user_info.model_dump()) + returned_user = LiteLLM_UserTable.model_validate(user_info.model_dump()) elif len(existing_user_row) > 1: raise HTTPException( status_code=400, @@ -354,7 +354,7 @@ async def add_new_member( include={"litellm_budget_table": True}, ) - returned_team_membership = LiteLLM_TeamMembership(**_returned_team_membership.model_dump()) + returned_team_membership = LiteLLM_TeamMembership.model_validate(_returned_team_membership.model_dump()) if returned_user is None: raise Exception("Unable to update user table with membership information!") diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index efd1d6b6cee..b2e36188681 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -14,6 +14,7 @@ from litellm.types.utils import SpecialEnums if TYPE_CHECKING: from fastapi import Request + from litellm.proxy._types import UserAPIKeyAuth from litellm.router import Router @@ -294,9 +295,8 @@ def get_credentials_for_model( def get_team_provider_credentials( llm_router: Optional["Router"], - team_models: List[str], + user_api_key_dict: "UserAPIKeyAuth", custom_llm_provider: str, - team_id: Optional[str] = None, ) -> Optional[dict]: """ Resolve upstream credentials for a provider-scoped file operation @@ -304,21 +304,61 @@ def get_team_provider_credentials( Priority: 1. The team's own (BYOK) deployment for this provider — a deployment whose - ``model_info.team_id`` matches ``team_id``. This keeps team-scoped listings - on the team's own provider account/key instead of a shared global one. - 2. Fallback: any deployment the team is granted access to for this provider, - expanding wildcard routes and the all-proxy-models sentinel. + ``model_info.team_id`` matches the caller's team. This keeps team-scoped + listings on the team's own provider account/key instead of a shared + global one. + 2. Fallback: any deployment the caller is granted access to for this + provider, expanding wildcard routes and the all-proxy-models sentinel. - Credential lookup is always scoped to the team's allowlist, so a team can - never resolve a provider key for a deployment it isn't authorized to use. + Credential lookup is scoped to both the team's allowlist and the key's own + model allowlist (``user_api_key_dict.models``), so neither a team nor a + restricted key within a team can resolve a provider key for a deployment + it isn't authorized to use. A key restricted to an explicit model list + only narrows the team scope; sentinel-bearing keys (all-proxy-models / + all-team-models) defer to the team scope instead of widening past it. Returns None when the router is unavailable or no authorized deployment matches, so the caller can fall back to default credential resolution. """ if llm_router is None: return None + from litellm.proxy._types import SpecialModelNames + from litellm.proxy.auth.model_checks import get_complete_model_list, get_key_models + + team_id = user_api_key_dict.team_id + team_models = user_api_key_dict.team_models or [] + + proxy_model_list = llm_router.get_model_names(team_id=team_id) + model_access_groups = llm_router.get_model_access_groups() + + raw_key_models = user_api_key_dict.models or [] + sentinel_values = { + SpecialModelNames.all_proxy_models.value, + SpecialModelNames.all_team_models.value, + } + key_is_restricted = bool(raw_key_models) and not (set(raw_key_models) & sentinel_values) + key_model_allowlist = ( + tuple( + dict.fromkeys( + get_key_models( + user_api_key_dict=user_api_key_dict, + proxy_model_list=proxy_model_list, + model_access_groups=model_access_groups, + ) + ) + ) + if key_is_restricted + else () + ) + key_model_allowlist_set = frozenset(key_model_allowlist) + + def _key_may_use(public_model_name: Optional[str]) -> bool: + if not key_model_allowlist_set: + return True + return public_model_name is not None and public_model_name in key_model_allowlist_set + def _provider_credentials(model_id: str) -> Optional[dict]: - credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id) + credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id, team_id=team_id) if credentials is not None and credentials.get("custom_llm_provider") == custom_llm_provider: return credentials return None @@ -332,27 +372,27 @@ def get_team_provider_credentials( deployment_id = model_info.get("id") if deployment_id is None: continue + if not _key_may_use(model_info.get("team_public_model_name") or deployment.get("model_name")): + continue credentials = _provider_credentials(deployment_id) if credentials is not None: return credentials - # 2. Fall back to deployments the team is allowed to access. The - # all-proxy-models sentinel isn't expanded by get_complete_model_list, so - # normalize it to an empty allowlist, which defers to the team-scoped - # proxy model list. A team with a restricted allowlist (e.g. anthropic - # only) therefore never resolves another provider's key. - from litellm.proxy._types import SpecialModelNames - from litellm.proxy.auth.model_checks import get_complete_model_list - + # 2. Fall back to deployments the caller is allowed to access. The key's + # effective allowlist (sentinels and access groups already expanded by + # get_key_models) wins when set; otherwise the team's allowlist applies. + # The all-proxy-models sentinel isn't expanded by + # get_complete_model_list, so normalize it to an empty allowlist, which + # defers to the team-scoped proxy model list. A team or key with a + # restricted allowlist (e.g. anthropic only) therefore never resolves + # another provider's key. grants_all_models = SpecialModelNames.all_proxy_models.value in team_models effective_team_models = [] if grants_all_models else team_models - proxy_model_list = llm_router.get_model_names(team_id=team_id) - model_access_groups = llm_router.get_model_access_groups() models_to_try = list( dict.fromkeys( get_complete_model_list( - key_models=[], + key_models=list(key_model_allowlist), team_models=effective_team_models, proxy_model_list=proxy_model_list, user_model=None, @@ -373,6 +413,28 @@ def get_team_provider_credentials( return None +def apply_team_provider_credentials( + data: dict, # mutable-ok: credentials are merged into the request payload in place, same contract as prepare_data_with_credentials + llm_router: Optional["Router"], + user_api_key_dict: "UserAPIKeyAuth", + custom_llm_provider: str, +) -> None: + """ + Resolve credentials for a provider-only request (no model pinned) via + ``get_team_provider_credentials`` and merge them into ``data`` in-place. + Leaves ``data`` untouched when no authorized deployment matches, so the + caller falls back to environment-variable credentials exactly as before. + """ + credentials = get_team_provider_credentials( + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) + if credentials is None: + return + prepare_data_with_credentials(data=data, credentials=credentials) + + def prepare_data_with_credentials( data: dict, credentials: dict, diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 0c9aa667751..f1bcfbafe58 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -43,10 +43,10 @@ from litellm.litellm_core_utils.cloud_storage_security import ( ) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, + apply_team_provider_credentials, encode_file_id_with_model, extract_file_creation_params, get_credentials_for_model, - get_team_provider_credentials, handle_model_based_routing, prepare_data_with_credentials, validate_managed_files_requirement, @@ -253,6 +253,12 @@ async def route_create_file( _create_file_request=_create_file_request, ) else: + apply_team_provider_credentials( + data=cast(dict, _create_file_request), # cast-ok: TypedDict is a plain dict at runtime; merged in place + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) # get configs for custom_llm_provider llm_provider_config = get_files_provider_config(custom_llm_provider=custom_llm_provider) if llm_provider_config is not None: @@ -735,6 +741,14 @@ async def get_file_content( check_file_id_encoding=True, ) + if not should_route: + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) + from litellm.proxy.openai_files_endpoints.file_content_streaming_handler import ( FileContentStreamingHandler, ) @@ -983,6 +997,12 @@ async def get_file( # Remove file_id from data to avoid "multiple values for keyword argument" error # data was initialized with {"file_id": file_id} data.pop("file_id", None) + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) response = await litellm.afile_retrieve( custom_llm_provider=custom_llm_provider, file_id=file_id, @@ -1183,6 +1203,12 @@ async def delete_file( ) else: data.pop("file_id", None) + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) response = await litellm.afile_delete( custom_llm_provider=custom_llm_provider, file_id=file_id, @@ -1354,14 +1380,12 @@ async def list_files( # No model/target_model_names pinned: resolve upstream credentials from # the team's deployment for this provider so the call is authenticated # against the team's own account (e.g. the team's openai deployment). - team_credentials = get_team_provider_credentials( + apply_team_provider_credentials( + data=data, llm_router=llm_router, - team_models=user_api_key_dict.team_models or [], + user_api_key_dict=user_api_key_dict, custom_llm_provider=custom_llm_provider, - team_id=user_api_key_dict.team_id, ) - if team_credentials is not None: - prepare_data_with_credentials(data=data, credentials=team_credentials) response = await litellm.afile_list( custom_llm_provider=custom_llm_provider, diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 7e573de261b..28d2c62f1f1 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -86,11 +86,16 @@ def is_passthrough_request_using_router_model(request_body: dict, llm_router: Op return False -def is_passthrough_request_streaming(request_body: dict) -> bool: +def is_passthrough_request_streaming(request_body: object) -> bool: """ - Returns True if the request is streaming + Returns True if the request is streaming. + + A JSON body need not be an object, so a list or scalar can reach here; it + carries no streaming flag. """ - return request_body.get("stream", False) + if not isinstance(request_body, dict): + return False + return bool(request_body.get("stream", False)) async def llm_passthrough_factory_proxy_route( @@ -551,8 +556,7 @@ async def is_streaming_request_fn(request: Request) -> bool: _request_body = await get_form_data(request) else: _request_body = await _read_request_body(request) - if _request_body.get("stream"): - return True + return is_passthrough_request_streaming(_request_body) return False @@ -1755,9 +1759,11 @@ async def _base_vertex_proxy_route( ## check for streaming target = str(updated_url) - is_streaming_request = False - if "stream" in str(updated_url): - is_streaming_request = True + if ":rawPredict" in target or ":streamRawPredict" in target: + is_streaming_request = await is_streaming_request_fn(request) + else: + is_streaming_request = "stream" in target + if is_streaming_request: target += "?alt=sse" ## CREATE PASS-THROUGH diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4486cd7de59..c72e3d4ee5b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -5398,7 +5398,7 @@ class ProxyConfig: # decrypt values for k, v in _litellm_params.items(): _litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v) - _litellm_params = LiteLLM_Params(**_litellm_params) + _litellm_params = LiteLLM_Params.model_validate(_litellm_params) else: verbose_proxy_logger.error( @@ -5429,7 +5429,7 @@ class ProxyConfig: # decrypt values for k, v in _litellm_params.items(): _litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v) - _litellm_params = LiteLLM_Params(**_litellm_params) + _litellm_params = LiteLLM_Params.model_validate(_litellm_params) else: verbose_proxy_logger.error( f"Invalid model added to proxy db. Invalid litellm params. litellm_params={_litellm_params}" @@ -6758,6 +6758,9 @@ class ProxyConfig: from litellm.proxy._experimental.mcp_server.oauth2_flow_backfill import ( backfill_null_oauth2_flows, ) + from litellm.proxy._experimental.mcp_server.oauth_issuer_stamp_backfill import ( + backfill_discovery_stamped_issuers, + ) try: if prisma_client is not None: @@ -6767,6 +6770,16 @@ class ProxyConfig: "litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db backfill - {}".format(str(e)) ) + try: + if prisma_client is not None: + await backfill_discovery_stamped_issuers(prisma_client) + except Exception as e: # noqa: BLE001 + verbose_proxy_logger.exception( + "litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db issuer stamp backfill - {}".format( + str(e) + ) + ) + try: await global_mcp_server_manager.reload_servers_from_database() except Exception as e: @@ -6778,6 +6791,31 @@ class ProxyConfig: if self._should_load_db_object(object_type="mcp"): await self._init_mcp_servers_in_db() + async def reload_mcp_servers_from_db(self) -> None: + """Registry refresh only, for the periodic job in store_model_in_db-off deployments. + + Deliberately narrower than ``init_mcp_servers_from_db``: the oauth2_flow backfill is a write + path that only needs to run once at startup, so the cadence here is purely the read-side + reload whose fast-path exemption retries failed OAuth discovery. Gated the same way, so an + admin who excluded mcp from supported_db_objects opts out of this too. + """ + if not self._should_load_db_object(object_type="mcp"): + return + from litellm.proxy._experimental.mcp_server.utils import is_mcp_available + + if not is_mcp_available(): + return + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + try: + await global_mcp_server_manager.reload_servers_from_database() + except Exception as e: # noqa: BLE001 # scheduled job: a reload failure must not kill the recurring retry + verbose_proxy_logger.exception( + "litellm.proxy.proxy_server.py::ProxyConfig:reload_mcp_servers_from_db - {}".format(str(e)) + ) + async def _init_agents_in_db(self, prisma_client: PrismaClient): from litellm.proxy.agent_endpoints.agent_registry import ( global_agent_registry as AGENT_REGISTRY, @@ -8099,6 +8137,22 @@ class ProxyStartupEvent: if store_model_in_db is not True: await proxy_config.init_mcp_servers_from_db() + if prisma_client is not None: + # DB-backed MCP servers are live objects in every mode, so the registry refresh that + # store_model_in_db=True deployments get via the add_deployment job must run here + # too; without it, a server whose OAuth discovery failed at startup is rebuilt only + # by a management write, since the reload fast path is the retry's only driver. + mcp_reload_interval_seconds = proxy_config_reload_interval_seconds + if not isinstance(mcp_reload_interval_seconds, int) or mcp_reload_interval_seconds <= 0: + mcp_reload_interval_seconds = 30 + scheduler.add_job( + proxy_config.reload_mcp_servers_from_db, + "interval", + seconds=mcp_reload_interval_seconds, + id="reload_mcp_servers_job", + replace_existing=True, + misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, + ) await cls._initialize_slack_alerting_jobs( scheduler=scheduler, @@ -11214,11 +11268,15 @@ async def get_all_team_models( if user_teams == "*": team_db_objects = await TeamRepository(prisma_client).table.find_many() - team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects] + team_db_objects_typed = [ + LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) for team_db_object in team_db_objects + ] else: team_db_objects = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_teams}}) - team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects] + team_db_objects_typed = [ + LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) for team_db_object in team_db_objects + ] team_models = _add_team_models_to_all_models( team_db_objects_typed=team_db_objects_typed, @@ -11292,7 +11350,7 @@ async def _populate_team_access_on_models( where={"user_id": user_api_key_dict.user_id} ) if user_db_object is not None: - user_object = LiteLLM_UserTable(**user_db_object.model_dump()) + user_object = LiteLLM_UserTable.model_validate(user_db_object.model_dump()) user_teams = user_object.teams or [] direct_access_models = get_direct_access_models( user_db_object=user_object, @@ -11769,6 +11827,22 @@ def _sort_models( return all_models +def _is_auto_router_model(model: Mapping[str, object]) -> bool: + """ + True for any auto-router deployment, i.e. every `auto_router/*` strategy + (semantic, complexity, adaptive, quality). + + Router._is_auto_router_deployment is deliberately narrower; it answers "is this the + *semantic* auto-router strategy" and returns False for the complexity and adaptive + prefixes, so it is not reusable here. + """ + litellm_params = model.get("litellm_params") + if not isinstance(litellm_params, Mapping): + return False + litellm_model = litellm_params.get("model") + return isinstance(litellm_model, str) and litellm_model.startswith("auto_router/") + + def _paginate_models_response( all_models: List[Dict[str, Any]], page: int, @@ -11827,7 +11901,7 @@ async def _load_team_object_for_model_filter(team_id: str, prisma_client: Prisma if team_db_object is None: verbose_proxy_logger.warning(f"Team {team_id} not found in database") return None - return LiteLLM_TeamTable(**team_db_object.model_dump()) + return LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) except Exception as e: verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}") return None @@ -12063,6 +12137,15 @@ async def model_info_v2( "asc", description="Sort order. Options: asc, desc", ), + exclude_auto_routers: bool | None = fastapi.Query( + False, + description=( + "Omit auto-router deployments (litellm model prefixed `auto_router/`). " + "They select among deployments rather than being deployments themselves, so a " + "caller rendering a deployment list can leave them out. Defaults to false, so " + "existing callers are unaffected" + ), + ), ): """ Paginated model metadata for proxy deployments (pricing, provider, team access). @@ -12230,6 +12313,11 @@ async def model_info_v2( user_api_key_dict=user_api_key_dict, ) + # `is True` because direct-call tests bypass FastAPI, so the Query default arrives as a + # truthy sentinel object rather than False. + if exclude_auto_routers is True: + all_models = [m for m in all_models if not _is_auto_router_model(m)] + # Update total count to include agents search_total_count = len(all_models) @@ -13005,7 +13093,7 @@ def _get_model_group_info( _model_group_info = llm_router.get_model_group_info(model_group=model) if _model_group_info is not None: - model_groups.append(ModelGroupInfoProxy(**_model_group_info.model_dump())) + model_groups.append(ModelGroupInfoProxy.model_validate(_model_group_info.model_dump())) else: model_group_info = ModelGroupInfoProxy( model_group=model, @@ -14724,7 +14812,7 @@ async def update_config_general_settings( ) try: - ConfigGeneralSettings(**{data.field_name: data.field_value}) + ConfigGeneralSettings.model_validate({data.field_name: data.field_value}) except Exception: raise HTTPException( status_code=400, @@ -15170,7 +15258,6 @@ async def get_config_list( "forward_client_headers_to_llm_api": {"type": "Boolean"}, "mcp_required_fields": {"type": "List"}, "cancel_on_disconnect": {"type": "Boolean"}, - "skip_user_budget_on_team_key": {"type": "Boolean"}, "disable_auto_add_proxy_admin_to_teams": {"type": "Boolean"}, } diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 013873179c6..12ecd2ec300 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -155,7 +155,6 @@ async def reserve_budget_for_request( proxy_logging_obj: ProxyLogging, end_user_id: Optional[str] = None, end_user_object: Optional[Any] = None, - skip_user_budget_on_team_key: bool = False, fail_closed_budget_enforcement: bool = False, ) -> Optional[dict]: if valid_token is None or not RouteChecks.is_llm_api_route(route=route): @@ -175,7 +174,6 @@ async def reserve_budget_for_request( proxy_logging_obj=proxy_logging_obj, end_user_id=end_user_id, end_user_object=end_user_object, - skip_user_budget_on_team_key=skip_user_budget_on_team_key, ) if not counters: return None @@ -333,7 +331,6 @@ async def _get_budget_counters( proxy_logging_obj: ProxyLogging, end_user_id: Optional[str] = None, end_user_object: Optional[Any] = None, - skip_user_budget_on_team_key: bool = False, ) -> List[_BudgetCounter]: counters: List[_BudgetCounter] = [] @@ -382,9 +379,8 @@ async def _get_budget_counters( ) ) - is_team_key = team_object is not None and team_object.team_id is not None if ( - not (is_team_key and skip_user_budget_on_team_key) + (team_object is None or team_object.team_id is None) and user_object is not None and user_object.user_id is not None and user_object.max_budget is not None diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 0c525ee9466..42788227acc 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -3308,8 +3308,36 @@ async def ui_view_session_spend_logs( detail="Database not connected", ) - # Build query conditions - where_conditions = {"session_id": session_id} + if _is_admin_view_safe(user_api_key_dict=user_api_key_dict): + scope_sql = "" + scope_params = () + where_conditions = {"session_id": session_id} + else: + try: + permitted_team_ids = ( + await _get_permitted_team_ids_for_spend_logs( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + ) + if _can_user_view_spend_log(user_api_key_dict=user_api_key_dict) + else [] + ) + except Exception: # noqa: BLE001 # mirror /spend/logs/ui: failed team lookup falls back to own-logs-only scope + permitted_team_ids = [] + if permitted_team_ids: + scope_sql = ' AND ("user" = $4 OR team_id = ANY($5::text[]))' + scope_params = (user_api_key_dict.user_id, permitted_team_ids) + where_conditions = { + "session_id": session_id, + "OR": [ + {"user": user_api_key_dict.user_id}, + {"team_id": {"in": permitted_team_ids}}, + ], + } + else: + scope_sql = ' AND "user" = $4' + scope_params = (user_api_key_dict.user_id,) + where_conditions = {"session_id": session_id, "user": user_api_key_dict.user_id} # Calculate pagination offsets skip = (page - 1) * page_size @@ -3318,7 +3346,7 @@ async def ui_view_session_spend_logs( total_records = await SpendLogsRepository(prisma_client).table.count(where=where_conditions) # Query with raw SQL to exclude heavy columns (messages, response, proxy_server_request) - sql_query = """ + sql_query = f""" SELECT request_id, call_type, api_key, spend, total_tokens, prompt_tokens, completion_tokens, "startTime", "endTime", @@ -3328,11 +3356,11 @@ async def ui_view_session_spend_logs( organization_id, end_user, requester_ip_address, session_id, status, mcp_namespaced_tool_name, agent_id FROM "LiteLLM_SpendLogs" - WHERE session_id = $1 + WHERE session_id = $1{scope_sql} ORDER BY "startTime" DESC LIMIT $2 OFFSET $3 """ - result = await prisma_client.db.query_raw(sql_query, session_id, page_size, skip) + result = await prisma_client.db.query_raw(sql_query, session_id, page_size, skip, *scope_params) total_pages = (total_records + page_size - 1) // page_size @@ -3548,7 +3576,7 @@ async def _can_team_member_view_log( team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) if team_row is None: return False - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): return True return _team_member_has_permission( @@ -3640,7 +3668,7 @@ async def _get_permitted_team_ids_for_spend_logs( permitted: List[str] = [] for team_row in team_rows: - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): permitted.append(team_obj.team_id) elif _team_member_has_permission( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 924189fed4b..1e80f15a761 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -109,6 +109,7 @@ from litellm.litellm_core_utils.core_helpers import coerce_token_limit from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.safe_json_loads import safe_json_loads +from litellm.llms import load_guardrail_translation_mappings from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import ( AlertType, @@ -172,6 +173,7 @@ from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionRe from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams if TYPE_CHECKING: + from mcp.types import CallToolResult from opentelemetry.trace import Span as _Span from prisma.client import TransactionManager @@ -2470,6 +2472,59 @@ class ProxyLogging: if raised: raise raised[0] + async def post_mcp_call_hook( + self, + response: "CallToolResult", + request_data: Mapping[str, Any], + user_api_key_dict: UserAPIKeyAuth | None = None, + ) -> "CallToolResult": + """ + Run guardrails configured for ``post_mcp_call`` against an MCP tool result. + + The MCP counterpart of ``post_call_success_hook``: guardrails that + implement ``apply_guardrail`` see the tool result's text through the + unified guardrail seam (``MCPGuardrailTranslationHandler``), so a text + guardrail can mask sensitive values in the result without any MCP-specific + code of its own. Guardrails that instead implement + ``async_post_mcp_tool_call_hook`` are dispatched by + ``Logging.async_post_mcp_tool_call_hook`` and are not run here. + + A guardrail that rejects the result raises, and the exception propagates + (matching the inbound ``pre_mcp_call`` behavior) rather than being + swallowed into an unguarded result. + """ + caps = ProxyLogging._callback_capabilities() + if not caps.has_guardrail: + return response + + handler_cls = load_guardrail_translation_mappings().get(CallTypes.call_mcp_tool) + if handler_cls is None: + verbose_proxy_logger.debug("MCP guardrail translation handler unavailable; skipping post_mcp_call hook") + return response + + for callback in caps.resolved_callbacks: + if not isinstance(callback, CustomGuardrail): + continue + if "apply_guardrail" not in type(callback).__dict__: + continue + if ( + callback.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_mcp_call) + is not True + ): + continue + response = await self._run_guardrail_with_metrics( + callback, + handler_cls().process_output_response( + response=response, + guardrail_to_apply=callback, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ), + "post_mcp_call", + ) + return response + async def post_call_response_headers_hook( self, data: dict, diff --git a/litellm/repositories/base_repository.py b/litellm/repositories/base_repository.py index 40aeb6df3de..755e4595c01 100644 --- a/litellm/repositories/base_repository.py +++ b/litellm/repositories/base_repository.py @@ -3,38 +3,58 @@ Base repository class with common functionality. """ from abc import ABC, abstractmethod -from typing import Any, Dict, Generic, List, Optional, Type, TypeVar +from collections.abc import Iterable, Mapping, Sequence +from typing import Any, Dict, Generic, List, Optional, Protocol, Tuple, Type, TypeVar, Union, runtime_checkable from pydantic import BaseModel T = TypeVar("T", bound=BaseModel) -def _record_to_dict(record: Any) -> Dict[str, Any]: - if isinstance(record, dict): - return record - if hasattr(record, "model_dump") and callable(record.model_dump): +@runtime_checkable +class SupportsModelDump(Protocol): + def model_dump(self) -> Dict[str, object]: ... + + +@runtime_checkable +class SupportsDict(Protocol): + def dict(self) -> Dict[str, object]: ... + + +DbRecord = Union[ + Mapping[str, object], + SupportsModelDump, + SupportsDict, + Sequence[Tuple[str, object]], +] + + +def record_to_dict(record: DbRecord) -> Mapping[str, object]: + """Project a database record into a mapping of column name to value.""" + if isinstance(record, SupportsModelDump): return record.model_dump() - if hasattr(record, "dict") and callable(record.dict): + if isinstance(record, SupportsDict): return record.dict() - return dict(record) + if isinstance(record, Mapping): + return record + return {key: value for key, value in record} class BaseRepository(ABC, Generic[T]): """Abstract base class for all repositories.""" - def __init__(self, prisma_client: Any): + def __init__(self, prisma_client: Any): # any-ok: PrismaClient is an untyped runtime wrapper self._prisma_client = prisma_client @property - def prisma_client(self) -> Any: + def prisma_client(self) -> Any: # any-ok: PrismaClient is an untyped runtime wrapper if self._prisma_client is None: raise RuntimeError("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys") return self._prisma_client @property @abstractmethod - def table(self) -> Any: + def table(self) -> Any: # any-ok: Prisma table actions are reached through the untyped client wrapper """Return the Prisma table for this repository.""" ... @@ -44,21 +64,15 @@ class BaseRepository(ABC, Generic[T]): """Return the domain model class for this repository.""" ... - def _to_model(self, record: Any) -> Optional[T]: + def _to_model(self, record: Optional[DbRecord]) -> Optional[T]: """Convert a database record to a domain model.""" if record is None: return None - return self.model_class(**_record_to_dict(record)) + return self.model_class.model_validate(record_to_dict(record)) - def _to_model_list(self, records: List[Any]) -> List[T]: + def _to_model_list(self, records: Iterable[Optional[DbRecord]]) -> List[T]: """Convert a list of database records to domain models.""" - result: List[T] = [] - for r in records: - if r is not None: - model = self._to_model(r) - if model is not None: - result.append(model) - return result + return [model for record in records if record is not None and (model := self._to_model(record)) is not None] async def find_by_id(self, id_value: str, id_field: str = "id") -> Optional[T]: """Find a record by its primary key.""" diff --git a/litellm/repositories/organization_repository.py b/litellm/repositories/organization_repository.py index d5f8c990001..99c4a881736 100644 --- a/litellm/repositories/organization_repository.py +++ b/litellm/repositories/organization_repository.py @@ -26,10 +26,8 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): async def find_by_alias(self, organization_alias: str) -> Optional[LiteLLM_OrganizationTable]: """Find an organization by alias.""" - records = await self.table.find_many(where={"organization_alias": organization_alias}) - if records: - return self._to_model(records[0]) - return None + organizations = await self.find_many(where={"organization_alias": organization_alias}) + return organizations[0] if organizations else None async def create_organization( self, diff --git a/litellm/repositories/project_repository.py b/litellm/repositories/project_repository.py index 86faaf2e13c..27cb346e1b1 100644 --- a/litellm/repositories/project_repository.py +++ b/litellm/repositories/project_repository.py @@ -24,15 +24,12 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): async def find_by_alias(self, project_alias: str) -> Optional[LiteLLM_ProjectTable]: """Find a project by alias.""" - records = await self.table.find_many(where={"project_alias": project_alias}) - if records: - return self._to_model(records[0]) - return None + projects = await self.find_many(where={"project_alias": project_alias}) + return projects[0] if projects else None async def find_by_team_id(self, team_id: str) -> List[LiteLLM_ProjectTable]: """Find all projects belonging to a team.""" - records = await self.table.find_many(where={"team_id": team_id}) - return self._to_model_list(records) + return await self.find_many(where={"team_id": team_id}) async def create_project( self, diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 68875bd7972..25437cfe49a 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -3,55 +3,59 @@ Team repository for database operations on LiteLLM_TeamTable. """ import json +from collections.abc import Mapping from datetime import datetime from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type from pydantic import TypeAdapter from litellm.models.team import LiteLLM_TeamTable, Member -from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.base_repository import ( + BaseRepository, + DbRecord, + record_to_dict, +) if TYPE_CHECKING: from prisma import Prisma _MEMBERS_WITH_ROLES_ADAPTER = TypeAdapter(list[Member]) +_JSON_ENCODED_TEAM_FIELDS = ( + "metadata", + "model_spend", + "model_max_budget", + "router_settings", + "budget_limits", + "members_with_roles", +) class TeamRepository(BaseRepository[LiteLLM_TeamTable]): """Repository for team database operations.""" @property - def table(self) -> Any: + def table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper return self.prisma_client.db.litellm_teamtable @property - def deleted_table(self) -> Any: + def deleted_table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper return self.prisma_client.db.litellm_deletedteamtable @property def model_class(self) -> Type[LiteLLM_TeamTable]: return LiteLLM_TeamTable - def _to_model(self, record: Any) -> Optional[LiteLLM_TeamTable]: + def _to_model(self, record: Optional[DbRecord]) -> Optional[LiteLLM_TeamTable]: """Convert a database record to a Team model.""" if record is None: return None - data = record.dict() if hasattr(record, "dict") else dict(record) + data = { + field: json.loads(value) if field in _JSON_ENCODED_TEAM_FIELDS and isinstance(value, str) else value + for field, value in record_to_dict(record).items() + } - json_fields = [ - "metadata", - "model_spend", - "model_max_budget", - "router_settings", - "budget_limits", - "members_with_roles", - ] - for field in json_fields: - if isinstance(data.get(field), str): - data[field] = json.loads(data[field]) - - return LiteLLM_TeamTable(**data) + return LiteLLM_TeamTable.model_validate(data) async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> List[Member]: """Return the team's members_with_roles, locking the row FOR UPDATE. @@ -103,8 +107,8 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): organization_id: Optional[str] = None, admins: Optional[List[str]] = None, members: Optional[List[str]] = None, - members_with_roles: Optional[Dict[str, Any]] = None, - metadata: Optional[Dict[str, Any]] = None, + members_with_roles: Optional[Mapping[str, object]] = None, + metadata: Optional[Mapping[str, object]] = None, max_budget: Optional[float] = None, soft_budget: Optional[float] = None, models: Optional[List[str]] = None, @@ -115,7 +119,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): object_permission_id: Optional[str] = None, ) -> LiteLLM_TeamTable: """Create a new team.""" - data: Dict[str, Any] = {"team_id": team_id} + data: Dict[str, object] = {"team_id": team_id} if team_alias is not None: data["team_alias"] = team_alias if organization_id is not None: @@ -154,8 +158,8 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): organization_id: Optional[str] = None, admins: Optional[List[str]] = None, members: Optional[List[str]] = None, - members_with_roles: Optional[Dict[str, Any]] = None, - metadata: Optional[Dict[str, Any]] = None, + members_with_roles: Optional[Mapping[str, object]] = None, + metadata: Optional[Mapping[str, object]] = None, max_budget: Optional[float] = None, soft_budget: Optional[float] = None, models: Optional[List[str]] = None, @@ -167,7 +171,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): object_permission_id: Optional[str] = None, ) -> Optional[LiteLLM_TeamTable]: """Update a team.""" - data: Dict[str, Any] = {} + data: Dict[str, object] = {} if team_alias is not None: data["team_alias"] = team_alias if organization_id is not None: @@ -228,9 +232,9 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return team - def _build_archive_data(self, team: LiteLLM_TeamTable) -> Dict[str, Any]: + def _build_archive_data(self, team: LiteLLM_TeamTable) -> Dict[str, object]: """Build archive data dict with only columns that exist in LiteLLM_DeletedTeamTable.""" - data: Dict[str, Any] = {"team_id": team.team_id} + data: Dict[str, object] = {"team_id": team.team_id} if team.team_alias is not None: data["team_alias"] = team.team_alias if team.organization_id is not None: diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py index 19352c1b3c4..f7795f15fd5 100644 --- a/litellm/repositories/verification_token_repository.py +++ b/litellm/repositories/verification_token_repository.py @@ -3,14 +3,18 @@ VerificationToken repository for database operations on LiteLLM_VerificationToke """ import json -from collections.abc import Iterator, Mapping +from collections.abc import Mapping from datetime import datetime -from typing import TYPE_CHECKING, Any, Protocol +from typing import TYPE_CHECKING, Any from litellm.models.verification_token import ( LiteLLM_VerificationToken, ) -from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.base_repository import ( + BaseRepository, + DbRecord, + record_to_dict, +) if TYPE_CHECKING: from prisma.models import ( @@ -19,11 +23,17 @@ if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient - -class _DictConvertible(Protocol): - def dict(self) -> dict[str, object]: ... - - def __iter__(self) -> Iterator[tuple[str, object]]: ... +_JSON_ENCODED_TOKEN_FIELDS = ( + "aliases", + "config", + "permissions", + "metadata", + "model_spend", + "model_max_budget", + "router_settings", + "budget_limits", + "litellm_budget_table", +) class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): @@ -46,31 +56,21 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): def model_class(self) -> type[LiteLLM_VerificationToken]: return LiteLLM_VerificationToken - def _to_model(self, record: _DictConvertible | None) -> LiteLLM_VerificationToken | None: + def _to_model(self, record: DbRecord | None) -> LiteLLM_VerificationToken | None: """Convert a database record to a VerificationToken model.""" if record is None: return None - data = record.dict() if hasattr(record, "dict") else dict(record) - - json_fields = [ - "aliases", - "config", - "permissions", - "metadata", - "model_spend", - "model_max_budget", - "router_settings", - "budget_limits", - "litellm_budget_table", - ] - for field in json_fields: - value = data.get(field) - if isinstance(value, str): - data[field] = json.loads(value) - - if data.get("org_id") is None and data.get("organization_id") is not None: - data["org_id"] = data["organization_id"] + decoded = { + field: json.loads(value) if field in _JSON_ENCODED_TOKEN_FIELDS and isinstance(value, str) else value + for field, value in record_to_dict(record).items() + } + organization_id = decoded.get("organization_id") + data = ( + decoded + if decoded.get("org_id") is not None or organization_id is None + else {**decoded, "org_id": organization_id} + ) return LiteLLM_VerificationToken.model_validate(data) diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index cb680dc8b86..9d30e40cd54 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -795,20 +795,33 @@ class LiteLLM_Proxy_MCP_Handler: proxy_logging_obj=proxy_logging_obj, ) + if proxy_logging_obj: + result = await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data=( + litellm_logging_obj.model_call_details + if litellm_logging_obj + else {"mcp_tool_name": tool_name} + ), + user_api_key_dict=user_api_key_auth, + ) + if litellm_logging_obj: try: litellm_logging_obj.post_call(original_response=result) - end_time = datetime.now() await litellm_logging_obj.async_post_mcp_tool_call_hook( kwargs=litellm_logging_obj.model_call_details, response_obj=result, start_time=start_time, - end_time=end_time, + end_time=datetime.now(), ) + except Exception: + verbose_logger.exception("Failed to run post-call logging for MCP tool call %s", tool_name) + try: await litellm_logging_obj.async_success_handler( result=result, start_time=start_time, - end_time=end_time, + end_time=datetime.now(), ) except Exception: verbose_logger.exception("Failed to log MCP tool call success for %s", tool_name) diff --git a/litellm/router.py b/litellm/router.py index 487d6a31226..ac00cdbc2b0 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -19,6 +19,7 @@ import threading import time import traceback from collections import defaultdict +from collections.abc import Mapping from functools import lru_cache from typing import ( TYPE_CHECKING, @@ -70,6 +71,7 @@ from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, coerce_token_limit, get_metadata_variable_name_from_kwargs, + get_or_create_metadata_bucket, ) from litellm.litellm_core_utils.coroutine_checker import coroutine_checker from litellm.litellm_core_utils.credential_accessor import CredentialAccessor @@ -109,6 +111,9 @@ from litellm.router_utils.batch_utils import ( replace_model_in_jsonl, should_replace_model_in_jsonl, ) +from litellm.router_utils.auto_router_model_naming import ( + classify_strategy_router_model, +) from litellm.router_utils.client_initalization_utils import InitalizeCachedClient from litellm.router_utils.clientside_credential_handler import ( get_dynamic_litellm_params, @@ -206,8 +211,10 @@ from litellm.types.utils import ( from litellm.types.utils import ModelInfo from litellm.types.utils import ModelInfo as ModelMapInfo from litellm.types.utils import ( + PROMPT_QUOTING_ROUTING_DECISION_FIELDS, ModelResponseStream, StandardLoggingPayload, + StandardLoggingRoutingDecision, Usage, ) from litellm.utils import ( @@ -270,6 +277,43 @@ def _cost_value_as_float(value: Union[str, int, float, None]) -> Optional[float] return None +def model_info_is_active_for_environment(model_info: Mapping[str, object] | None) -> bool: + """Single owner of the environment-gating rule: a deployment whose model_info names + `supported_environments` loads only on pods whose LITELLM_ENVIRONMENT is in that list. + `Router.deployment_is_active_for_environment` delegates here, and the model-write + endpoints consult the same rule to tell a deliberately inactive model from one that + was dropped by a failed reload.""" + if model_info is None: + return True + supported_environments = model_info.get("supported_environments") + if supported_environments is None: + return True + if not isinstance(supported_environments, (list, tuple)): + raise ValueError( + f"supported_environments must be a list of {VALID_LITELLM_ENVIRONMENTS}. " + f"but set as: {supported_environments} for model_info: {model_info}" + ) + litellm_environment = get_secret_str(secret_name="LITELLM_ENVIRONMENT") + if litellm_environment is None: + raise ValueError("Set 'supported_environments' for model but not 'LITELLM_ENVIRONMENT' set in .env") + + if litellm_environment not in VALID_LITELLM_ENVIRONMENTS: + raise ValueError( + f"LITELLM_ENVIRONMENT must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {litellm_environment}" + ) + + for _env in supported_environments: + if _env not in VALID_LITELLM_ENVIRONMENTS: + raise ValueError( + f"supported_environments must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {_env} " + f"for model_info: {model_info}" + ) + + if litellm_environment in supported_environments: + return True + return False + + _PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT") @@ -7585,15 +7629,7 @@ class Router: but NOT "auto_router/complexity_router" or "auto_router/adaptive_router" (which use the complexity-router and adaptive-router strategies). """ - if litellm_params.model.startswith("auto_router/complexity_router"): - return False # This is handled by complexity_router - if litellm_params.model.startswith("auto_router/adaptive_router"): - return False # This is handled by adaptive_router - if litellm_params.model.startswith("auto_router/quality_router"): - return False # This is handled by quality_router - if litellm_params.model.startswith("auto_router/"): - return True - return False + return classify_strategy_router_model(litellm_params.model) == "semantic" @staticmethod def _deployment_tags(deployment: Deployment) -> tuple[str, ...]: @@ -7648,9 +7684,7 @@ class Router: Returns True if the litellm_params model starts with "auto_router/complexity_router" """ - if litellm_params.model.startswith("auto_router/complexity_router"): - return True - return False + return classify_strategy_router_model(litellm_params.model) == "complexity" def init_complexity_router_deployment(self, deployment: Deployment): """ @@ -7700,7 +7734,7 @@ class Router: def _is_adaptive_router_deployment(self, litellm_params: LiteLLM_Params) -> bool: """True when this deployment opts in via the `auto_router/adaptive_router` model prefix.""" - return litellm_params.model.startswith("auto_router/adaptive_router") + return classify_strategy_router_model(litellm_params.model) == "adaptive" def _deployment_participates_in_adaptive_routing(self, litellm_params: LiteLLM_Params) -> bool: """True when this deployment owns an `adaptive_routers` entry once finalized: @@ -7926,9 +7960,7 @@ class Router: Returns True if the litellm_params model starts with "auto_router/quality_router". """ - if litellm_params.model.startswith("auto_router/quality_router"): - return True - return False + return classify_strategy_router_model(litellm_params.model) == "quality" def init_quality_router_deployment(self, deployment: Deployment): """ @@ -7982,30 +8014,7 @@ class Router: - ValueError: If LITELLM_ENVIRONMENT is not set in .env or not one of the valid values - ValueError: If supported_environments is not set in model_info or not one of the valid values """ - if ( - deployment.model_info is None - or "supported_environments" not in deployment.model_info - or deployment.model_info["supported_environments"] is None - ): - return True - litellm_environment = get_secret_str(secret_name="LITELLM_ENVIRONMENT") - if litellm_environment is None: - raise ValueError("Set 'supported_environments' for model but not 'LITELLM_ENVIRONMENT' set in .env") - - if litellm_environment not in VALID_LITELLM_ENVIRONMENTS: - raise ValueError( - f"LITELLM_ENVIRONMENT must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {litellm_environment}" - ) - - for _env in deployment.model_info["supported_environments"]: - if _env not in VALID_LITELLM_ENVIRONMENTS: - raise ValueError( - f"supported_environments must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {_env} for deployment: {deployment}" - ) - - if litellm_environment in deployment.model_info["supported_environments"]: - return True - return False + return model_info_is_active_for_environment(model_info=deployment.model_info) def set_model_list(self, model_list: list): original_model_list = copy.deepcopy(model_list) @@ -8630,6 +8639,33 @@ class Router: raise Exception("Model Name invalid - {}".format(type(model))) return None + @staticmethod + def _deployment_usable_by_team(model: Union[Mapping, Deployment], team_id: str | None) -> bool: + """ + A team-scoped deployment (``model_info.team_id`` set) is only usable by + callers from that same team; deployments without a team owner are shared. + """ + model_info = model.get("model_info") if isinstance(model, dict) else model.model_info + owner_team_id = model_info.get("team_id") if model_info is not None else None + return owner_team_id is None or owner_team_id == team_id + + def _get_model_group_deployment_usable_by_team( + self, model_group_name: str, team_id: str | None + ) -> Deployment | None: + """ + Like ``get_deployment_by_model_group_name``, but skips deployments owned + by other teams so a shared model name never resolves another team's + credentials. + """ + indices = self.model_name_to_deployment_indices.get(model_group_name) or () + usable = ( + self.model_list[idx] for idx in indices if self._deployment_usable_by_team(self.model_list[idx], team_id) + ) + first_usable = next(usable, None) + if first_usable is None: + return None + return Deployment(**first_usable) if isinstance(first_usable, dict) else first_usable + 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 @@ -8664,7 +8700,10 @@ class Router: model_id: Model ID or model name from model_list (e.g., "gpt-4o-litellm") team_id: Optional team id of the caller. When set, team-scoped deployments (indexed by team public model name, including team - wildcard models like "openai/*") are also considered. + wildcard models like "openai/*") are also considered. Name and + wildcard lookups never resolve a deployment owned by a + different team, so shared model names can't leak another + team's credentials. Returns: Dictionary containing api_key, api_base, custom_llm_provider, etc. @@ -8681,7 +8720,7 @@ class Router: # If not found, try by model_group_name if deployment is None: - deployment = self.get_deployment_by_model_group_name(model_group_name=model_id) + deployment = self._get_model_group_deployment_usable_by_team(model_group_name=model_id, team_id=team_id) # If not found, check team-scoped deployments whose team public model # name exactly matches model_id (wildcard team names are matched via @@ -8698,7 +8737,12 @@ class Router: if deployment is None: team_pattern_router = self.team_pattern_routers.get(team_id) if team_id is not None else None team_wildcard_models = (team_pattern_router.route(model_id) or []) if team_pattern_router else [] - potential_wildcard_models = team_wildcard_models or self.pattern_router.route(model_id) or [] + global_wildcard_models = [ + wildcard_model + for wildcard_model in (self.pattern_router.route(model_id) or []) + if self._deployment_usable_by_team(wildcard_model, team_id) + ] + potential_wildcard_models = team_wildcard_models or global_wildcard_models if potential_wildcard_models: # Use the first matching wildcard deployment deployment_dict = potential_wildcard_models[0] @@ -9519,7 +9563,12 @@ class Router: return None # Strategy 1: Check if model_id directly matches a model_name or deployment ID - if model_id in self.model_names or self.has_model_id(model_id): + if model_id in self.model_names: + return model_id + if self.has_model_id(model_id): + deployment = self.get_deployment(model_id=model_id) + if deployment is not None and deployment.model_name: + return deployment.model_name return model_id # Strategy 2: Search through router's model_list to find by litellm_params.model @@ -11112,6 +11161,7 @@ class Router: router_strategy = self._select_pre_routing_strategy(model=model, request_kwargs=request_kwargs) if router_strategy is None: + self._record_routing_decision(request_kwargs=request_kwargs, routing_decision=None) return None pre_routing_hook_response = await router_strategy.async_pre_routing_hook( @@ -11121,6 +11171,10 @@ class Router: input=input, specific_deployment=specific_deployment, ) + self._record_routing_decision( + request_kwargs=request_kwargs, + routing_decision=(pre_routing_hook_response.routing_decision if pre_routing_hook_response else None), + ) # `model` (the alias, e.g. "smart-router") is never the deployment actually # called - apply the alias's own litellm_params (besides `model` itself, @@ -11139,6 +11193,68 @@ class Router: return pre_routing_hook_response + @staticmethod + def _record_routing_decision( + request_kwargs: dict, + routing_decision: StandardLoggingRoutingDecision | None, + ) -> None: + """Make the request's metadata describe THIS routing attempt, and only this one. + + Fallbacks re-enter the hook with the same `request_kwargs`, so an attempt that + picks a plain model group after an auto-router group failed must clear the + earlier decision; leaving it would attribute the first router's tier and cause + to the deployment that actually served the request. Every attempt therefore + writes or clears, never just writes. + """ + if routing_decision is None: + for bucket in (request_kwargs.get("metadata"), request_kwargs.get("litellm_metadata")): + if isinstance(bucket, dict): + bucket.pop("routing_decision", None) + return + + # `get_or_create_metadata_bucket` is the single owner of "which dict holds + # proxy-internal metadata": it picks `litellm_metadata` when present (so the + # decision never lands in the `metadata` dict that routes like /v1/messages + # forward to the provider) and replaces a non-dict value rather than silently + # skipping the write. + _, metadata_bucket = get_or_create_metadata_bucket(request_kwargs) + metadata_bucket["routing_decision"] = Router._redact_prompt_text_if_needed( + request_kwargs=request_kwargs, routing_decision=routing_decision + ) + + @staticmethod + def _redact_prompt_text_if_needed( + request_kwargs: Mapping[str, Any], + routing_decision: StandardLoggingRoutingDecision, + ) -> StandardLoggingRoutingDecision: + """Drop verbatim prompt text from the record when message logging is redacted. + + An operator who turns message logging off has said prompt content must not reach + the logs, so the fields that quote the prompt (the matched keywords, and the + signals that name them) are omitted. Derived values are kept, because a tier, a + cause, a score or an escalation flag aggregates the prompt rather than + reproducing any of it, and dropping them would leave the row unexplainable for + no privacy gain. Applied here rather than in each strategy so a strategy added + later cannot bypass it. + """ + from litellm.litellm_core_utils.redact_messages import ( + should_redact_message_logging, + ) + + if not should_redact_message_logging( + { + "litellm_params": request_kwargs, + "standard_callback_dynamic_params": request_kwargs.get("standard_callback_dynamic_params"), + } + ): + return routing_decision + kept = { + field: value + for field, value in routing_decision.items() + if field not in PROMPT_QUOTING_ROUTING_DECISION_FIELDS + } + return cast(StandardLoggingRoutingDecision, kept) # cast-ok: dropping optional keys preserves the type + def get_available_deployment( self, model: str, diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index ec84eb1decf..e8fcec2667d 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -23,6 +23,7 @@ from litellm._logging import verbose_router_logger from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_last_user_message, ) +from litellm.types.utils import StandardLoggingRoutingDecision from litellm.router_strategy.adaptive_router.bandit import ( BanditCell, apply_delta, @@ -193,7 +194,17 @@ class AdaptiveRouter: if isinstance(kwargs_metadata, dict): kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = chosen_model - return PreRoutingHookResponse(model=chosen_model, messages=messages) + return PreRoutingHookResponse( + model=chosen_model, + messages=messages, + routing_decision=StandardLoggingRoutingDecision( + router_model_name=self.router_name, + router_type="adaptive", + routed_model=chosen_model, + cause="bandit", + request_type=request_type.value, + ), + ) # ---- Pick model ------------------------------------------------------ diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index e5268b5107b..933c6d170cf 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -18,14 +18,21 @@ from __future__ import annotations import asyncio import random import re -from typing import TYPE_CHECKING, Any, Literal, Union, cast +from collections.abc import Mapping +from typing import TYPE_CHECKING, Any, Literal, NamedTuple, Union, cast from pydantic import BaseModel from litellm._logging import verbose_router_logger from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger -from litellm.types.utils import ModelResponse +from litellm.llms.base_llm.base_utils import type_to_response_format_param +from litellm.types.utils import ( + ModelResponse, + RoutingDecisionCause, + StandardLoggingRoutingDecision, + StandardLoggingRoutingDecisionTierBoundaries, +) from .config import ( DEFAULT_CODE_KEYWORDS, @@ -112,6 +119,16 @@ def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any] } +def _effective_turn_off_message_logging(request_kwargs: Mapping[str, Any] | None) -> bool | None: + from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( + initialize_standard_callback_dynamic_params, + ) + + return initialize_standard_callback_dynamic_params(dict(request_kwargs) if request_kwargs else {}).get( + "turn_off_message_logging" + ) + + class DimensionScore: """Represents a score for a single dimension with optional signal.""" @@ -123,6 +140,27 @@ class DimensionScore: self.signal = signal +class KeywordOverride(NamedTuple): + """A keyword_tier_rules match: the winning tier and, on the lexical path, the keyword that fired.""" + + tier: ComplexityTier + matched_keyword: str | None + + +class ClassificationOutcome(NamedTuple): + """What the classifier decided and which mechanism actually produced it. + + `cause` reflects the path that ran, not the configured classifier_type: an LLM + classifier that fails falls back to the heuristic scorer and reports it. + `score` is None on the LLM path, which produces a tier label and no score. + """ + + tier: ComplexityTier + score: float | None + signals: tuple[str, ...] + cause: Literal["heuristic_scorer", "reasoning_override", "llm_classifier"] + + class ComplexityRouter(CustomLogger): """ Complexity router that classifies requests and routes to appropriate models. @@ -244,6 +282,7 @@ class ComplexityRouter(CustomLogger): def _score_keyword_match( self, text: str, + disclosable_text: str, keywords: list[str], name: str, signal_label: str, @@ -252,6 +291,15 @@ class ComplexityRouter(CustomLogger): ) -> tuple[DimensionScore, int]: """Score based on keyword matches using word boundary matching. + Scoring reads `text`, which for most dimensions includes the system prompt. + The signal names only the terms that also appear in `disclosable_text`, the + caller's own message: signals are persisted to the request's spend log, which + the caller can read, so naming a term matched solely in the system prompt would + let a caller recover configured terms from a prompt it cannot see. Terms it did + not supply are reported as a count instead, which explains the score without + disclosing anything. `disclosable_text` is required rather than defaulted so a + future dimension has to state which text it is willing to quote. + Returns: Tuple of (DimensionScore, match_count) so callers can reuse the count. """ @@ -260,18 +308,13 @@ class ComplexityRouter(CustomLogger): matches = [kw for kw in keywords if self._keyword_matches(text, kw)] match_count = len(matches) + if match_count < low_threshold: + return DimensionScore(name, score_none, None), match_count - if match_count >= high_threshold: - return ( - DimensionScore(name, score_high, f"{signal_label} ({', '.join(matches[:3])})"), - match_count, - ) - if match_count >= low_threshold: - return ( - DimensionScore(name, score_low, f"{signal_label} ({', '.join(matches[:3])})"), - match_count, - ) - return DimensionScore(name, score_none, None), match_count + disclosable = [kw for kw in matches if self._keyword_matches(disclosable_text, kw)] + detail = ", ".join(disclosable[:3]) if disclosable else f"{match_count} matches" + score = score_high if match_count >= high_threshold else score_low + return DimensionScore(name, score, f"{signal_label} ({detail})"), match_count def _score_multi_step(self, text: str) -> DimensionScore: """Score based on multi-step patterns.""" @@ -288,8 +331,19 @@ class ComplexityRouter(CustomLogger): return DimensionScore("questionComplexity", 0, None) def classify(self, prompt: str, system_prompt: str | None = None) -> tuple[ComplexityTier, float, list[str]]: + """Classify a prompt by complexity, discarding which rule decided the tier. + + Kept for callers that only need the tier and score; `_score_and_classify` is the + single computation behind both, so the two can never disagree. """ - Classify a prompt by complexity. + tier, score, signals, _cause = self._score_and_classify(prompt, system_prompt) + return tier, score, list(signals) + + def _score_and_classify( + self, prompt: str, system_prompt: str | None = None + ) -> tuple[ComplexityTier, float, tuple[str, ...], Literal["heuristic_scorer", "reasoning_override"]]: + """ + Classify a prompt by complexity, reporting whether the score chose the tier. Args: prompt: The user's prompt/message. @@ -315,6 +369,7 @@ class ComplexityRouter(CustomLogger): # Score all dimensions, capturing match counts where needed code_score, _ = self._score_keyword_match( full_text, + user_text, self.code_keywords, "codePresence", "code", @@ -322,6 +377,7 @@ class ComplexityRouter(CustomLogger): (0, 0.5, 1.0), ) reasoning_score, reasoning_match_count = self._score_keyword_match( + user_text, user_text, self.reasoning_keywords, "reasoningMarkers", @@ -331,6 +387,7 @@ class ComplexityRouter(CustomLogger): ) technical_score, _ = self._score_keyword_match( full_text, + user_text, self.technical_keywords, "technicalTerms", "technical", @@ -339,6 +396,7 @@ class ComplexityRouter(CustomLogger): ) simple_score, _ = self._score_keyword_match( full_text, + user_text, self.simple_keywords, "simpleIndicators", "simple", @@ -366,48 +424,112 @@ class ComplexityRouter(CustomLogger): # Check for reasoning override (2+ reasoning markers) # Reuse match count from _score_keyword_match to avoid scanning twice if reasoning_match_count >= 2: - return ComplexityTier.REASONING, weighted_score, signals + return ComplexityTier.REASONING, weighted_score, tuple(signals), "reasoning_override" # Map score to tier - boundaries = self.config.tier_boundaries - simple_medium = boundaries.get("simple_medium", 0.15) - medium_complex = boundaries.get("medium_complex", 0.35) - complex_reasoning = boundaries.get("complex_reasoning", 0.60) - - if weighted_score < simple_medium: + boundaries = self._effective_tier_boundaries() + if weighted_score < boundaries["simple_medium"]: tier = ComplexityTier.SIMPLE - elif weighted_score < medium_complex: + elif weighted_score < boundaries["medium_complex"]: tier = ComplexityTier.MEDIUM - elif weighted_score < complex_reasoning: + elif weighted_score < boundaries["complex_reasoning"]: tier = ComplexityTier.COMPLEX else: tier = ComplexityTier.REASONING - return tier, weighted_score, signals + return tier, weighted_score, tuple(signals), "heuristic_scorer" + + def _effective_tier_boundaries(self) -> StandardLoggingRoutingDecisionTierBoundaries: + """The tier boundaries in effect, with the documented defaults filled in. + + Shared by score-to-tier mapping and the per-request routing decision snapshot, + so a logged decision always reflects the boundaries that actually applied. + """ + boundaries = self.config.tier_boundaries + return StandardLoggingRoutingDecisionTierBoundaries( + simple_medium=boundaries.get("simple_medium", 0.15), + medium_complex=boundaries.get("medium_complex", 0.35), + complex_reasoning=boundaries.get("complex_reasoning", 0.60), + ) + + def _build_routing_decision( + self, + *, + routed_model: str, + cause: RoutingDecisionCause, + tier: ComplexityTier | None = None, + score: float | None = None, + signals: tuple[str, ...] | None = None, + matched_keyword: str | None = None, + escalation_keyword: str | None = None, + escalated: bool = False, + classifier_model: str | None = None, + ) -> StandardLoggingRoutingDecision: + """Assemble the per-request provenance record for this router's decision. + + Optional facts are omitted rather than set to None, so a spend log row only + carries the keys that applied to its path. `tier_boundaries` rides with + `score` because the score is only interpretable against the boundaries that + mapped it to a tier. + """ + decision = StandardLoggingRoutingDecision( + router_model_name=self.model_name, + router_type="complexity", + routed_model=routed_model, + cause=cause, + ) + if tier is not None: + decision["tier"] = tier.value + if score is not None: + decision["score"] = score + decision["tier_boundaries"] = self._effective_tier_boundaries() + if signals: + # Stored as a list because this record is serialized to JSON for the spend + # log and read back as an array by the dashboard; a sequence type that only + # happens to survive the serializer would make the wire shape depend on it. + decision["signals"] = list(signals) + if matched_keyword is not None: + decision["matched_keyword"] = matched_keyword + if escalation_keyword is not None: + # Two separate facts: the caller asked to escalate, and whether the tier + # actually moved. A request that escalates from an already-highest tier has + # nowhere to go, so it records the keyword with escalated=False rather than + # dropping the ask (which reads as an ordinary route) or claiming a bump + # that never happened. Every path reports both the same way. + decision["escalation_keyword"] = escalation_keyword + decision["escalated"] = escalated + if classifier_model is not None: + decision["classifier_model"] = classifier_model + return decision async def aclassify( self, prompt: str, system_prompt: str | None = None, request_kwargs: dict[str, Any] | None = None, - ) -> tuple[ComplexityTier, float, list[str]]: + ) -> ClassificationOutcome: """ Classify a prompt by complexity, using the LLM classifier when configured. Falls back to the local heuristic scorer if classifier_type is "heuristic", or if the LLM call fails, times out, or returns an unparseable response. + The outcome's `cause` reports which path actually classified the request. """ if self.config.classifier_type != "llm" or self.config.classifier_llm_config is None: - return self.classify(prompt, system_prompt) + tier, score, signals, cause = self._score_and_classify(prompt, system_prompt) + return ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause) try: tier = await self._classify_with_llm(prompt, system_prompt, request_kwargs) - return tier, 1.0, [f"llm-classifier:{tier.value}"] + return ClassificationOutcome( + tier=tier, score=None, signals=(f"llm-classifier:{tier.value}",), cause="llm_classifier" + ) except Exception as e: # noqa: BLE001 -- external LLM call can fail in many distinct ways (timeout, provider error, validation, parse error); any failure must fall back to the heuristic scorer verbose_router_logger.warning( f"ComplexityRouter: LLM classifier failed ({e}), falling back to heuristic scoring" ) - return self.classify(prompt, system_prompt) + tier, score, signals, cause = self._score_and_classify(prompt, system_prompt) + return ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause) async def _classify_with_llm( self, @@ -427,7 +549,17 @@ class ComplexityRouter(CustomLogger): # attributed to the calling key/team instead of being dropped. Excludes the # parent request's budget reservation, which the routed completion (not this # internal classifier call) is responsible for reconciling. - metadata = _classifier_call_metadata((request_kwargs or {}).get("litellm_metadata")) + request_metadata = (request_kwargs or {}).get("litellm_metadata") or (request_kwargs or {}).get("metadata") + metadata = _classifier_call_metadata(request_metadata) + turn_off_message_logging = _effective_turn_off_message_logging(request_kwargs) + + proxy_server_request = { + "body": { + "model": llm_config.model, + "messages": [{"role": "user", "content": classification_prompt}], + "response_format": type_to_response_format_param(TierClassification), + } + } response: ModelResponse = await self.litellm_router_instance.acompletion( model=llm_config.model, @@ -435,6 +567,8 @@ class ComplexityRouter(CustomLogger): response_format=TierClassification, timeout=llm_config.timeout_ms / 1000, metadata=metadata, + proxy_server_request=proxy_server_request, + turn_off_message_logging=turn_off_message_logging, ) content = response.choices[0].message.content if not content: @@ -675,16 +809,16 @@ class ComplexityRouter(CustomLogger): } return best_model - def _escalation_triggered(self, user_message: str) -> bool: - """Whether the prompt asks to escalate to a stronger model. + def _matched_escalation_keyword(self, user_message: str) -> str | None: + """The escalation keyword the prompt contains, or None when escalation is off. Matching is a case-sensitive substring test so the default "LITELLM ESCALATE" only fires on the deliberate, shouted form and not on incidental lowercase mentions of the word (e.g. "how do I escalate this ticket"). """ if not self.escalation_keywords: - return False - return any(keyword in user_message for keyword in self.escalation_keywords) + return None + return next((keyword for keyword in self.escalation_keywords if keyword in user_message), None) def _tier_for_model(self, model: str) -> ComplexityTier | None: """Return the most-severe configured tier whose pool contains this model.""" @@ -722,7 +856,7 @@ class ComplexityRouter(CustomLogger): return pinned_model return self.get_model_for_tier(escalated_tier) - def _lexical_tier_override(self, user_message: str) -> ComplexityTier | None: + def _lexical_tier_override(self, user_message: str) -> KeywordOverride | None: """When keyword_tier_rules match literally, the most-severe matched tier wins. Escalating to the highest tier (rather than the first rule in the list) keeps @@ -733,12 +867,15 @@ class ComplexityRouter(CustomLogger): if not rules: return None text = user_message.lower() - matched_tiers = [ - rule.tier for rule in rules if any(self._keyword_matches(text, keyword) for keyword in rule.keywords) + matches = [ + KeywordOverride(tier=rule.tier, matched_keyword=matched_keyword) + for rule in rules + if (matched_keyword := next((kw for kw in rule.keywords if self._keyword_matches(text, kw)), None)) + is not None ] - if not matched_tiers: + if not matches: return None - return max(matched_tiers, key=TIER_SEVERITY_ORDER.index) + return max(matches, key=lambda match: TIER_SEVERITY_ORDER.index(match.tier)) def _get_or_create_semantic_routelayer(self) -> SemanticRouter: """Build (once) a SemanticRouter with one route per tier, utterances = that tier's keywords.""" @@ -821,8 +958,16 @@ class ComplexityRouter(CustomLogger): # key/team budget. Key/team attribution fields are preserved for spend logging. metadata = _classifier_call_metadata(request_kwargs.get("metadata")) litellm_metadata = _classifier_call_metadata(request_kwargs.get("litellm_metadata")) + turn_off_message_logging = _effective_turn_off_message_logging(request_kwargs) + proxy_server_request = {"body": {"model": self.config.embedding_model, "input": [user_message]}} query_vector = ( - await encoder.aencode_queries([user_message], metadata=metadata, litellm_metadata=litellm_metadata) + await encoder.aencode_queries( + [user_message], + metadata=metadata, + litellm_metadata=litellm_metadata, + proxy_server_request=proxy_server_request, + turn_off_message_logging=turn_off_message_logging, + ) )[0] route_choice = await routelayer.acall(vector=query_vector) @@ -835,7 +980,7 @@ class ComplexityRouter(CustomLogger): except ValueError: return None - async def _resolve_keyword_tier_override(self, user_message: str, request_kwargs: dict) -> ComplexityTier | None: + async def _resolve_keyword_tier_override(self, user_message: str, request_kwargs: dict) -> KeywordOverride | None: """Resolve a keyword_tier_rule override, semantically or lexically per config. Returns None (no override -> fall through to the scorer) not only when no rule @@ -847,12 +992,17 @@ class ComplexityRouter(CustomLogger): if not self.config.semantic_keyword_matching: return self._lexical_tier_override(user_message) try: - return await self._semantic_tier_override(user_message, request_kwargs) + semantic_tier = await self._semantic_tier_override(user_message, request_kwargs) except Exception as e: # noqa: BLE001 -- embedding call can fail many ways (timeout, provider/network/parse error); any failure must fall back to scoring, never fail the request verbose_router_logger.warning( f"ComplexityRouter: semantic keyword matching failed ({e}), falling back to complexity scoring" ) return None + if semantic_tier is None: + return None + # A semantic match is a similarity hit against the rule's utterances, not a + # literal keyword, so there is no single matched keyword to report. + return KeywordOverride(tier=semantic_tier, matched_keyword=None) def _resolve_messages( self, @@ -971,6 +1121,7 @@ class ComplexityRouter(CustomLogger): pinned_model = await self.litellm_router_instance.cache.async_get_cache(key=cache_key) if isinstance(pinned_model, str): routed_model: str | None = pinned_model + pin_escalation_keyword: str | None = None if self.escalation_keywords: resolved_messages = self._resolve_messages(messages, request_kwargs) user_message = ( @@ -978,7 +1129,9 @@ class ComplexityRouter(CustomLogger): if resolved_messages else None ) - if user_message is not None and self._escalation_triggered(user_message): + if user_message is not None: + pin_escalation_keyword = self._matched_escalation_keyword(user_message) + if pin_escalation_keyword is not None: routed_model = self._escalated_pin(pinned_model) if routed_model is not None: # Refresh the TTL on every hit so an active session doesn't lose its @@ -996,7 +1149,8 @@ class ComplexityRouter(CustomLogger): kwargs_metadata = request_kwargs.setdefault("metadata", {}) if isinstance(kwargs_metadata, dict): kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = routed_model - cause = "session_affinity_escalation" if routed_model != pinned_model else "session_affinity_pin" + escalated = routed_model != pinned_model + cause: RoutingDecisionCause = "session_affinity_escalation" if escalated else "session_affinity_pin" verbose_router_logger.info( f"ComplexityRouter: routing decision cause={cause}, routed_model={routed_model}" ) @@ -1004,6 +1158,12 @@ class ComplexityRouter(CustomLogger): return PreRoutingHookResponse( model=routed_model, messages=messages if has_original_messages else None, + routing_decision=self._build_routing_decision( + routed_model=routed_model, + cause=cause, + escalation_keyword=pin_escalation_keyword, + escalated=escalated, + ), ) response = await self._classify_and_route( @@ -1074,29 +1234,45 @@ class ComplexityRouter(CustomLogger): return PreRoutingHookResponse( model=routed_model, messages=messages if has_original_messages else None, + routing_decision=self._build_routing_decision(routed_model=routed_model, cause="default_fallback"), ) - escalate = self._escalation_triggered(user_message) + escalation_keyword = self._matched_escalation_keyword(user_message) - override_tier = await self._resolve_keyword_tier_override(user_message, request_kwargs) - if override_tier is not None: - routed_tier = self._escalate_tier(override_tier) if escalate else override_tier + override = await self._resolve_keyword_tier_override(user_message, request_kwargs) + if override is not None: + routed_tier = self._escalate_tier(override.tier) if escalation_keyword is not None else override.tier + keyword_escalated = routed_tier != override.tier routed_model = await self._pick_model_for_tier(routed_tier, messages, resolved_messages, request_kwargs) - base_cause = "semantic_keyword_match" if self.config.semantic_keyword_matching else "literal_keyword_match" - cause = f"{base_cause}+escalation" if escalate else base_cause + keyword_cause: RoutingDecisionCause = ( + "semantic_keyword_match" if self.config.semantic_keyword_matching else "literal_keyword_match" + ) verbose_router_logger.info( - f"ComplexityRouter: routing decision cause={cause}, " + f"ComplexityRouter: routing decision cause={keyword_cause}, escalated={keyword_escalated}, " f"tier={routed_tier.value}, routed_model={routed_model}" ) return PreRoutingHookResponse( model=routed_model, messages=messages if has_original_messages else None, + routing_decision=self._build_routing_decision( + routed_model=routed_model, + cause=keyword_cause, + tier=routed_tier, + matched_keyword=override.matched_keyword, + escalation_keyword=escalation_keyword, + escalated=keyword_escalated, + ), ) - tier, score, signals = await self.aclassify(user_message, system_prompt, request_kwargs) - if escalate: + outcome = await self.aclassify(user_message, system_prompt, request_kwargs) + tier, score, signals = outcome.tier, outcome.score, outcome.signals + classified_tier = tier + if escalation_keyword is not None: tier = self._escalate_tier(tier) - signals = [*signals, "escalation"] + escalated = tier != classified_tier + if escalated: + signals = (*signals, "escalation") + score_repr = f"{score:.3f}" if score is not None else "n/a" if self.config.adaptive: routed_model = self._soft_floor_pick(tier, user_message, request_kwargs) adaptive = self._ensure_adaptive_router() @@ -1106,18 +1282,33 @@ class ComplexityRouter(CustomLogger): chosen_key = getattr(self, "_adaptive_chosen_model_key", "adaptive_router_chosen_model") kwargs_metadata[chosen_key] = routed_model verbose_router_logger.info( - f"ComplexityRouter[adaptive]: routing decision cause=complexity_scorer, " - f"tier={tier.value}, score={score:.3f}, " + f"ComplexityRouter[adaptive]: routing decision cause={outcome.cause}, " + f"tier={tier.value}, score={score_repr}, " f"signals={signals}, routed_model={routed_model}" ) else: routed_model = await self._pick_model_for_tier(tier, messages, resolved_messages, request_kwargs) verbose_router_logger.info( - f"ComplexityRouter: routing decision cause=complexity_scorer, tier={tier.value}, " - f"score={score:.3f}, signals={signals}, routed_model={routed_model}" + f"ComplexityRouter: routing decision cause={outcome.cause}, tier={tier.value}, " + f"score={score_repr}, signals={signals}, routed_model={routed_model}" ) + classifier_model = ( + self.config.classifier_llm_config.model + if outcome.cause == "llm_classifier" and self.config.classifier_llm_config is not None + else None + ) return PreRoutingHookResponse( model=routed_model, messages=messages if has_original_messages else None, + routing_decision=self._build_routing_decision( + routed_model=routed_model, + cause=outcome.cause, + tier=tier, + score=score, + signals=signals, + escalation_keyword=escalation_keyword, + escalated=escalated, + classifier_model=classifier_model, + ), ) diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 23476fe7dcc..ffe8245b012 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -73,13 +73,20 @@ class LowestLatencyLoggingHandler(CustomLogger): precise_minute = f"{current_date}-{current_hour}-{current_minute}" response_ms = end_time - start_time + if isinstance(response_ms, timedelta): + # normalize to float seconds up-front: non-chat responses + # (embeddings, speech, image) skip the ModelResponse branch + # below, and a raw timedelta appended to the latency list + # breaks JSON serialization when the router cache syncs to + # Redis (issue #33169) + response_ms = response_ms.total_seconds() time_to_first_token_response_time = None if kwargs.get("stream", None) is not None and kwargs["stream"] is True: # only log ttft for streaming request time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time - final_value: Union[float, timedelta] = response_ms + final_value: float = response_ms time_to_first_token: Optional[float] = None total_tokens = 0 @@ -89,15 +96,12 @@ class LowestLatencyLoggingHandler(CustomLogger): completion_tokens = _usage.completion_tokens total_tokens = _usage.total_tokens - # Handle both timedelta and float response times - if isinstance(response_ms, timedelta): - response_seconds = response_ms.total_seconds() - else: - response_seconds = response_ms + # response_ms is already normalized to float seconds above + response_seconds = response_ms - final_value = safe_divide_seconds(response_seconds, completion_tokens) - if final_value is not None: - final_value = float(final_value) + normalized_value = safe_divide_seconds(response_seconds, completion_tokens) + if normalized_value is not None: + final_value = float(normalized_value) else: final_value = response_seconds @@ -262,12 +266,19 @@ class LowestLatencyLoggingHandler(CustomLogger): precise_minute = f"{current_date}-{current_hour}-{current_minute}" response_ms = end_time - start_time + if isinstance(response_ms, timedelta): + # normalize to float seconds up-front: non-chat responses + # (embeddings, speech, image) skip the ModelResponse branch + # below, and a raw timedelta appended to the latency list + # breaks JSON serialization when the router cache syncs to + # Redis (issue #33169) + response_ms = response_ms.total_seconds() time_to_first_token_response_time = None if kwargs.get("stream", None) is not None and kwargs["stream"] is True: # only log ttft for streaming request time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time - final_value: Union[float, timedelta] = response_ms + final_value: float = response_ms total_tokens = 0 time_to_first_token: Optional[float] = None @@ -277,17 +288,14 @@ class LowestLatencyLoggingHandler(CustomLogger): completion_tokens = _usage.completion_tokens total_tokens = _usage.total_tokens - # Handle both timedelta and float response times - if isinstance(response_ms, timedelta): - response_seconds = response_ms.total_seconds() - else: - response_seconds = response_ms + # response_ms is already normalized to float seconds above + response_seconds = response_ms - final_value = safe_divide_seconds(response_seconds, completion_tokens) - if final_value is not None: - final_value = float(final_value) + normalized_value = safe_divide_seconds(response_seconds, completion_tokens) + if normalized_value is not None: + final_value = float(normalized_value) else: - final_value = response_ms + final_value = response_seconds if time_to_first_token_response_time is not None: if isinstance(time_to_first_token_response_time, timedelta): diff --git a/litellm/router_strategy/quality_router/quality_router.py b/litellm/router_strategy/quality_router/quality_router.py index 15c26f2c278..fd4a91d76bd 100644 --- a/litellm/router_strategy/quality_router/quality_router.py +++ b/litellm/router_strategy/quality_router/quality_router.py @@ -23,6 +23,7 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.router_strategy.complexity_router.complexity_router import ( ComplexityRouter, ) +from litellm.types.utils import StandardLoggingRoutingDecision from .config import QualityRouterConfig, RoutingPreferences @@ -357,6 +358,12 @@ class QualityRouter(CustomLogger): return PreRoutingHookResponse( model=self.config.default_model, messages=messages, + routing_decision=StandardLoggingRoutingDecision( + router_model_name=self.model_name, + router_type="quality", + routed_model=self.config.default_model, + cause="default_fallback", + ), ) # Try keyword override first — it short-circuits complexity classification. @@ -380,9 +387,20 @@ class QualityRouter(CustomLogger): "complexity_tier": None, }, ) + routing_decision = StandardLoggingRoutingDecision( + router_model_name=self.model_name, + router_type="quality", + routed_model=routed_model, + cause="keyword", + matched_keyword=matched_keyword, + ) + keyword_quality_tier = self._model_quality.get(routed_model) + if keyword_quality_tier is not None: + routing_decision["tier"] = str(keyword_quality_tier) return PreRoutingHookResponse( model=routed_model, messages=messages, + routing_decision=routing_decision, ) # No keyword match → complexity classification flow. @@ -419,4 +437,13 @@ class QualityRouter(CustomLogger): return PreRoutingHookResponse( model=routed_model, messages=messages, + routing_decision=StandardLoggingRoutingDecision( + router_model_name=self.model_name, + router_type="quality", + routed_model=routed_model, + cause="quality_tier", + tier=str(int(quality_tier)), + score=score, + signals=list(signals), + ), ) diff --git a/litellm/router_utils/auto_router_model_naming.py b/litellm/router_utils/auto_router_model_naming.py new file mode 100644 index 00000000000..72865501030 --- /dev/null +++ b/litellm/router_utils/auto_router_model_naming.py @@ -0,0 +1,101 @@ +"""Naming contract for strategy-router (auto-router) pseudo-models. + +A deployment whose ``litellm_params.model`` starts with ``auto_router/`` does not +name a provider model; the string is the discriminator that selects which +pre-routing strategy owns the deployment. This module is the single source of +truth for classifying that string (``Router._is_*_router_deployment`` delegates +here) and for checking that a client-supplied write leaves the deployment +coherent, so management endpoints can reject corruption with a 400 instead of +the router silently dropping the deployment at load time under +``ignore_invalid_deployments``. +""" + +from typing import Literal, Mapping + +AUTO_ROUTER_MODEL_PREFIX = "auto_router/" + +StrategyRouterKind = Literal["semantic", "complexity", "adaptive", "quality"] + +STRATEGY_ROUTER_PARAM_FIELDS: frozenset[str] = frozenset( + { + "auto_router_config", + "auto_router_config_path", + "auto_router_default_model", + "auto_router_embedding_model", + "complexity_router_config", + "complexity_router_default_model", + "adaptive_router_config", + "quality_router_config", + "quality_router_default_model", + } +) + +_REQUIRED_FIELD_GROUPS: Mapping[StrategyRouterKind, tuple[tuple[str, ...], ...]] = { + "semantic": ( + ("auto_router_config", "auto_router_config_path"), + ("auto_router_default_model",), + ("auto_router_embedding_model",), + ), + "complexity": (("complexity_router_config", "complexity_router_default_model"),), + "adaptive": (("adaptive_router_config",),), + "quality": (("quality_router_config", "quality_router_default_model"),), +} + + +def classify_strategy_router_model(model: str) -> StrategyRouterKind | None: + """Classify a ``litellm_params.model`` string the way the Router does. + + Returns None for regular provider models. Mirrors Router registration + exactly: reserved names are matched by prefix, everything else under + ``auto_router/`` is a semantic router. + """ + if not model.startswith(AUTO_ROUTER_MODEL_PREFIX): + return None + remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :] + if remainder.startswith("complexity_router"): + return "complexity" + if remainder.startswith("adaptive_router"): + return "adaptive" + if remainder.startswith("quality_router"): + return "quality" + return "semantic" + + +def validate_strategy_router_model_write(model: str, present_fields: frozenset[str]) -> str | None: + """Check that writing ``model`` leaves a deployment the router can load. + + ``present_fields`` is the set of strategy-router param fields that are + non-None on the deployment after the write (stored fields merged with the + incoming ones). Returns a human-readable violation, or None when coherent. + """ + kind = classify_strategy_router_model(model) + if kind is None: + offending = sorted(present_fields & STRATEGY_ROUTER_PARAM_FIELDS) + if offending: + return ( + f"litellm_params.model='{model}' does not start with '{AUTO_ROUTER_MODEL_PREFIX}' but the " + f"deployment carries auto-router settings ({', '.join(offending)}), so the router could not " + f"load it. Keep the '{AUTO_ROUTER_MODEL_PREFIX}' prefix; to change the name clients call, " + "edit the public model_name instead." + ) + return None + remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :] + if remainder.startswith(AUTO_ROUTER_MODEL_PREFIX): + return ( + f"litellm_params.model='{model}' repeats the '{AUTO_ROUTER_MODEL_PREFIX}' prefix, so the router " + f"could not load it. Use '{remainder}'; to change the name clients call, edit the public " + "model_name instead." + ) + if not remainder: + return ( + f"litellm_params.model='{model}' is missing the router name after the '{AUTO_ROUTER_MODEL_PREFIX}' prefix." + ) + missing = tuple( + " or ".join(group) for group in _REQUIRED_FIELD_GROUPS[kind] if not any(f in present_fields for f in group) + ) + if missing: + return ( + f"litellm_params.model='{model}' selects the {kind} router, which requires " + f"{'; '.join(missing)} in litellm_params." + ) + return None diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index a324e71e289..af419d8cb6f 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -1043,6 +1043,7 @@ class GuardrailEventHooks(str, Enum): logging_only = "logging_only" pre_mcp_call = "pre_mcp_call" during_mcp_call = "during_mcp_call" + post_mcp_call = "post_mcp_call" realtime_input_transcription = "realtime_input_transcription" diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 318ba1f5956..905efc84b5a 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -190,6 +190,7 @@ class UserAPIKeyLabelNames(Enum): ORG_ALIAS = "org_alias" MCP_TOOL_NAME = "mcp_tool_name" MCP_SERVER_NAME = "mcp_server_name" + SERVICE_TIER = "service_tier" DEFINED_PROMETHEUS_METRICS = Literal[ @@ -286,6 +287,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.USER.value, UserAPIKeyLabelNames.MODEL_ID.value, UserAPIKeyLabelNames.API_PROVIDER.value, + UserAPIKeyLabelNames.SERVICE_TIER.value, ] litellm_llm_api_time_to_first_token_metric = [ @@ -299,6 +301,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.USER.value, UserAPIKeyLabelNames.MODEL_ID.value, UserAPIKeyLabelNames.API_PROVIDER.value, + UserAPIKeyLabelNames.SERVICE_TIER.value, ] litellm_request_total_latency_metric = [ @@ -312,6 +315,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, UserAPIKeyLabelNames.MODEL_ID.value, UserAPIKeyLabelNames.API_PROVIDER.value, + UserAPIKeyLabelNames.SERVICE_TIER.value, ] litellm_request_queue_time_seconds = [ @@ -453,6 +457,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.REQUESTED_MODEL.value, UserAPIKeyLabelNames.MODEL_ID.value, UserAPIKeyLabelNames.API_PROVIDER.value, + UserAPIKeyLabelNames.SERVICE_TIER.value, ] litellm_input_tokens_metric = [ @@ -878,6 +883,7 @@ class UserAPIKeyLabelValues: org_alias: Optional[str] = None mcp_tool_name: Optional[str] = None mcp_server_name: Optional[str] = None + service_tier: Optional[str] = None # Added for test compatibility. def __init__(self, **kwargs: Any) -> None: diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index fb3ddeebf52..da1ff7eda67 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -16,7 +16,7 @@ GeminiEmbeddingInput = Union[EmbeddingInput, List[List[str]]] class FunctionResponse(TypedDict, total=False): # `id` correlates this response with the originating `functionCall` part. - # Supported on Google AI Studio Gemini 3.5+; Vertex AI rejects this field. + # Supported on Gemini 3+; older Gemini models reject this field. id: str name: Required[str] response: Optional[dict] @@ -24,8 +24,8 @@ class FunctionResponse(TypedDict, total=False): class FunctionCall(TypedDict, total=False): - # `id` correlates the corresponding `functionResponse` on Google AI Studio - # Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field. + # `id` correlates the corresponding `functionResponse` on Gemini 3+. + # Older Gemini models omit/reject this field. id: str name: Required[str] args: Optional[dict] @@ -58,8 +58,8 @@ class PartType(TypedDict, total=False): class HttpxFunctionCall(TypedDict, total=False): - # `id` correlates the corresponding `functionResponse` on Google AI Studio - # Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field. + # `id` correlates the corresponding `functionResponse` on Gemini 3+. + # Older Gemini models omit/reject this field. id: str name: Required[str] args: dict diff --git a/litellm/types/proxy/management_endpoints/scim_v2.py b/litellm/types/proxy/management_endpoints/scim_v2.py index 8c434481975..e7f5c85e6c6 100644 --- a/litellm/types/proxy/management_endpoints/scim_v2.py +++ b/litellm/types/proxy/management_endpoints/scim_v2.py @@ -18,6 +18,9 @@ SCIM_ENTERPRISE_METADATA_KEY = "scim_enterprise" SCIM_ENTITLEMENTS_METADATA_KEY = "scim_entitlements" SCIM_ROLES_METADATA_KEY = "scim_roles" +SCIM_MANAGED_TEAM_METADATA_KEY = "scim_managed" +SCIM_TEAM_DATA_METADATA_KEY = "scim_data" + class LiteLLM_UserScimMetadata(BaseModel): """ @@ -131,6 +134,15 @@ class SCIMUser(SCIMResource): class SCIMMember(BaseModel): value: str # User ID display: Optional[str] = None # Username or email + type: str | None = None + + @field_validator("type", mode="before") + @classmethod + def normalize_type(cls, v: object) -> str | None: + """Anything that is not a string carries no canonical type, and rejecting the + request over it would be a regression: before this field existed the value was + parsed away silently.""" + return v if isinstance(v, str) else None class SCIMGroup(SCIMResource): diff --git a/litellm/types/router.py b/litellm/types/router.py index 28e4a8272e8..837a93367a2 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -28,7 +28,7 @@ from .completion import CompletionRequest from .embedding import EmbeddingRequest from .llms.openai import OpenAIFileObject from .search import SearchProvider -from .utils import CustomPricingLiteLLMParams, ModelResponse +from .utils import CustomPricingLiteLLMParams, ModelResponse, StandardLoggingRoutingDecision class ConfigurableClientsideParamsCustomAuth(TypedDict): @@ -839,6 +839,7 @@ class PreRoutingHookResponse(BaseModel): model: str messages: Optional[List[Dict[str, Any]]] + routing_decision: StandardLoggingRoutingDecision | None = None _PreRoutingStrategyT_co = TypeVar("_PreRoutingStrategyT_co", covariant=True) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index e4dfac48141..9df44c6202c 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -10,6 +10,7 @@ from typing import ( Literal, Mapping, Optional, + Sequence, Union, get_args, ) @@ -2674,6 +2675,73 @@ class StandardLoggingPromptManagementMetadata(TypedDict): prompt_integration: str +class StandardLoggingRoutingDecisionTierBoundaries(TypedDict): + """Snapshot of the complexity scorer's tier boundaries at decision time, so a + historical spend log row stays explainable after the router config changes.""" + + simple_medium: float + medium_complex: float + complex_reasoning: float + + +RoutingDecisionCause = Literal[ + "heuristic_scorer", + # The scorer found 2+ reasoning markers and forced REASONING regardless of score. + # A distinct cause rather than a marker inside `signals`, because it is the fact + # that tells a reader the score did NOT choose the tier; encoding it as free text + # meant anything that filtered `signals` silently changed what the row claimed. + "reasoning_override", + "llm_classifier", + "literal_keyword_match", + "semantic_keyword_match", + "session_affinity_pin", + "session_affinity_escalation", + "default_fallback", + "keyword", + "quality_tier", + "bandit", +] + + +class StandardLoggingRoutingDecision(TypedDict, total=False): + """Per-request provenance for a pre-routing strategy (auto-router) decision.""" + + router_model_name: str + router_type: Literal["complexity", "adaptive", "quality"] + routed_model: str + cause: RoutingDecisionCause + tier: str + request_type: str + score: float + signals: Sequence[str] + matched_keyword: str + escalation_keyword: str + classifier_model: str + escalated: bool + tier_boundaries: StandardLoggingRoutingDecisionTierBoundaries + + +# Fields whose values quote the caller's prompt. Dropped when an operator turns message +# logging off. Every other field aggregates the prompt without reproducing it and is kept, +# so a redacted row stays explainable. `test_every_routing_decision_field_is_classified` +# fails if a field is added to the record without being placed in one set or the other. +PROMPT_QUOTING_ROUTING_DECISION_FIELDS: FrozenSet[str] = frozenset({"signals", "matched_keyword", "escalation_keyword"}) +DERIVED_ROUTING_DECISION_FIELDS: FrozenSet[str] = frozenset( + { + "router_model_name", + "router_type", + "routed_model", + "cause", + "tier", + "request_type", + "score", + "classifier_model", + "escalated", + "tier_boundaries", + } +) + + class StandardLoggingMetadata(StandardLoggingUserAPIKeyMetadata): """ Specific metadata k,v pairs logged to integration for easier cost tracking and prompt management @@ -2687,6 +2755,7 @@ class StandardLoggingMetadata(StandardLoggingUserAPIKeyMetadata): prompt_management_metadata: Optional[StandardLoggingPromptManagementMetadata] mcp_tool_call_metadata: Optional[StandardLoggingMCPToolCall] vector_store_request_metadata: Optional[List[StandardLoggingVectorStoreRequest]] + routing_decision: StandardLoggingRoutingDecision | None applied_guardrails: Optional[List[str]] usage_object: Optional[dict] cold_storage_object_key: Optional[str] # S3/GCS object key for cold storage retrieval diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index db28118d52b..c4628fecdb8 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3454,22 +3454,16 @@ }, "azure_ai/gpt-5.4-mini": { "cache_read_input_token_cost": 7.5e-08, - "cache_read_input_token_cost_above_272k_tokens": 1.5e-07, "cache_read_input_token_cost_priority": 1.5e-07, - "cache_read_input_token_cost_above_272k_tokens_priority": 3e-07, "input_cost_per_token": 7.5e-07, - "input_cost_per_token_above_272k_tokens": 1.5e-06, "input_cost_per_token_priority": 1.5e-06, - "input_cost_per_token_above_272k_tokens_priority": 3e-06, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.5e-06, - "output_cost_per_token_above_272k_tokens": 6.75e-06, "output_cost_per_token_priority": 9e-06, - "output_cost_per_token_above_272k_tokens_priority": 1.35e-05, "source": "https://ai.azure.com/catalog/models/gpt-5.4-mini", "supported_endpoints": [ "/v1/chat/completions", @@ -3500,22 +3494,16 @@ }, "azure_ai/gpt-5.4-mini-2026-03-17": { "cache_read_input_token_cost": 7.5e-08, - "cache_read_input_token_cost_above_272k_tokens": 1.5e-07, "cache_read_input_token_cost_priority": 1.5e-07, - "cache_read_input_token_cost_above_272k_tokens_priority": 3e-07, "input_cost_per_token": 7.5e-07, - "input_cost_per_token_above_272k_tokens": 1.5e-06, "input_cost_per_token_priority": 1.5e-06, - "input_cost_per_token_above_272k_tokens_priority": 3e-06, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.5e-06, - "output_cost_per_token_above_272k_tokens": 6.75e-06, "output_cost_per_token_priority": 9e-06, - "output_cost_per_token_above_272k_tokens_priority": 1.35e-05, "source": "https://ai.azure.com/catalog/models/gpt-5.4-mini", "supported_endpoints": [ "/v1/chat/completions", @@ -3546,22 +3534,16 @@ }, "azure_ai/gpt-5.4-nano": { "cache_read_input_token_cost": 2e-08, - "cache_read_input_token_cost_above_272k_tokens": 4e-08, "cache_read_input_token_cost_priority": 4e-08, - "cache_read_input_token_cost_above_272k_tokens_priority": 8e-08, "input_cost_per_token": 2e-07, - "input_cost_per_token_above_272k_tokens": 4e-07, "input_cost_per_token_priority": 4e-07, - "input_cost_per_token_above_272k_tokens_priority": 8e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.25e-06, - "output_cost_per_token_above_272k_tokens": 1.875e-06, "output_cost_per_token_priority": 2.5e-06, - "output_cost_per_token_above_272k_tokens_priority": 3.75e-06, "source": "https://ai.azure.com/catalog/models/gpt-5.4-nano", "supported_endpoints": [ "/v1/chat/completions", @@ -3592,22 +3574,16 @@ }, "azure_ai/gpt-5.4-nano-2026-03-17": { "cache_read_input_token_cost": 2e-08, - "cache_read_input_token_cost_above_272k_tokens": 4e-08, "cache_read_input_token_cost_priority": 4e-08, - "cache_read_input_token_cost_above_272k_tokens_priority": 8e-08, "input_cost_per_token": 2e-07, - "input_cost_per_token_above_272k_tokens": 4e-07, "input_cost_per_token_priority": 4e-07, - "input_cost_per_token_above_272k_tokens_priority": 8e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.25e-06, - "output_cost_per_token_above_272k_tokens": 1.875e-06, "output_cost_per_token_priority": 2.5e-06, - "output_cost_per_token_above_272k_tokens_priority": 3.75e-06, "source": "https://ai.azure.com/catalog/models/gpt-5.4-nano", "supported_endpoints": [ "/v1/chat/completions", @@ -7201,7 +7177,7 @@ "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 7.5e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7236,7 +7212,7 @@ "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 7.5e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7271,7 +7247,7 @@ "cache_read_input_token_cost": 2e-08, "input_cost_per_token": 2e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7306,7 +7282,7 @@ "cache_read_input_token_cost": 2e-08, "input_cost_per_token": 2e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -13567,6 +13543,56 @@ } ] }, + "dashscope/qwen3.7-max": { + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token": 2.5e-06, + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 7.5e-06, + "source": "https://www.alibabacloud.com/help/en/model-studio/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "dashscope/qwen3.7-plus": { + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "source": "https://www.alibabacloud.com/help/en/model-studio/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tiered_pricing": [ + { + "cache_read_input_token_cost": 8e-08, + "input_cost_per_token": 4e-07, + "output_cost_per_token": 1.6e-06, + "range": [ + 0, + 256000.0 + ] + }, + { + "cache_read_input_token_cost": 2.4e-07, + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 4.8e-06, + "range": [ + 256000.0, + 1000000.0 + ] + } + ] + }, "dashscope/qwq-plus": { "input_cost_per_token": 8e-07, "litellm_provider": "dashscope", @@ -24314,7 +24340,7 @@ "input_cost_per_token_batches": 3.75e-07, "input_cost_per_token_priority": 1.5e-06, "litellm_provider": "openai", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -24360,7 +24386,7 @@ "input_cost_per_token_batches": 3.75e-07, "input_cost_per_token_priority": 1.5e-06, "litellm_provider": "openai", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -24404,7 +24430,7 @@ "input_cost_per_token_flex": 1e-07, "input_cost_per_token_batches": 1e-07, "litellm_provider": "openai", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -24447,7 +24473,7 @@ "input_cost_per_token_flex": 1e-07, "input_cost_per_token_batches": 1e-07, "litellm_provider": "openai", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json new file mode 100644 index 00000000000..988f56655fc --- /dev/null +++ b/model_prices_and_context_window.schema.json @@ -0,0 +1,742 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "title": "LiteLLM model_prices_and_context_window.json", + "description": "Schema for LiteLLM's model price and context window registry (https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json). Every top-level key except 'sample_spec' and 'fallback_generalizations' is a model id, optionally prefixed with its provider (e.g. 'azure/gpt-5.4'), mapping to a model entry. All costs are USD per unit. New optional fields are added regularly, so consumers should ignore unknown fields rather than reject them.", + "type": "object", + "properties": { + "sample_spec": { + "type": "object", + "description": "Documentation placeholder illustrating the entry shape; not a real model and not schema-conformant (several values are prose)." + }, + "fallback_generalizations": { + "type": "object", + "description": "Regex rules that generalize unknown model ids to known families; not a model entry.", + "properties": { + "rules": { + "type": "array", + "items": { + "type": "object", + "properties": { + "name": { + "type": "string" + }, + "pattern": { + "type": "string" + }, + "description": { + "type": "string" + } + }, + "required": [ + "name", + "pattern" + ], + "additionalProperties": true + } + } + }, + "additionalProperties": false + } + }, + "additionalProperties": { + "$ref": "#/$defs/modelEntry" + }, + "$defs": { + "modelEntry": { + "type": "object", + "description": "Pricing, limits, and capability flags for one model. Fields other than litellm_provider are optional; boolean capability flags are simply omitted when unknown or false.", + "required": [ + "litellm_provider" + ], + "properties": { + "annotation_cost_per_page": { + "type": "number", + "minimum": 0 + }, + "audio_transcription_config": { + "type": "string" + }, + "bedrock_converse_supports_strict_tools": { + "type": "boolean" + }, + "bedrock_output_config_effort_ceiling": { + "type": "string", + "description": "Highest reasoning effort the Bedrock output_config accepts for this model.", + "enum": [ + "low", + "medium", + "high", + "max", + "xhigh" + ] + }, + "cache_creation_input_audio_token_cost": { + "type": "number", + "minimum": 0 + }, + "cache_creation_input_token_cost": { + "type": "number", + "minimum": 0, + "description": "USD per token written to the provider's prompt cache." + }, + "cache_creation_input_token_cost_above_1hr": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "cache_creation_input_token_cost_above_200k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "cache_creation_input_token_cost_above_272k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "cache_creation_input_token_cost_flex": { + "type": "number", + "minimum": 0, + "description": "Flex service-tier rate for the same-named base field." + }, + "cache_creation_input_token_cost_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "cache_read_input_audio_token_cost": { + "type": "number", + "minimum": 0 + }, + "cache_read_input_token_cost": { + "type": "number", + "minimum": 0, + "description": "USD per prompt token served from the provider's prompt cache." + }, + "cache_read_input_token_cost_above_200k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "cache_read_input_token_cost_above_200k_tokens_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "cache_read_input_token_cost_above_272k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "cache_read_input_token_cost_above_272k_tokens_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "cache_read_input_token_cost_above_512k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "cache_read_input_token_cost_flex": { + "type": "number", + "minimum": 0, + "description": "Flex service-tier rate for the same-named base field." + }, + "cache_read_input_token_cost_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "citation_cost_per_token": { + "type": "number", + "minimum": 0 + }, + "code_interpreter_cost_per_session": { + "type": "number", + "minimum": 0 + }, + "comment": { + "type": "string" + }, + "deprecation_date": { + "type": "string", + "description": "Date the provider deprecates the model, YYYY-MM-DD.", + "format": "date", + "pattern": "^\\d{4}-(0[1-9]|1[0-2])-(0[1-9]|[12]\\d|3[01])$" + }, + "gemini_audio_only_live": { + "type": "boolean" + }, + "gemini_native_audio": { + "type": "boolean" + }, + "input_cost_per_audio_per_second": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_audio_per_second_above_128k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_cost_per_audio_token": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_audio_token_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "input_cost_per_character": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_character_above_128k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_cost_per_image": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_image_above_128k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_cost_per_image_token": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_pixel": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_query": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_request": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_second": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_token": { + "type": "number", + "minimum": 0, + "description": "USD per prompt token." + }, + "input_cost_per_token_above_128k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_cost_per_token_above_200k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_cost_per_token_above_200k_tokens_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "input_cost_per_token_above_256k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_cost_per_token_above_272k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_cost_per_token_above_272k_tokens_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "input_cost_per_token_above_512k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_cost_per_token_batches": { + "type": "number", + "minimum": 0, + "description": "USD per prompt token via the provider's batch API." + }, + "input_cost_per_token_cache_hit": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_token_flex": { + "type": "number", + "minimum": 0, + "description": "Flex service-tier rate for the same-named base field." + }, + "input_cost_per_token_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "input_cost_per_video_per_second": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_video_per_second_above_128k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_cost_per_video_per_second_above_15s_interval": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_cost_per_video_per_second_above_8s_interval": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "input_dbu_cost_per_token": { + "type": "number", + "minimum": 0 + }, + "litellm_provider": { + "type": "string", + "description": "LiteLLM provider slug; one of https://docs.litellm.ai/docs/providers." + }, + "max_input_tokens": { + "type": "integer", + "minimum": 0, + "description": "Maximum prompt/context tokens the model accepts." + }, + "max_output_tokens": { + "type": "integer", + "minimum": 0, + "description": "Maximum tokens the model can generate in one response." + }, + "max_tokens": { + "type": "integer", + "minimum": 0, + "description": "Legacy field: max output tokens if the provider specifies it, else max input tokens." + }, + "metadata": { + "type": "object", + "description": "Free-form notes about the entry (e.g. pricing derivation)." + }, + "mode": { + "type": "string", + "description": "Primary API surface / task type of the model.", + "enum": [ + "audio_speech", + "audio_transcription", + "chat", + "completion", + "embedding", + "image_edit", + "image_generation", + "moderation", + "ocr", + "realtime", + "rerank", + "responses", + "search", + "vector_store", + "video_generation" + ] + }, + "ocr_cost_per_credit": { + "type": "number", + "minimum": 0 + }, + "ocr_cost_per_page": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_audio_token": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_character": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_character_above_128k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "output_cost_per_image": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_image_token": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_pixel": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_reasoning_token": { + "type": "number", + "minimum": 0, + "description": "USD per reasoning/thinking token, when billed separately." + }, + "output_cost_per_second": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_second_1080p": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_token": { + "type": "number", + "minimum": 0, + "description": "USD per generated token." + }, + "output_cost_per_token_above_128k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "output_cost_per_token_above_200k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "output_cost_per_token_above_200k_tokens_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "output_cost_per_token_above_256k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "output_cost_per_token_above_272k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "output_cost_per_token_above_272k_tokens_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "output_cost_per_token_above_512k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, + "output_cost_per_token_batches": { + "type": "number", + "minimum": 0, + "description": "USD per generated token via the provider's batch API." + }, + "output_cost_per_token_flex": { + "type": "number", + "minimum": 0, + "description": "Flex service-tier rate for the same-named base field." + }, + "output_cost_per_token_priority": { + "type": "number", + "minimum": 0, + "description": "Priority service-tier rate for the same-named base field." + }, + "output_cost_per_video_per_second": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_video_token": { + "type": "number", + "minimum": 0 + }, + "output_dbu_cost_per_token": { + "type": "number", + "minimum": 0 + }, + "output_vector_size": { + "type": "integer", + "minimum": 0, + "description": "Embedding dimension for embedding models." + }, + "prompt_cache_min_tokens": { + "type": "integer", + "minimum": 0, + "description": "Smallest prefix the provider will actually cache; absent means the provider default applies." + }, + "provider_specific_entry": { + "type": "object", + "description": "Provider-internal routing hints (e.g. bedrock_invocation_schema)." + }, + "regional_processing_uplift_multiplier_eu": { + "type": "number", + "minimum": 1, + "description": "Multiplier applied to all token costs for EU data residency (e.g. 1.10 = +10%)." + }, + "regional_processing_uplift_multiplier_us": { + "type": "number", + "minimum": 1, + "description": "Multiplier applied to all token costs for US data residency (e.g. 1.10 = +10%)." + }, + "rpm": { + "type": "integer", + "minimum": 0, + "description": "Provider default requests-per-minute limit." + }, + "search_context_cost_per_query": { + "type": "object", + "description": "USD cost per web search query, keyed by search context size.", + "properties": { + "search_context_size_low": { + "type": "number", + "minimum": 0 + }, + "search_context_size_medium": { + "type": "number", + "minimum": 0 + }, + "search_context_size_high": { + "type": "number", + "minimum": 0 + } + }, + "additionalProperties": false + }, + "source": { + "type": "string", + "description": "URL of the provider pricing/model page this entry was taken from." + }, + "supported_endpoints": { + "type": "array", + "description": "OpenAI-style API routes this model can be called through, e.g. /v1/chat/completions.", + "items": { + "type": "string" + } + }, + "supported_modalities": { + "type": "array", + "description": "Input modalities the model accepts.", + "items": { + "type": "string", + "enum": [ + "text", + "image", + "audio", + "video" + ] + } + }, + "supported_output_modalities": { + "type": "array", + "description": "Output modalities the model can produce.", + "items": { + "type": "string", + "enum": [ + "text", + "image", + "audio", + "video", + "code" + ] + } + }, + "supported_regions": { + "type": "array", + "description": "Cloud regions the model is available in ('global' or region ids).", + "items": { + "type": "string" + } + }, + "supports_adaptive_thinking": { + "type": "boolean" + }, + "supports_assistant_prefill": { + "type": "boolean" + }, + "supports_audio_input": { + "type": "boolean" + }, + "supports_audio_output": { + "type": "boolean" + }, + "supports_computer_use": { + "type": "boolean" + }, + "supports_embedding_image_input": { + "type": "boolean" + }, + "supports_function_calling": { + "type": "boolean" + }, + "supports_image_input": { + "type": "boolean" + }, + "supports_image_size": { + "type": "boolean" + }, + "supports_low_reasoning_effort": { + "type": "boolean" + }, + "supports_max_reasoning_effort": { + "type": "boolean" + }, + "supports_mid_conversation_system": { + "type": "boolean" + }, + "supports_minimal_reasoning_effort": { + "type": "boolean" + }, + "supports_multimodal": { + "type": "boolean" + }, + "supports_native_streaming": { + "type": "boolean" + }, + "supports_native_structured_output": { + "type": "boolean" + }, + "supports_none_reasoning_effort": { + "type": "boolean" + }, + "supports_nova_canvas_image_edit": { + "type": "boolean" + }, + "supports_output_config": { + "type": "boolean" + }, + "supports_parallel_function_calling": { + "type": "boolean" + }, + "supports_parallel_tool_use_config": { + "type": "boolean" + }, + "supports_pdf_input": { + "type": "boolean" + }, + "supports_prompt_caching": { + "type": "boolean" + }, + "supports_reasoning": { + "type": "boolean" + }, + "supports_response_schema": { + "type": "boolean" + }, + "supports_sampling_params": { + "type": "boolean" + }, + "supports_speed": { + "type": "boolean" + }, + "supports_system_messages": { + "type": "boolean" + }, + "supports_tool_choice": { + "type": "boolean" + }, + "supports_url_context": { + "type": "boolean" + }, + "supports_video_input": { + "type": "boolean" + }, + "supports_vision": { + "type": "boolean" + }, + "supports_web_search": { + "type": "boolean" + }, + "supports_xhigh_reasoning_effort": { + "type": "boolean" + }, + "tiered_pricing": { + "type": "array", + "description": "Context-length or result-count tiered rates; each tier's costs apply within its range.", + "items": { + "type": "object", + "properties": { + "range": { + "type": "array", + "description": "[min, max] prompt-token span this tier applies to.", + "items": { + "type": "number", + "minimum": 0 + }, + "minItems": 2, + "maxItems": 2 + }, + "max_results_range": { + "type": "array", + "description": "[min, max] result-count span this tier applies to (search models).", + "items": { + "type": "number", + "minimum": 0 + }, + "minItems": 2, + "maxItems": 2 + }, + "input_cost_per_token": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_token": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_reasoning_token": { + "type": "number", + "minimum": 0 + }, + "cache_read_input_token_cost": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_query": { + "type": "number", + "minimum": 0 + } + }, + "additionalProperties": false + } + }, + "tpm": { + "type": "integer", + "minimum": 0, + "description": "Provider default tokens-per-minute limit." + }, + "use_openai_responses_path": { + "type": "boolean" + }, + "uses_embed_content": { + "type": "boolean" + }, + "web_search_billing_unit": { + "type": "string", + "description": "Whether web search is billed per query or per prompt.", + "enum": [ + "per_query", + "per_prompt" + ] + } + }, + "additionalProperties": true + } + } +} diff --git a/pyproject.toml b/pyproject.toml index e15bf1351dd..93fb32da464 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm" -version = "1.95.0" +version = "1.96.0" description = "Library to easily interface with LLM API providers" readme = "README.md" requires-python = ">=3.10, <3.15" @@ -302,7 +302,7 @@ members = ["enterprise", "litellm-proxy-extras"] profile = "black" [tool.commitizen] -version = "1.95.0" +version = "1.96.0" version_files = [ "pyproject.toml:^version", ] diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index addee5fc68a..dfdc4efe800 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -1,6 +1,6 @@ { "ANN001": { - "limit": 3142 + "limit": 3118 }, "ANN002": { "limit": 69 @@ -24,7 +24,7 @@ "limit": 130 }, "ANN401": { - "limit": 2015 + "limit": 2010 }, "ASYNC230": { "limit": 14 @@ -33,7 +33,7 @@ "limit": 4 }, "B006": { - "limit": 190 + "limit": 188 }, "B008": { "limit": 505 @@ -42,7 +42,7 @@ "limit": 84 }, "B010": { - "limit": 197 + "limit": 194 }, "B018": { "limit": 5 @@ -60,7 +60,7 @@ "limit": 4 }, "BLE001": { - "limit": 2902 + "limit": 2899 }, "C401": { "limit": 11 @@ -81,7 +81,7 @@ "limit": 4 }, "C901": { - "limit": 316 + "limit": 314 }, "D419": { "limit": 9 @@ -135,7 +135,7 @@ "limit": 30 }, "PERF401": { - "limit": 144 + "limit": 142 }, "PERF402": { "limit": 9 @@ -180,7 +180,7 @@ "limit": 34 }, "PLR1714": { - "limit": 265 + "limit": 261 }, "PLR1730": { "limit": 10 @@ -189,7 +189,7 @@ "limit": 4 }, "PLW0127": { - "limit": 44 + "limit": 43 }, "PLW0133": { "limit": 4 @@ -222,7 +222,7 @@ "limit": 38 }, "RET504": { - "limit": 717 + "limit": 716 }, "RUF010": { "limit": 874 @@ -261,7 +261,7 @@ "limit": 24 }, "SIM101": { - "limit": 63 + "limit": 61 }, "SIM102": { "limit": 324 @@ -273,7 +273,7 @@ "limit": 6 }, "SIM114": { - "limit": 113 + "limit": 111 }, "SIM115": { "limit": 5 @@ -288,7 +288,7 @@ "limit": 4 }, "SIM210": { - "limit": 12 + "limit": 11 }, "SIM211": { "limit": 4 @@ -309,22 +309,22 @@ "limit": 2652 }, "TRY002": { - "limit": 548 + "limit": 547 }, "TRY004": { "limit": 98 }, "TRY201": { - "limit": 424 + "limit": 420 }, "TRY203": { - "limit": 123 + "limit": 121 }, "TRY300": { - "limit": 883 + "limit": 879 }, "UP006": { - "limit": 12147 + "limit": 12138 }, "UP007": { "limit": 2526 @@ -348,7 +348,7 @@ "limit": 5 }, "UP032": { - "limit": 629 + "limit": 626 }, "UP034": { "limit": 4 @@ -363,6 +363,6 @@ "limit": 105 }, "UP045": { - "limit": 17824 + "limit": 17805 } } diff --git a/tests/batches_tests/test_batch_custom_pricing.py b/tests/batches_tests/test_batch_custom_pricing.py index 3dc1d116e8d..c2159b564a8 100644 --- a/tests/batches_tests/test_batch_custom_pricing.py +++ b/tests/batches_tests/test_batch_custom_pricing.py @@ -12,8 +12,7 @@ import litellm import pytest from litellm.batches.batch_utils import ( - _batch_cost_calculator, - _get_batch_job_cost_from_file_content, + _aggregate_batch_cost_usage_models, calculate_batch_cost_and_usage, ) from litellm.cost_calculator import batch_cost_calculator @@ -113,28 +112,12 @@ def test_batch_cost_calculator_uses_custom_model_info(): ), f"Expected completion cost {expected_completion}, got {completion_cost}" -def test_get_batch_job_cost_from_file_content_uses_custom_model_info(): - """_get_batch_job_cost_from_file_content should thread model_info to completion_cost.""" +def test_aggregate_batch_cost_uses_custom_model_info(): + """_aggregate_batch_cost_usage_models should thread model_info to batch_cost_calculator.""" file_content = [_make_batch_output_line(prompt_tokens=10, completion_tokens=5)] - cost = _get_batch_job_cost_from_file_content( - file_content_dictionary=file_content, - custom_llm_provider="openai", - model_info=CUSTOM_MODEL_INFO, - ) - - expected = (10 * 0.00125) + (5 * 0.005) - assert cost == pytest.approx( - expected - ), f"Expected total cost {expected}, got {cost}" - - -def test_batch_cost_calculator_func_uses_custom_model_info(): - """_batch_cost_calculator should thread model_info.""" - file_content = [_make_batch_output_line(prompt_tokens=10, completion_tokens=5)] - - cost = _batch_cost_calculator( - file_content_dictionary=file_content, + cost, _, _ = _aggregate_batch_cost_usage_models( + entries=file_content, custom_llm_provider="openai", model_info=CUSTOM_MODEL_INFO, ) diff --git a/tests/batches_tests/test_batch_rate_limits.py b/tests/batches_tests/test_batch_rate_limits.py index e1fe8782ef9..2c804d21ace 100644 --- a/tests/batches_tests/test_batch_rate_limits.py +++ b/tests/batches_tests/test_batch_rate_limits.py @@ -913,7 +913,7 @@ async def test_batch_logging_azure_credentials_regression(): from unittest.mock import AsyncMock, MagicMock, patch from litellm.batches.batch_utils import ( _extract_file_access_credentials, - _get_batch_output_file_content_as_dictionary, + _fetch_batch_output_file_content, _handle_completed_batch, ) from litellm.types.llms.openai import Batch, HttpxBinaryResponseContent @@ -996,7 +996,7 @@ async def test_batch_logging_azure_credentials_regression(): with patch( "litellm.files.main.afile_content", side_effect=mock_afile_content_tracker ): - result = await _get_batch_output_file_content_as_dictionary( + result = await _fetch_batch_output_file_content( batch=mock_batch, custom_llm_provider="azure", litellm_params=azure_credentials, @@ -1092,7 +1092,7 @@ async def test_batch_logging_azure_credentials_regression(): ) # Call without litellm_params (should still work for OpenAI) - result = await _get_batch_output_file_content_as_dictionary( + result = await _fetch_batch_output_file_content( batch=mock_batch, custom_llm_provider="openai", litellm_params=None, diff --git a/tests/batches_tests/test_batches_logging_unit_tests.py b/tests/batches_tests/test_batches_logging_unit_tests.py index 33a0a87dd92..62b6f5b08e4 100644 --- a/tests/batches_tests/test_batches_logging_unit_tests.py +++ b/tests/batches_tests/test_batches_logging_unit_tests.py @@ -19,10 +19,8 @@ import litellm from litellm import create_batch, create_file from litellm._logging import verbose_logger from litellm.batches.batch_utils import ( - _batch_cost_calculator, + _aggregate_batch_cost_usage_models, _get_file_content_as_dictionary, - _get_batch_job_cost_from_file_content, - _get_batch_job_total_usage_from_file_content, _get_batch_job_usage_from_response_body, _get_response_from_batch_job_output_file, _batch_response_was_successful, @@ -139,9 +137,10 @@ def test_get_file_content_as_dictionary(sample_file_content): def test_get_batch_job_total_usage_from_file_content(sample_file_content_dict): - usage = _get_batch_job_total_usage_from_file_content( - sample_file_content_dict, custom_llm_provider="openai" - ) + with patch("litellm.completion_cost", return_value=0.0): + _, usage, _ = _aggregate_batch_cost_usage_models( + entries=sample_file_content_dict, custom_llm_provider="openai" + ) assert usage.total_tokens == 62 # 30 + 32 assert usage.prompt_tokens == 42 # 20 + 22 assert usage.completion_tokens == 20 # 10 + 10 @@ -157,8 +156,8 @@ async def test_batch_cost_calculator(sample_file_content_dict): so we expect the cost to be 0.5 * 2 = 1.0 """ with patch("litellm.completion_cost", return_value=0.5): - cost = _batch_cost_calculator( - file_content_dictionary=sample_file_content_dict, + cost, _, _ = _aggregate_batch_cost_usage_models( + entries=sample_file_content_dict, custom_llm_provider="openai", ) assert cost == 1.0 # 0.5 * 2 successful responses @@ -278,9 +277,12 @@ async def test_handle_completed_batch_computes_real_cost_from_output_file( created_at=1234567890, ) + sample_file_content_bytes = "\n".join( + json.dumps(row) for row in sample_file_content_dict + ).encode() with patch( - "litellm.batches.batch_utils._get_batch_output_file_content_as_dictionary", - new=AsyncMock(return_value=sample_file_content_dict), + "litellm.batches.batch_utils._fetch_batch_output_file_content", + new=AsyncMock(return_value=sample_file_content_bytes), ): cost, usage, models = await _handle_completed_batch( batch=batch, custom_llm_provider="openai" diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 0bc3cebdd5a..244e17b46a1 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -57,6 +57,9 @@ IGNORE_FUNCTIONS = [ "_filter_mcp_argument_value", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by blocking the MCP call at the cap. "_redact_scanned_content", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by returning "[REDACTED]" at the cap. "_iter_fallback_targets", # max depth set (2 * ROUTER_MAX_FALLBACKS); fails closed by raising ValueError at the cap. + "json_string_leaves", # max depth set (MAX_STRUCTURED_CONTENT_SCAN_DEPTH); fails closed by raising at the cap so nothing goes unscanned. + "with_json_string_leaves", # transitively bounded: only runs on a tree json_string_leaves already walked under the cap. + "json_unrewritable_labels", # max depth set (MAX_STRUCTURED_CONTENT_SCAN_DEPTH); returns the None sentinel at the cap so the caller blocks. ] diff --git a/tests/e2e/coverage_registry/quota_management.yaml b/tests/e2e/coverage_registry/quota_management.yaml index a8d0749cd8d..eb620395c46 100644 --- a/tests/e2e/coverage_registry/quota_management.yaml +++ b/tests/e2e/coverage_registry/quota_management.yaml @@ -23,7 +23,7 @@ - {id: quota_management.budget.key.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: key, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_duration zeroes key spend after the window; a blocked key serves again"} - {id: quota_management.budget.team.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: team, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_duration zeroes a team's spend after the window; every key on the team serves again"} - {id: quota_management.budget.organization.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: organization, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "An org budget resets after its window; keys under the org serve again"} -- {id: quota_management.budget.internal_user.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: internal_user, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "An internal user's budget resets after its window; their personal and team-member keys serve again"} +- {id: quota_management.budget.internal_user.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: internal_user, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "An internal user's budget resets after its window; their personal keys serve again"} - {id: quota_management.budget.team_member.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: team_member, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "Member per-team budget reset keeps advancing window after window"} - {id: quota_management.budget.key_multi_window.blocks_then_resets, module: quota_management, tier: P1, behavior: budget, variant: key_multi_window, assertions: [blocks_then_resets], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_limits enforce within a short window and serve again in the next"} - {id: quota_management.budget.key_multi_window.resets_windows_independently, module: quota_management, tier: P2, behavior: budget, variant: key_multi_window, assertions: [resets_windows_independently], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "Each window of a multi-window budget resets on its own schedule"} diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py index 5a54a4f0bbc..93861d19922 100644 --- a/tests/e2e/guardrails/guardrails_client.py +++ b/tests/e2e/guardrails/guardrails_client.py @@ -63,16 +63,6 @@ class OpenAIModerationParamsBody(GuardrailParamsBase): model: str | None = None -class PresidioParamsBody(GuardrailParamsBase): - guardrail: Literal["presidio"] = "presidio" - presidio_analyzer_api_base: str | None = None - presidio_anonymizer_api_base: str | None = None - # apply_to_output masks PII the model itself emitted, which also makes the - # guardrail run post_call. logging_only masks what the proxy logs. - apply_to_output: bool | None = None - logging_only: bool | None = None - - class BlockCodeExecutionParamsBody(GuardrailParamsBase): guardrail: Literal["block_code_execution"] = "block_code_execution" @@ -81,7 +71,6 @@ GuardrailParamsBody = ( ContentFilterParamsBody | BedrockGuardrailParamsBody | OpenAIModerationParamsBody - | PresidioParamsBody | BlockCodeExecutionParamsBody ) diff --git a/tests/e2e/guardrails/test_presidio_guardrail_e2e.py b/tests/e2e/guardrails/test_presidio_guardrail_e2e.py deleted file mode 100644 index 9742dfc6ae7..00000000000 --- a/tests/e2e/guardrails/test_presidio_guardrail_e2e.py +++ /dev/null @@ -1,141 +0,0 @@ -"""Live e2e: the built-in Presidio PII guardrail masks PII on the request and on -the model output. - -Presidio replaces detected PII with `` placeholders (e.g. -``) via a real analyzer + anonymizer. Two modes are checked -independently, each opted into per request (default_on=False) so it never touches -unrelated traffic: - -- pre_call: the prompt is anonymized before it reaches the model, so a - repeat-verbatim request comes back with the placeholder, never the raw email -- post_call (apply_to_output): PII the model itself emits is masked on the way - out, so the caller never receives the raw value the model produced - -A third mode, logging_only, is not covered here: the raw email stayed in the OTEL -span's `gen_ai.input.messages` on every attempt over a full poll deadline while -these two modes masked correctly, so that cell is tracked in LIT-4841 rather than -asserted against known-failing behavior. - -Analyzer/anonymizer bases come from PRESIDIO_ANALYZER_API_BASE / -PRESIDIO_ANONYMIZER_API_BASE (compose provides the in-network hosts; point them at -locally published container ports for a host run). The chat backend is a gemini -deployment created for the test. -""" - -from __future__ import annotations - -import os -import time -from collections.abc import Callable - -import pytest - -from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker -from e2e_http import unwrap -from guardrails_client import GuardrailMode, GuardrailsClient, PresidioParamsBody -from lifecycle import ResourceManager -from models import ChatResponse - -pytestmark = pytest.mark.e2e - -RAW_EMAIL = "alice.example.person@example.com" -PLACEHOLDER = "" - -ECHO_REQUEST = f"Repeat the following text back exactly, verbatim, with no changes: My email is {RAW_EMAIL}" -EMIT_REQUEST = f"Output exactly this one line and nothing else: Please contact {RAW_EMAIL} today" - - -def _content(response: ChatResponse) -> str: - if not response.choices: - return "" - message = response.choices[0].message - return (message.content if message else None) or "" - - -def _presidio_params( - mode: GuardrailMode, *, apply_to_output: bool = False, logging_only: bool = False -) -> PresidioParamsBody: - analyzer = os.environ["PRESIDIO_ANALYZER_API_BASE"] - anonymizer = os.environ["PRESIDIO_ANONYMIZER_API_BASE"] - return PresidioParamsBody( - mode=mode, - default_on=False, - presidio_analyzer_api_base=analyzer, - presidio_anonymizer_api_base=anonymizer, - apply_to_output=apply_to_output, - logging_only=logging_only, - ) - - -def _poll_until_masked(call: Callable[[], str]) -> str: - """Retry a call until the guardrail masks its PII, returning the last content. - - Registering a guardrail is a control-plane write; the data-plane worker that - serves /chat/completions only picks it up on its next periodic DB sync (~30s - in proxy_server.py), so a call issued the instant after the create runs - against a worker that has no guardrail yet and passes the raw value through. - That is in-flight propagation, not a masking failure. Polling to the deadline - waits it out, so the assertions that follow judge the synced state; if the - mask never lands the last unmasked content is returned and they still fail. - """ - deadline = time.monotonic() + POLL_TIMEOUT - last = call() - while time.monotonic() < deadline: - if PLACEHOLDER in last and RAW_EMAIL not in last: - return last - time.sleep(POLL_INTERVAL) - last = call() - return last - - -class TestPresidioGuardrail: - @pytest.mark.covers( - "guardrail.presidio.pre_call.masks", - exercised_on=["chat_completions"], - ) - def test_pre_call_masks_pii_before_the_model_sees_it( - self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str - ) -> None: - model = client.create_backend_model(resources, prefix="e2e-presidio-pre") - name = f"e2e-presidio-pre-{unique_marker()}" - guardrail_id = client.register(name, _presidio_params("pre_call")) - resources.defer(lambda: client.delete_guardrail(guardrail_id)) - - echoed = _poll_until_masked( - lambda: _content( - unwrap(client.chat(scoped_key, model, ECHO_REQUEST, guardrails=[name], max_tokens=128)) - ) - ) - assert RAW_EMAIL not in echoed, ( - "pre_call masking must strip the raw email before the model sees it, but the " - f"model echoed it back: {echoed[:300]!r}" - ) - assert PLACEHOLDER in echoed, ( - "the model should have echoed the masked placeholder the guardrail substituted, " - f"got: {echoed[:300]!r}" - ) - - @pytest.mark.covers( - "guardrail.presidio.post_call.masks", - exercised_on=["chat_completions"], - ) - def test_post_call_masks_pii_in_model_output( - self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str - ) -> None: - model = client.create_backend_model(resources, prefix="e2e-presidio-post") - name = f"e2e-presidio-post-{unique_marker()}" - guardrail_id = client.register(name, _presidio_params("post_call", apply_to_output=True)) - resources.defer(lambda: client.delete_guardrail(guardrail_id)) - - out = _poll_until_masked( - lambda: _content( - unwrap(client.chat(scoped_key, model, EMIT_REQUEST, guardrails=[name], max_tokens=128)) - ) - ) - assert RAW_EMAIL not in out, ( - "post_call masking must strip PII the model emitted, but the raw email reached the " - f"caller: {out[:300]!r}" - ) - assert PLACEHOLDER in out, ( - f"the masked placeholder should replace the model's PII output, got: {out[:300]!r}" - ) diff --git a/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py b/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py index 8b93afb4752..918739863ce 100644 --- a/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py +++ b/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py @@ -79,9 +79,14 @@ class TestBudgetBlocksPerLevel: ) @pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit") - def test_user_budget_enforced_across_all_their_keys( + def test_user_budget_enforced_across_their_personal_keys( self, client: BudgetClient, resources: ResourceManager ) -> None: + """A user's max_budget follows the person across their personal keys, so a + second untouched key is not a fresh allowance. It stops at the team + boundary: the same user's team-scoped key is governed by the team and + team-member budgets, both uncapped here, so it is the control that must + keep serving while the personal keys are refused.""" user_id = client.create_user(max_budget=TINY_CAP) resources.defer(lambda: client.delete_user(user_id)) first_key = client.generate_key(user_id=user_id) @@ -95,12 +100,17 @@ class TestBudgetBlocksPerLevel: resources.defer(lambda: client.delete_key(team_key)) _assert_blocked_429(client, first_key) - for label, key in (("second personal key", second_key), ("team-member key", team_key)): - result = _chat(client, key) - assert is_budget_block(result) and result.status_code == 429, ( - f"the {label} of a user over budget must get the same 429 budget_exceeded, " - f"got {result.status_code}: {result.body[:200]}" - ) + second = _chat(client, second_key) + assert is_budget_block(second) and second.status_code == 429, ( + f"the second personal key of a user over budget must get the same 429 budget_exceeded, " + f"got {second.status_code}: {second.body[:200]}" + ) + team_result = _chat(client, team_key) + assert not is_budget_block(team_result), ( + f"the team-scoped key of a user over their personal budget must keep serving; " + f"got {team_result.status_code}: {team_result.body[:200]}" + ) + require_successful_call(team_result) @pytest.mark.covers("quota_management.budget.end_user.blocks_over_limit") def test_end_user_budget_blocks_attributed_calls( diff --git a/tests/e2e/quota_management/budgets/test_budget_reset_e2e.py b/tests/e2e/quota_management/budgets/test_budget_reset_e2e.py index 793b22a47c7..b7b7f269c47 100644 --- a/tests/e2e/quota_management/budgets/test_budget_reset_e2e.py +++ b/tests/e2e/quota_management/budgets/test_budget_reset_e2e.py @@ -102,21 +102,6 @@ class TestBudgetResetPerLevel: _drive_to_block(client, key) _poll_until_serves_again(client, key) - @pytest.mark.covers("quota_management.budget.internal_user.resets_after_window") - def test_team_member_key_user_budget_resets_after_window( - self, client: BudgetClient, resources: ResourceManager - ) -> None: - user_id = client.create_user(max_budget=TINY_CAP, budget_duration=WINDOW) - resources.defer(lambda: client.delete_user(user_id)) - team_id = client.create_team(alias=f"e2e-user-team-reset-{unique_marker()}") - resources.defer(lambda: client.delete_team(team_id)) - client.add_team_member(team_id, user_id, max_budget_in_team=100.0) - key = client.generate_key(team_id=team_id, user_id=user_id) - resources.defer(lambda: client.delete_key(key)) - - _drive_to_block(client, key) - _poll_until_serves_again(client, key) - class TestKeyBudgetResetAcrossKeyKinds: """The tiny max_budget and its 30s window sit on the key itself while the user, diff --git a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index 1b1ce3e0f2d..9acb87750e9 100644 --- a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -428,6 +428,7 @@ def test_set_latency_metrics(prometheus_logger): model="gpt-5-mini", model_id="model-123", api_provider="openai", + service_tier=None, ) prometheus_logger.litellm_llm_api_time_to_first_token_metric.labels().observe.assert_called_once_with( 0.5 @@ -447,6 +448,7 @@ def test_set_latency_metrics(prometheus_logger): model="gpt-5-mini", model_id="model-123", api_provider="openai", + service_tier=None, ) prometheus_logger.litellm_llm_api_latency_metric.labels().observe.assert_called_once_with( 1.5 @@ -466,6 +468,7 @@ def test_set_latency_metrics(prometheus_logger): model="gpt-5-mini", model_id="model-123", api_provider="openai", + service_tier=None, ) prometheus_logger.litellm_request_total_latency_metric.labels().observe.assert_called_once_with( 2.0 @@ -634,6 +637,7 @@ def test_increment_top_level_request_and_spend_metrics(prometheus_logger): client_ip=None, user_agent=None, requested_model=None, + service_tier=None, ) prometheus_logger.litellm_spend_metric.labels().inc.assert_called_once_with(0.1) diff --git a/tests/litellm_utils_tests/test_secret_manager.py b/tests/litellm_utils_tests/test_secret_manager.py index 0a2419d0bea..0f95fd75c53 100644 --- a/tests/litellm_utils_tests/test_secret_manager.py +++ b/tests/litellm_utils_tests/test_secret_manager.py @@ -19,7 +19,8 @@ sys.path.insert( import pytest import litellm from litellm.llms.azure.azure import get_azure_ad_token_from_oidc -from litellm.llms.bedrock.chat import BedrockConverseLLM, BedrockLLM +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.bedrock.chat import BedrockConverseLLM from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2 from litellm.secret_managers.main import ( get_secret, @@ -160,7 +161,7 @@ def test_oidc_circle_v1_with_amazon(): aws_role_name = "arn:aws:iam::335785316107:role/litellm-github-unit-tests-circleci-v1-assume-only" aws_web_identity_token = "oidc/circleci/" - bllm = BedrockLLM() + bllm = BaseAWSLLM() creds = bllm.get_credentials( aws_region_name="ca-west-1", aws_web_identity_token=aws_web_identity_token, diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index f4c307e9c8a..8ab2feaf896 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -33,7 +33,7 @@ from litellm import ( completion_cost, embedding, ) -from litellm.llms.bedrock.chat import BedrockLLM +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt from base_llm_unit_tests import BaseLLMChatTest, BaseAnthropicChatTest @@ -225,7 +225,7 @@ def bedrock_session_token_creds(): aws_region_name = os.environ["AWS_REGION_NAME"] aws_session_token = os.environ.get("AWS_SESSION_TOKEN") - bllm = BedrockLLM() + bllm = BaseAWSLLM() if aws_session_token is not None: # For local testing creds = bllm.get_credentials( @@ -3573,40 +3573,11 @@ def test_bedrock_openai_model_id_extraction(): print(f"✓ Model ID extracted and encoded: {model_id}") -def test_bedrock_openai_convert_messages_to_prompt(): - """ - Test that convert_messages_to_prompt returns empty string for OpenAI models. - """ - from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM - - bedrock_llm = BedrockLLM() - messages = [ - {"role": "system", "content": "You are helpful"}, - {"role": "user", "content": "Hello"}, - ] - - prompt, chat_history = bedrock_llm.convert_messages_to_prompt( - model="test-model", messages=messages, provider="openai", custom_prompt_dict={} +def test_bedrock_openai_response_parsing(): + from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import ( + AmazonBedrockOpenAIConfig, ) - # OpenAI models use messages directly, no prompt conversion - assert prompt == "" - assert chat_history is None - print("✓ convert_messages_to_prompt returns empty for OpenAI") - - -def test_bedrock_openai_response_parsing(): - """ - Test that OpenAI responses are correctly parsed. - """ - from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM - from litellm import ModelResponse - from unittest.mock import Mock - import json - - bedrock_llm = BedrockLLM() - - # Mock OpenAI-style response openai_response = { "choices": [ { @@ -3627,34 +3598,24 @@ def test_bedrock_openai_response_parsing(): mock_response.status_code = 200 mock_response.headers = {} - model_response = ModelResponse() - mock_logging = Mock() - - result = bedrock_llm.process_response( + result = AmazonBedrockOpenAIConfig().transform_response( model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test", - response=mock_response, - model_response=model_response, - stream=False, - logging_obj=mock_logging, - optional_params={}, - api_key="", - data={}, + raw_response=mock_response, + model_response=ModelResponse(), + logging_obj=Mock(), + request_data={}, messages=[{"role": "user", "content": "What is the capital of France?"}], - print_verbose=lambda x: None, + optional_params={}, + litellm_params={}, encoding=None, ) - # Verify response content assert result.choices[0].message.content == "The capital of France is Paris." assert result.choices[0].finish_reason == "stop" - - # Verify usage assert result.usage.prompt_tokens == 10 assert result.usage.completion_tokens == 8 assert result.usage.total_tokens == 18 - print("✓ OpenAI response parsing works correctly") - def test_bedrock_openai_request_transformation(): """ @@ -3846,43 +3807,20 @@ def test_bedrock_openai_multiple_message_types(): def test_bedrock_openai_error_handling(): - """ - Test that errors from OpenAI models are properly handled. - """ - from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM - from litellm import ModelResponse + from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import ( + AmazonBedrockOpenAIConfig, + ) from litellm.llms.bedrock.common_utils import BedrockError - from unittest.mock import Mock - import json - bedrock_llm = BedrockLLM() + error = AmazonBedrockOpenAIConfig().get_error_class( + error_message="ValidationException: bad request", + status_code=422, + headers={}, + ) - # Mock error response - mock_response = Mock() - mock_response.json.side_effect = Exception("Invalid JSON") - mock_response.text = "Invalid response" - mock_response.status_code = 422 - - model_response = ModelResponse() - mock_logging = Mock() - - with pytest.raises(BedrockError) as exc_info: - bedrock_llm.process_response( - model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test", - response=mock_response, - model_response=model_response, - stream=False, - logging_obj=mock_logging, - optional_params={}, - api_key="", - data={}, - messages=[], - print_verbose=lambda x: None, - encoding=None, - ) - - assert exc_info.value.status_code == 422 - print("✓ Error handling works correctly") + assert isinstance(error, BedrockError) + assert error.status_code == 422 + assert "ValidationException: bad request" in str(error) # ============================================================================ diff --git a/tests/proxy_migration_tests/test_replica_identity_full.py b/tests/proxy_migration_tests/test_replica_identity_full.py new file mode 100644 index 00000000000..6a88e6994e9 --- /dev/null +++ b/tests/proxy_migration_tests/test_replica_identity_full.py @@ -0,0 +1,159 @@ +"""Coverage for the opt-in REPLICA IDENTITY FULL post-migration step. + +The DB-backed tests run against the same Postgres the migration suite uses, in +a throwaway schema so they cannot disturb the migrated tables. +""" + +import os +import uuid + +import pytest + +from litellm_proxy_extras.replica_identity import ( + REPLICA_IDENTITY_FULL_ENV_VAR, + apply_replica_identity_full, +) +from litellm_proxy_extras.utils import ProxyExtrasDBManager + +psycopg = pytest.importorskip("psycopg") + +requires_db = pytest.mark.skipif( + "DATABASE_URL" not in os.environ, + reason="requires a postgres database (DATABASE_URL)", +) + + +def _base_url() -> str: + return os.environ["DATABASE_URL"].split("?")[0] + + +def _replica_identities(schema: str) -> dict: + with psycopg.connect(_base_url(), autocommit=True) as conn: + rows = conn.execute( + "SELECT c.relname, c.relreplident FROM pg_class c " + "JOIN pg_namespace n ON n.oid = c.relnamespace " + "WHERE n.nspname = %s AND c.relkind = 'r'", + (schema,), + ).fetchall() + return dict(rows) + + +@pytest.fixture +def scratch_schema(monkeypatch): + """A schema holding two LiteLLM tables and one foreign table, all at the default.""" + schema = f"replica_identity_{uuid.uuid4().hex[:8]}" + with psycopg.connect(_base_url(), autocommit=True) as conn: + conn.execute(f'CREATE SCHEMA "{schema}"') + conn.execute( + f'CREATE TABLE "{schema}"."LiteLLM_ScratchTable" (id TEXT PRIMARY KEY, note TEXT)' + ) + conn.execute(f'CREATE TABLE "{schema}"."LiteLLM_ScratchSibling" (id TEXT PRIMARY KEY)') + conn.execute(f'CREATE TABLE "{schema}"."ScratchForeignTable" (id TEXT PRIMARY KEY)') + + monkeypatch.setenv("DATABASE_URL", f"{_base_url()}?schema={schema}") + yield schema + + with psycopg.connect(_base_url(), autocommit=True) as conn: + conn.execute(f'DROP SCHEMA "{schema}" CASCADE') + + +@requires_db +def test_applies_full_to_litellm_tables_only(scratch_schema, monkeypatch): + monkeypatch.setenv(REPLICA_IDENTITY_FULL_ENV_VAR, "true") + + assert ProxyExtrasDBManager.apply_replica_identity_full_if_requested() is True + + identities = _replica_identities(scratch_schema) + assert identities["LiteLLM_ScratchTable"] == "f" + assert identities["LiteLLM_ScratchSibling"] == "f" + assert identities["ScratchForeignTable"] == "d" + + +@requires_db +def test_a_locked_table_does_not_block_the_others(scratch_schema, monkeypatch): + """ALTER TABLE needs an exclusive lock, so a table busy with a long read has + to be skipped for the next run instead of stalling every other table behind it.""" + monkeypatch.setenv(REPLICA_IDENTITY_FULL_ENV_VAR, "true") + + with psycopg.connect(_base_url()) as holder: + holder.execute(f'SELECT * FROM "{scratch_schema}"."LiteLLM_ScratchTable"') + assert ProxyExtrasDBManager.apply_replica_identity_full_if_requested() is True + + identities = _replica_identities(scratch_schema) + assert identities["LiteLLM_ScratchTable"] == "d" + assert identities["LiteLLM_ScratchSibling"] == "f" + + +@requires_db +def test_leaves_tables_alone_when_not_requested(scratch_schema, monkeypatch): + monkeypatch.delenv(REPLICA_IDENTITY_FULL_ENV_VAR, raising=False) + + assert ProxyExtrasDBManager.apply_replica_identity_full_if_requested() is False + assert _replica_identities(scratch_schema)["LiteLLM_ScratchTable"] == "d" + + +@requires_db +def test_is_idempotent_across_runs(scratch_schema, monkeypatch): + monkeypatch.setenv(REPLICA_IDENTITY_FULL_ENV_VAR, "true") + + assert ProxyExtrasDBManager.apply_replica_identity_full_if_requested() is True + assert ProxyExtrasDBManager.apply_replica_identity_full_if_requested() is True + + assert _replica_identities(scratch_schema)["LiteLLM_ScratchTable"] == "f" + + +@requires_db +def test_reports_failure_without_raising(scratch_schema, monkeypatch): + """A run that cannot execute the statement must not take the migration down.""" + monkeypatch.setenv(REPLICA_IDENTITY_FULL_ENV_VAR, "true") + monkeypatch.setattr( + ProxyExtrasDBManager, + "_get_prisma_dir", + staticmethod(lambda: "/nonexistent/prisma/dir"), + ) + + assert ProxyExtrasDBManager.apply_replica_identity_full_if_requested() is False + assert _replica_identities(scratch_schema)["LiteLLM_ScratchTable"] == "d" + + +def test_reports_an_unrunnable_prisma_cli_without_raising(tmp_path): + """A deployment without the Prisma CLI on PATH must still finish its + migration run instead of dying on the optional replication step.""" + assert ( + apply_replica_identity_full( + schema_path=str(tmp_path / "schema.prisma"), + prisma_command=str(tmp_path / "no-such-prisma"), + prisma_env={}, + ) + is False + ) + + +def test_setup_database_applies_after_a_successful_migration_run(monkeypatch): + applied = [] + monkeypatch.setattr( + ProxyExtrasDBManager, "_run_migrations", staticmethod(lambda **kwargs: True) + ) + monkeypatch.setattr( + ProxyExtrasDBManager, + "apply_replica_identity_full_if_requested", + staticmethod(lambda: applied.append(True)), + ) + + assert ProxyExtrasDBManager.setup_database(use_migrate=True) is True + assert applied == [True] + + +def test_setup_database_skips_replica_identity_when_migrations_fail(monkeypatch): + applied = [] + monkeypatch.setattr( + ProxyExtrasDBManager, "_run_migrations", staticmethod(lambda **kwargs: False) + ) + monkeypatch.setattr( + ProxyExtrasDBManager, + "apply_replica_identity_full_if_requested", + staticmethod(lambda: applied.append(True)), + ) + + assert ProxyExtrasDBManager.setup_database(use_migrate=True) is False + assert applied == [] diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index d2218b08386..ad852c16905 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -2184,7 +2184,8 @@ def test_team_alias_stale_bypass_disabled_by_default(monkeypatch): pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None class _MockRouter: - team_model_to_deployment_indices = {("team-1", "gpt-4o"): [0]} + model_name_to_deployment_indices = {"model_name_team-1_legacy-uuid": [0]} + team_model_to_deployment_indices = {("team-1", "gpt-4o"): [1]} test_data = {"model": "gpt-4o"} user_api_key_dict = UserAPIKeyAuth( @@ -2209,7 +2210,8 @@ def test_team_alias_stale_bypass_enabled_by_flag(monkeypatch): pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None class _MockRouter: - team_model_to_deployment_indices = {("team-1", "gpt-4o"): [0]} + model_name_to_deployment_indices = {"model_name_team-1_legacy-uuid": [0]} + team_model_to_deployment_indices = {("team-1", "gpt-4o"): [1]} test_data = {"model": "gpt-4o"} user_api_key_dict = UserAPIKeyAuth( diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index 59c4caefa33..01dbb65a648 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -219,8 +219,8 @@ async def test_aaauser_personal_budgets(key_ownership): """ Set a personal budget on a user - User budget is enforced regardless of key ownership (personal or team). - Both cases should raise BudgetExceededError when the user is over budget. + - have it only apply when key belongs to user -> raises BudgetExceededError + - if key belongs to team, have key respect team budget -> allows call to go through """ import asyncio import time @@ -278,9 +278,12 @@ async def test_aaauser_personal_budgets(key_ownership): == valid_token ) - with pytest.raises(ProxyException) as exc_info: + if key_ownership == "user_key": + with pytest.raises(ProxyException) as exc_info: + await user_api_key_auth(request=request, api_key="Bearer " + user_key) + assert exc_info.value.type == ProxyErrorTypes.budget_exceeded + else: await user_api_key_auth(request=request, api_key="Bearer " + user_key) - assert exc_info.value.type == ProxyErrorTypes.budget_exceeded @pytest.mark.asyncio diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index a969d21a681..bcc70fae67c 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -2659,6 +2659,23 @@ def test_resolve_model_name_from_model_id(): result = router.resolve_model_name_from_model_id("gpt-5-mini") assert result == "gpt-5-mini" + # Test case 10: model_id is a deployment ID (hash) that differs from the + # public model_name. Regression for #32580: managed batch/file IDs embed the + # deployment model_id, and it must resolve back to the public model_name so + # team model-access checks compare against the model group, not the hash. + model_list = [ + { + "model_name": "bedrock-batch-model", + "litellm_params": { + "model": "bedrock/global.anthropic.claude-haiku-4-5-20251001-v1:0", + }, + "model_info": {"id": "8d0eaa7e6c6f54a425dfd0062cb6b0dc"}, + }, + ] + router = Router(model_list=model_list) + result = router.resolve_model_name_from_model_id("8d0eaa7e6c6f54a425dfd0062cb6b0dc") + assert result == "bedrock-batch-model" + def test_get_valid_args(): """Test get_valid_args static method returns valid Router.__init__ arguments""" diff --git a/tests/test_litellm/a2a_protocol/test_main.py b/tests/test_litellm/a2a_protocol/test_main.py index 2a675616245..c5e48620fde 100644 --- a/tests/test_litellm/a2a_protocol/test_main.py +++ b/tests/test_litellm/a2a_protocol/test_main.py @@ -104,3 +104,32 @@ async def test_streaming_trace_id_prefers_logging_trace_id(): pass assert captured["extra_headers"]["X-LiteLLM-Trace-Id"] == "trace-from-logging" + + +def test_streaming_logging_obj_carries_call_type_into_model_call_details(): + """The streaming logging object is built by hand rather than through + ``update_environment_variables``, which is the only place ``call_type`` normally + reaches ``model_call_details``. Callbacks read the call type from there, so + without this the streamed turn arrives at every logger with no call type at all + and OTel's GenAI metrics label it ``chat`` instead of ``invoke_agent``.""" + from a2a.compat.v0_3.types import MessageSendParams, SendStreamingMessageRequest + + from litellm.a2a_protocol.main import _build_streaming_logging_obj + + request = SendStreamingMessageRequest( + id="rpc-call-type", + params=MessageSendParams( + message={"messageId": "m1", "role": "user", "parts": [{"kind": "text", "text": "hi"}]} + ), + ) + + logging_obj = _build_streaming_logging_obj( + request=request, + agent_name="some-agent", + agent_id=None, + litellm_params=None, + metadata=None, + proxy_server_request=None, + ) + + assert logging_obj.model_call_details["call_type"] == "asend_message_streaming" diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index c7aecac477e..ea9dcea4e72 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -14,10 +14,13 @@ maps (litellm.completion_cost, batch_cost_calculator), the tokenizer deterministic stand-ins so the arithmetic under test is the only variable. """ +import json import os import sys +import httpx import pytest +import respx sys.path.insert(0, os.path.abspath("../../../..")) @@ -200,29 +203,34 @@ def test_estimate_tokens_never_zero_for_short_rows(): # =========================================================================== # -# _get_batch_models_from_file_content (output file) +# _aggregate_batch_cost_usage_models: models (output file) # =========================================================================== # -def test_output_models_uses_model_name_override(): - # model_name short-circuits: content is ignored entirely. - assert bu._get_batch_models_from_file_content([_success_row(model="ignored")], model_name="forced-model") == [ - "forced-model" - ] +def test_output_models_uses_model_name_override(monkeypatch): + monkeypatch.setattr(litellm, "completion_cost", lambda **kw: 0.0) + _, _, models = bu._aggregate_batch_cost_usage_models( + entries=[_success_row(model="ignored")], custom_llm_provider="openai", model_name="forced-model" + ) + assert models == ["forced-model"] -def test_output_models_collects_from_successful_only(): +def test_output_models_collects_from_successful_only(monkeypatch): + monkeypatch.setattr(litellm, "completion_cost", lambda **kw: 0.0) rows = [ _success_row(model="gpt-4o"), _failed_row(model="should-be-skipped"), _success_row(model="claude-3"), ] - assert bu._get_batch_models_from_file_content(rows) == ["gpt-4o", "claude-3"] + _, _, models = bu._aggregate_batch_cost_usage_models(entries=rows, custom_llm_provider="openai") + assert models == ["gpt-4o", "claude-3"] -def test_output_models_skips_successful_without_model(): +def test_output_models_skips_successful_without_model(monkeypatch): + monkeypatch.setattr(litellm, "completion_cost", lambda **kw: 0.0) rows = [{"response": {"status_code": 200, "body": {}}}] - assert bu._get_batch_models_from_file_content(rows) == [] + _, _, models = bu._aggregate_batch_cost_usage_models(entries=rows, custom_llm_provider="openai") + assert models == [] # =========================================================================== # @@ -235,6 +243,8 @@ def test_extract_credentials_only_known_keys(): "api_key": "sk-1", "api_base": "https://b", "vertex_project": "proj", + "gcs_bucket_name": "my-bucket", + "bucket_name": "my-alias-bucket", "model": "gpt-4o", # not a credential key "unrelated": "x", } @@ -242,6 +252,8 @@ def test_extract_credentials_only_known_keys(): "api_key": "sk-1", "api_base": "https://b", "vertex_project": "proj", + "gcs_bucket_name": "my-bucket", + "bucket_name": "my-alias-bucket", } @@ -261,6 +273,8 @@ def test_extract_credentials_all_supported_keys(): "vertex_project", "vertex_location", "vertex_credentials", + "gcs_bucket_name", + "bucket_name", "timeout", "max_retries", } @@ -372,17 +386,18 @@ def test_count_entry_uses_model_name_fallback(monkeypatch): # =========================================================================== # -# _get_batch_job_total_usage_from_file_content (output usage aggregation) +# _aggregate_batch_cost_usage_models: usage (output usage aggregation) # =========================================================================== # -def test_total_usage_sums_successful_only(): +def test_total_usage_sums_successful_only(monkeypatch): + monkeypatch.setattr(litellm, "completion_cost", lambda **kw: 0.0) rows = [ _success_row(usage=_usage(10, 5)), # 15 _failed_row(), # excluded _success_row(usage=_usage(20, 10)), # 30 ] - usage = bu._get_batch_job_total_usage_from_file_content(rows) + _, usage, _ = bu._aggregate_batch_cost_usage_models(entries=rows, custom_llm_provider="openai") assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( 30, 15, @@ -391,7 +406,9 @@ def test_total_usage_sums_successful_only(): def test_total_usage_empty_is_zero(): - usage = bu._get_batch_job_total_usage_from_file_content([]) + cost, usage, models = bu._aggregate_batch_cost_usage_models(entries=[], custom_llm_provider="openai") + assert cost == 0.0 + assert models == [] assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( 0, 0, @@ -400,7 +417,7 @@ def test_total_usage_empty_is_zero(): # =========================================================================== # -# _get_batch_job_cost_from_file_content (cost maps mocked) +# _aggregate_batch_cost_usage_models: cost (cost maps mocked) # =========================================================================== # @@ -419,7 +436,7 @@ def test_cost_from_content_completion_cost_path(monkeypatch): _success_row(usage=_usage(20, 10)), ] - total = bu._get_batch_job_cost_from_file_content(rows, custom_llm_provider="openai") + total, _, _ = bu._aggregate_batch_cost_usage_models(entries=rows, custom_llm_provider="openai") assert total == 1.0 # 2 successful * 0.5 assert len(calls) == 2 # failed row not costed @@ -435,8 +452,8 @@ def test_cost_from_content_model_info_path(monkeypatch): _success_row(usage=_usage(20, 10)), ] - total = bu._get_batch_job_cost_from_file_content( - rows, + total, _, _ = bu._aggregate_batch_cost_usage_models( + entries=rows, custom_llm_provider="openai", model_info={"input_cost_per_token": 0.0}, # type: ignore[arg-type] # truthy -> model_info path ) @@ -444,32 +461,65 @@ def test_cost_from_content_model_info_path(monkeypatch): assert total == pytest.approx(0.6) # 2 * (0.1 + 0.2) +def test_aggregate_consumes_entries_in_a_single_pass(monkeypatch): + """A one-shot generator: any implementation that iterates the entries twice + (e.g. separate cost and usage passes) sees nothing on the second pass and + returns wrong totals for at least one of cost/usage/models.""" + monkeypatch.setattr(litellm, "completion_cost", lambda **kw: 0.5) + one_shot = (row for row in [_success_row(usage=_usage(10, 5)), _failed_row(), _success_row(usage=_usage(20, 10))]) + + cost, usage, models = bu._aggregate_batch_cost_usage_models(entries=one_shot, custom_llm_provider="openai") + + assert cost == 1.0 + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (30, 15, 45) + assert models == ["gpt-4o", "gpt-4o"] + + # =========================================================================== # -# _batch_cost_calculator (dispatch: vertex-disable-transform vs generic) +# calculate_batch_cost_and_usage (dispatch: vertex-disable-transform vs generic) # =========================================================================== # -def test_batch_cost_calculator_generic_path(monkeypatch): - monkeypatch.setattr(bu, "_get_batch_job_cost_from_file_content", lambda **kw: 4.2) - assert bu._batch_cost_calculator([], custom_llm_provider="openai", model_name="gpt-4o") == 4.2 - - -def test_batch_cost_calculator_vertex_disable_transform_path(monkeypatch): +@pytest.mark.asyncio +async def test_calculate_vertex_disable_transform_path(monkeypatch): monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False) monkeypatch.setattr( bu, "calculate_vertex_ai_batch_cost_and_usage", - lambda content, model: (9.9, Usage()), + lambda content, model: (9.9, Usage(prompt_tokens=1, completion_tokens=2, total_tokens=3)), ) # generic path must NOT be taken monkeypatch.setattr( bu, - "_get_batch_job_cost_from_file_content", + "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run"), ) - cost = bu._batch_cost_calculator([], custom_llm_provider="vertex_ai", model_name="gemini-2.0-flash-001") + cost, usage, models = await bu.calculate_batch_cost_and_usage( + file_content_dictionary=[], custom_llm_provider="vertex_ai", model_name="gemini-2.0-flash-001" + ) assert cost == 9.9 + assert usage.total_tokens == 3 + assert models == ["gemini-2.0-flash-001"] + + +@pytest.mark.asyncio +async def test_calculate_vertex_disable_transform_needs_model_name(monkeypatch): + """Without a model_name the raw-vertex path cannot price lines; the generic + aggregation path must run even with the disable flag set.""" + monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False) + monkeypatch.setattr( + bu, + "calculate_vertex_ai_batch_cost_and_usage", + lambda content, model: pytest.fail("raw vertex path should not run"), + ) + + cost, usage, models = await bu.calculate_batch_cost_and_usage( + file_content_dictionary=[], custom_llm_provider="vertex_ai" + ) + assert cost == 0.0 + assert usage.total_tokens == 0 + assert models == [] # =========================================================================== # @@ -579,24 +629,19 @@ def test_vertex_cost_error_in_line_is_swallowed(monkeypatch): @pytest.mark.asyncio async def test_calculate_batch_cost_and_usage_orchestration(monkeypatch): rows = [_success_row(model="gpt-4o", usage=_usage(10, 5))] - monkeypatch.setattr(bu, "_batch_cost_calculator", lambda **kw: 2.5) - monkeypatch.setattr( - bu, - "_get_batch_job_total_usage_from_file_content", - lambda **kw: Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), - ) + monkeypatch.setattr(litellm, "completion_cost", lambda **kw: 2.5) cost, usage, models = await bu.calculate_batch_cost_and_usage( file_content_dictionary=rows, custom_llm_provider="openai" ) assert cost == 2.5 - assert usage.total_tokens == 15 - assert models == ["gpt-4o"] # real _get_batch_models_from_file_content + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (10, 5, 15) + assert models == ["gpt-4o"] # =========================================================================== # -# _get_batch_output_file_content_as_dictionary (file fetch + credential merge) +# _fetch_batch_output_file_content (file fetch + credential merge) # =========================================================================== # @@ -615,16 +660,217 @@ def _batch(output_file_id): ) +def _vertex_openai_row(custom_id, model, prompt_tokens, completion_tokens): + return { + "id": f"batch_req_{custom_id}", + "custom_id": custom_id, + "response": { + "status_code": 200, + "request_id": custom_id, + "body": { + "id": f"chatcmpl-{custom_id}", + "object": "chat.completion", + "model": model, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": _usage(prompt_tokens, completion_tokens), + }, + }, + "error": None, + } + + +def _vertex_jsonl(rows): + return "\n".join(json.dumps(row) for row in rows).encode() + + @pytest.mark.asyncio -async def test_output_file_content_vertex_raises(): - with pytest.raises(ValueError, match="Vertex AI does not support"): - await bu._get_batch_output_file_content_as_dictionary(_batch("of"), custom_llm_provider="vertex_ai") +async def test_output_file_content_vertex_fetches_via_afile_content(monkeypatch): + import litellm.files.main as files_main + + rows = [_vertex_openai_row("request-1", "gemini-3.6-flash", 10, 5)] + captured: dict = {} + + async def fake_afile_content(**kw): + captured.update(kw) + return type("R", (), {"content": _vertex_jsonl(rows)})() + + monkeypatch.setattr(files_main, "afile_content", fake_afile_content) + + result = await bu._fetch_batch_output_file_content( + _batch("gs://litellm-bucket/output/predictions.jsonl"), + custom_llm_provider="vertex_ai", + litellm_params={ + "vertex_project": "proj-1", + "vertex_location": "us-central1", + "vertex_credentials": "/path/to/creds.json", + "gcs_bucket_name": "litellm-bucket", + "model": "vertex_ai/gemini-3.6-flash", + }, + ) + + assert bu._get_file_content_as_dictionary(result) == rows + assert captured["file_id"] == "gs://litellm-bucket/output/predictions.jsonl" + assert captured["custom_llm_provider"] == "vertex_ai" + assert captured["vertex_project"] == "proj-1" + assert captured["vertex_location"] == "us-central1" + assert captured["vertex_credentials"] == "/path/to/creds.json" + assert captured["gcs_bucket_name"] == "litellm-bucket" + assert "model" not in captured + + +@pytest.mark.asyncio +async def test_output_file_content_vertex_unified_file_id_extracts_gcs_uri(monkeypatch): + import base64 + + import litellm.files.main as files_main + + captured: dict = {} + + async def fake_afile_content(**kw): + captured.update(kw) + return type("R", (), {"content": b'{"a": 1}'})() + + monkeypatch.setattr(files_main, "afile_content", fake_afile_content) + unified_id = ( + "litellm_proxy:application/jsonl;unified_id,uuid-1;target_model_names,vertex-model;" + "llm_output_file_id,gs://litellm-bucket/output/predictions.jsonl;llm_output_file_model_id,model-1" + ) + encoded_id = base64.urlsafe_b64encode(unified_id.encode()).decode().rstrip("=") + + await bu._fetch_batch_output_file_content(_batch(encoded_id), custom_llm_provider="vertex_ai") + + assert captured["file_id"] == "gs://litellm-bucket/output/predictions.jsonl" + assert captured["custom_llm_provider"] == "vertex_ai" + + +def _vertex_predictions_row(custom_id, prompt_tokens, completion_tokens): + return { + "request": { + "contents": [{"role": "user", "parts": [{"text": "hi"}]}], + "labels": {"litellm_custom_id": custom_id}, + }, + "status": "", + "response": { + "candidates": [ + { + "content": {"role": "model", "parts": [{"text": "ok"}]}, + "finishReason": "STOP", + } + ], + "usageMetadata": { + "promptTokenCount": prompt_tokens, + "candidatesTokenCount": completion_tokens, + "totalTokenCount": prompt_tokens + completion_tokens, + }, + "modelVersion": "gemini-3.6-flash", + }, + "processed_time": "2026-07-30T00:00:00.000000+00:00", + } + + +@pytest.fixture +def respx_interceptable_httpx_client(monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.asyncio +@respx.mock +async def test_output_file_content_vertex_managed_uri_accepted_by_real_validation(respx_interceptable_httpx_client): + managed_output_uri = ( + "gs://litellm-bucket/litellm-vertex-files/publishers/google/models/" + "gemini-3.6-flash/abc-123/prediction-model/predictions.jsonl" + ) + rows = [ + _vertex_predictions_row("request-1", 10, 5), + _vertex_predictions_row("request-2", 20, 10), + ] + route = respx.get(url__regex=r"https://storage\.googleapis\.com/storage/v1/b/litellm-bucket/o/.*").mock( + return_value=httpx.Response(200, content=_vertex_jsonl(rows)) + ) + + file_content = await bu._fetch_batch_output_file_content( + _batch(managed_output_uri), + custom_llm_provider="vertex_ai", + litellm_params={ + "api_key": "test-token", + "vertex_project": "proj-1", + "vertex_location": "us-central1", + "gcs_bucket_name": "litellm-bucket", + }, + ) + result = bu._get_file_content_as_dictionary(file_content) + + assert route.call_count == 1 + request = route.calls.last.request + assert request.url.raw_path == ( + b"/storage/v1/b/litellm-bucket/o/" + b"litellm-vertex-files%2Fpublishers%2Fgoogle%2Fmodels%2Fgemini-3.6-flash" + b"%2Fabc-123%2Fprediction-model%2Fpredictions.jsonl?alt=media" + ) + assert [row["custom_id"] for row in result] == ["request-1", "request-2"] + assert all(row["response"]["status_code"] == 200 for row in result) + assert all(row["response"]["body"]["model"] == "gemini-3.6-flash" for row in result) + assert [row["response"]["body"]["usage"]["prompt_tokens"] for row in result] == [10, 20] + assert [row["response"]["body"]["usage"]["completion_tokens"] for row in result] == [5, 10] + + +@pytest.mark.asyncio +@respx.mock +async def test_output_file_content_vertex_foreign_bucket_rejected_by_real_validation(): + with pytest.raises(Exception, match="does not match the configured storage bucket"): + await bu._fetch_batch_output_file_content( + _batch("gs://attacker-bucket/litellm-vertex-files/x/predictions.jsonl"), + custom_llm_provider="vertex_ai", + litellm_params={ + "api_key": "test-token", + "vertex_project": "proj-1", + "vertex_location": "us-central1", + "gcs_bucket_name": "litellm-bucket", + }, + ) + + assert respx.mock.calls.call_count == 0 + + +@pytest.mark.asyncio +async def test_handle_completed_vertex_batch_computes_cost_usage_and_models(monkeypatch): + import litellm.files.main as files_main + + rows = [ + _vertex_openai_row("request-1", "gemini-3.6-flash", 10, 5), + _vertex_openai_row("request-2", "gemini-3.6-flash", 20, 10), + ] + + async def fake_afile_content(**kw): + return type("R", (), {"content": _vertex_jsonl(rows)})() + + monkeypatch.setattr(files_main, "afile_content", fake_afile_content) + + cost, usage, models = await bu._handle_completed_batch( + _batch("gs://litellm-bucket/output/predictions.jsonl"), + custom_llm_provider="vertex_ai", + litellm_params={"vertex_project": "proj-1", "vertex_location": "us-central1"}, + ) + + assert cost > 0 + assert cost == pytest.approx(30 * 7.5e-07 + 15 * 3.75e-06) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (30, 15, 45) + assert models == ["gemini-3.6-flash", "gemini-3.6-flash"] @pytest.mark.asyncio async def test_output_file_content_no_output_file_id_raises(): with pytest.raises(ValueError, match="Output file id is None"): - await bu._get_batch_output_file_content_as_dictionary(_batch(None), custom_llm_provider="openai") + await bu._fetch_batch_output_file_content(_batch(None), custom_llm_provider="openai") @pytest.mark.asyncio @@ -641,13 +887,13 @@ async def test_output_file_content_fetches_and_parses(monkeypatch): monkeypatch.setattr(files_main, "afile_content", fake_afile_content) monkeypatch.setattr(cu, "_is_base64_encoded_unified_file_id", lambda fid: False) - result = await bu._get_batch_output_file_content_as_dictionary( + result = await bu._fetch_batch_output_file_content( _batch("file-out"), custom_llm_provider="azure", litellm_params={"api_key": "sk-az", "api_base": "https://az", "model": "x"}, ) - assert result == [{"a": 1}, {"b": 2}] + assert result == b'{"a": 1}\n{"b": 2}' # afile_content received the file id + extracted credentials (not "model"). assert captured["file_id"] == "file-out" assert captured["custom_llm_provider"] == "azure" @@ -676,13 +922,13 @@ async def test_output_file_content_unified_file_id_extraction(monkeypatch): lambda fid: "litellm_proxy;llm_output_file_id,real-file-99;rest", ) - await bu._get_batch_output_file_content_as_dictionary(_batch("encoded-blob"), custom_llm_provider="openai") + await bu._fetch_batch_output_file_content(_batch("encoded-blob"), custom_llm_provider="openai") assert captured["file_id"] == "real-file-99" # =========================================================================== # -# _handle_completed_batch (async orchestrator: fetch -> cost/usage/models) +# _handle_completed_batch (async orchestrator: fetch -> single-pass aggregate) # =========================================================================== # @@ -690,49 +936,48 @@ async def test_output_file_content_unified_file_id_extraction(monkeypatch): async def test_handle_completed_batch_orchestration(monkeypatch): rows = [_success_row(model="gpt-4o", usage=_usage(10, 5))] - async def fake_get_content(batch, custom_llm_provider, litellm_params=None): - return rows + async def fake_fetch(batch, custom_llm_provider, litellm_params=None): + return _vertex_jsonl(rows) - monkeypatch.setattr(bu, "_get_batch_output_file_content_as_dictionary", fake_get_content) - monkeypatch.setattr(bu, "_batch_cost_calculator", lambda **kw: 3.3) - monkeypatch.setattr( - bu, - "_get_batch_job_total_usage_from_file_content", - lambda **kw: Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), - ) + monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch) + monkeypatch.setattr(litellm, "completion_cost", lambda **kw: 3.3) cost, usage, models = await bu._handle_completed_batch(_batch("of"), custom_llm_provider="openai") assert cost == 3.3 - assert usage.total_tokens == 15 + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (10, 5, 15) assert models == ["gpt-4o"] -# =========================================================================== # -# Remaining branch: vertex usage disable-transform path. -# -# NOTE: the error path of _get_batch_job_cost_from_file_content (its `raise e`) -# is intentionally NOT tested: the preceding line logs via -# `verbose_logger.error("...", e)`, which passes the exception as a logging -# format-arg with no placeholder and itself raises TypeError under -# logging.raiseExceptions, masking the original error. Asserting that masked -# behavior would lock a source bug; left uncovered on purpose. -# =========================================================================== # +@pytest.mark.asyncio +async def test_handle_completed_batch_vertex_disable_transform_path(monkeypatch): + raw_rows = [{"response": {"usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 2}}}] + async def fake_fetch(batch, custom_llm_provider, litellm_params=None): + return _vertex_jsonl(raw_rows) -def test_total_usage_vertex_disable_transform_path(monkeypatch): + monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch) monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False) - monkeypatch.setattr( - bu, - "calculate_vertex_ai_batch_cost_and_usage", - lambda content, model: ( - 0.0, - Usage(prompt_tokens=1, completion_tokens=2, total_tokens=3), - ), + seen: dict = {} + + def fake_vertex_calc(content, model): + seen["content"] = content + seen["model"] = model + return 7.7, Usage(prompt_tokens=1, completion_tokens=2, total_tokens=3) + + monkeypatch.setattr(bu, "calculate_vertex_ai_batch_cost_and_usage", fake_vertex_calc) + + cost, usage, models = await bu._handle_completed_batch( + _batch("gs://litellm-bucket/output/predictions.jsonl"), + custom_llm_provider="vertex_ai", + model_name="gemini-x", ) - usage = bu._get_batch_job_total_usage_from_file_content([], custom_llm_provider="vertex_ai", model_name="gemini-x") + assert cost == 7.7 assert usage.total_tokens == 3 + assert models == ["gemini-x"] + assert seen["content"] == raw_rows + assert seen["model"] == "gemini-x" def _anthropic_usage(input_tokens, output_tokens, cache_creation=0, cache_read=0): @@ -832,46 +1077,54 @@ def test_bedrock_cost_uses_deployment_model_name(): "recordId": "1", "modelOutput": {"model": "claude-sonnet-4-6", "usage": {"input_tokens": 13, "output_tokens": 5}}, } - cost = bu._get_batch_job_cost_from_file_content( - file_content_dictionary=[row], + cost, _, models = bu._aggregate_batch_cost_usage_models( + entries=[row], custom_llm_provider="bedrock", model_name="us.anthropic.claude-sonnet-4-6", model_info={}, ) assert cost > 0 + assert models == ["us.anthropic.claude-sonnet-4-6"] -def test_anthropic_total_usage_sums_succeeded_only(): +def test_anthropic_total_usage_sums_succeeded_only(monkeypatch): + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (0.0, 0.0)) rows = [ _anthropic_succeeded_row(usage=_anthropic_usage(10, 5)), _anthropic_errored_row(), _anthropic_succeeded_row(usage=_anthropic_usage(20, 10, cache_read=100)), ] - usage = bu._get_batch_job_total_usage_from_file_content(rows, custom_llm_provider="anthropic") + _, usage, _ = bu._aggregate_batch_cost_usage_models(entries=rows, custom_llm_provider="anthropic") assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (130, 15, 145) -def test_anthropic_total_usage_aggregates_cache_token_details(): +def test_anthropic_total_usage_aggregates_cache_token_details(monkeypatch): + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (0.0, 0.0)) rows = [ _anthropic_succeeded_row(usage=_anthropic_usage(1000, 200, cache_creation=2000, cache_read=8000)), _anthropic_errored_row(), _anthropic_succeeded_row(usage=_anthropic_usage(50, 20, cache_creation=300, cache_read=700)), ] - usage = bu._get_batch_job_total_usage_from_file_content(rows, custom_llm_provider="anthropic") + _, usage, _ = bu._aggregate_batch_cost_usage_models(entries=rows, custom_llm_provider="anthropic") assert usage.prompt_tokens_details.cached_tokens == 8700 assert usage.prompt_tokens_details.cache_creation_tokens == 2300 assert usage.cache_read_input_tokens == 8700 assert usage.cache_creation_input_tokens == 2300 -def test_total_usage_without_cache_tokens_has_no_prompt_details(): +def test_total_usage_without_cache_tokens_has_no_prompt_details(monkeypatch): + monkeypatch.setattr(litellm, "completion_cost", lambda **kw: 0.0) rows = [ { "custom_id": "req-1", "response": {"status_code": 200, "body": {"model": "gpt-5.2", "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}}, } ] - usage = bu._get_batch_job_total_usage_from_file_content(rows, custom_llm_provider="openai") + _, usage, _ = bu._aggregate_batch_cost_usage_models(entries=rows, custom_llm_provider="openai") assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (10, 5, 15) assert usage.prompt_tokens_details is None @@ -884,8 +1137,8 @@ def test_anthropic_cost_applies_batch_discount_and_cache_pricing(): _anthropic_errored_row(), ] - total = bu._get_batch_job_cost_from_file_content( - rows, + total, _, _ = bu._aggregate_batch_cost_usage_models( + entries=rows, custom_llm_provider="anthropic", model_info=_ANTHROPIC_MODEL_INFO, # type: ignore[arg-type] ) @@ -910,8 +1163,8 @@ def test_anthropic_cost_without_model_info_uses_batch_cost_calculator(monkeypatc lambda **kw: pytest.fail("anthropic rows must not go through completion_cost"), ) - total = bu._get_batch_job_cost_from_file_content( - [_anthropic_succeeded_row()], custom_llm_provider="anthropic" + total, _, _ = bu._aggregate_batch_cost_usage_models( + entries=[_anthropic_succeeded_row()], custom_llm_provider="anthropic" ) assert total == pytest.approx(0.3) @@ -920,12 +1173,16 @@ def test_anthropic_cost_without_model_info_uses_batch_cost_calculator(monkeypatc assert seen[0]["usage"].prompt_tokens == 10 -def test_anthropic_batch_models_collected_from_succeeded_rows(): +def test_anthropic_batch_models_collected_from_succeeded_rows(monkeypatch): + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (0.0, 0.0)) rows = [ _anthropic_succeeded_row(model="claude-sonnet-4-5-20250929"), _anthropic_errored_row(), ] - assert bu._get_batch_models_from_file_content(rows, None, "anthropic") == ["claude-sonnet-4-5-20250929"] + _, _, models = bu._aggregate_batch_cost_usage_models(entries=rows, custom_llm_provider="anthropic") + assert models == ["claude-sonnet-4-5-20250929"] @pytest.mark.asyncio diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index a2e18a62638..59200719197 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -59,9 +59,24 @@ def test_delete_cache_applies_namespace(namespace, monkeypatch, redis_no_ping): @pytest.mark.asyncio -async def test_redis_client_init_with_socket_timeout(monkeypatch, redis_no_ping): - monkeypatch.setenv("REDIS_HOST", "my-fake-host") - redis_cache = RedisCache(socket_timeout=1.0) +@pytest.mark.parametrize( + "redis_config", + [ + pytest.param({"host": "my-fake-host"}, id="host_port"), + pytest.param({"url": "redis://my-fake-host:6379"}, id="url"), + ], +) +async def test_redis_client_init_with_socket_timeout(monkeypatch, redis_no_ping, redis_config): + """socket_timeout has to reach the connection however Redis was configured. + + A url config used to drop every connection kwarg, so redis-py was left with + socket_timeout (and socket_connect_timeout, which falls back to it) unset. A + Redis host that drops packets instead of refusing them then blocks each caller + indefinitely, and the circuit breaker never trips because no call ever returns. + """ + monkeypatch.delenv("REDIS_URL", raising=False) + monkeypatch.delenv("REDIS_HOST", raising=False) + redis_cache = RedisCache(socket_timeout=1.0, **redis_config) assert redis_cache.redis_kwargs["socket_timeout"] == 1.0 client = redis_cache.init_async_client() assert client is not None @@ -428,3 +443,168 @@ def test_delete_cache_namespaces_key(namespace, expected, monkeypatch, redis_no_ redis_cache.redis_client = mock_client redis_cache.delete_cache(key="k") mock_client.delete.assert_called_once_with(expected) + + +def _closed_port() -> int: + """A port with nothing listening, so Redis calls fail fast and deterministically.""" + import socket + + with socket.socket() as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "call_method", + [ + pytest.param(lambda c: c.async_get_cache("lit4930"), id="async_get_cache"), + pytest.param(lambda c: c.async_batch_get_cache(["lit4930"]), id="async_batch_get_cache"), + pytest.param(lambda c: c.async_set_cache("lit4930", "v"), id="async_set_cache"), + pytest.param(lambda c: c.async_get_ttl("lit4930"), id="async_get_ttl"), + ], +) +async def test_circuit_breaker_opens_when_method_swallows_redis_failure(redis_no_ping, call_method): + """A guarded method that swallows its own Redis error must still count as a failure. + + These methods catch connection errors and return a default so callers degrade instead + of failing, which is correct. But that returns cleanly through the circuit breaker + guard, and counting it as a success reset the failure streak on every call, so the + breaker could never open. An unreachable Redis then stayed in the pool and every + request kept paying the full socket timeout on it. + """ + from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD + + cache = RedisCache(host="127.0.0.1", port=_closed_port(), socket_timeout=0.5) + + for _ in range(REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD): + await call_method(cache) + + with pytest.raises(Exception, match="circuit breaker is open"): + await call_method(cache) + + +@pytest.mark.asyncio +async def test_circuit_breaker_success_still_resets_the_failure_streak(redis_no_ping): + """A reachable Redis must keep the breaker closed, however many earlier calls failed. + + The guard now records success only when nothing failed while the method ran, so this + pins the other half of that contract: a call that genuinely reaches Redis has to clear + the streak, or a healthy Redis would eventually be evicted from the pool. + """ + from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD + + cache = RedisCache(host="127.0.0.1", port=_closed_port(), socket_timeout=0.5) + + for _ in range(REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD - 1): + await cache.async_get_cache("lit4930") + assert cache._circuit_breaker.is_open() is False + + reachable_redis = AsyncMock() + reachable_redis.get.return_value = None + with patch.object(cache, "init_async_client", return_value=reachable_redis): + await cache.async_get_cache("lit4930") + + for _ in range(REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD - 1): + await cache.async_get_cache("lit4930") + + assert cache._circuit_breaker.is_open() is False, "one success must clear the streak" + + +@pytest.mark.asyncio +async def test_circuit_breaker_covers_lua_script_execution(redis_no_ping): + """Lua script execution must feed the breaker like every other Redis call. + + The v3 rate limiter issues all of its Redis traffic through async_register_script, so + leaving that path unguarded meant the coordination calls during an outage never + counted toward taking Redis out of the pool and kept paying a full socket timeout + each, which is the traffic the outage hurts most. + """ + from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD + + cache = RedisCache(host="127.0.0.1", port=_closed_port(), socket_timeout=0.5) + run_script = cache.async_register_script("return 1") + + for _ in range(REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD): + with pytest.raises(Exception): + await run_script(keys=["lit4930"], args=[1]) + + with pytest.raises(Exception, match="circuit breaker is open"): + await run_script(keys=["lit4930"], args=[1]) + + +@pytest.mark.asyncio +async def test_concurrent_success_is_not_cancelled_by_another_calls_failure(): + """One caller's failure must not discard a different caller's success. + + A breaker is shared by every concurrent caller, so tracking "did this call fail" on the + breaker itself cannot tell my failure from someone else's. A Redis that is still + answering would then be evicted from the pool by unrelated in-flight failures, which is + the opposite of the outage this guard exists to handle. + """ + from redis.exceptions import ConnectionError as RedisConnectionError + + from litellm.caching.redis_cache import ( + RedisCircuitBreaker, + _record_swallowed_redis_failure, + _run_under_circuit_breaker, + ) + + breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60) + + # The failure has to land after both calls are already in flight, which is the only + # ordering where a shared counter confuses the two. Failing before the healthy call + # starts would leave its snapshot correct and prove nothing. + async def swallows_a_failure(): + await asyncio.sleep(0.02) + _record_swallowed_redis_failure(breaker, RedisConnectionError("redis unreachable")) + return None + + async def succeeds_while_the_other_fails(): + await asyncio.sleep(0.05) + return "ok" + + rounds = breaker.failure_threshold + 1 + for _ in range(rounds): + await asyncio.gather( + _run_under_circuit_breaker(breaker, "failing", swallows_a_failure), + _run_under_circuit_breaker(breaker, "healthy", succeeds_while_the_other_fails), + ) + + assert breaker._failure_count < breaker.failure_threshold, "the healthy call must clear the streak" + assert breaker.is_open() is False, "a Redis answering every round must stay in the pool" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "error, opens_breaker", + [ + pytest.param("ConnectionError", True, id="connection_refused_is_unhealthy"), + pytest.param("TimeoutError", True, id="timeout_is_unhealthy"), + pytest.param("BusyLoadingError", True, id="loading_is_unhealthy"), + pytest.param("ResponseError", False, id="wrong_type_command_is_not"), + pytest.param("DataError", False, id="bad_data_is_not"), + ], +) +async def test_only_connectivity_failures_open_the_breaker(error, opens_breaker): + """Command and data errors must not count against Redis health. + + They say nothing about connectivity, and a caller able to provoke them (an INCR against + a non-numeric value, say) could otherwise trip the shared breaker on demand and drop + rate limiting to per-process counters, which spreading traffic across replicas outruns. + """ + import redis.exceptions + + from litellm.caching.redis_cache import RedisCircuitBreaker, _run_under_circuit_breaker + + breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60) + raised = getattr(redis.exceptions, error)("boom") + + async def failing_call(): + raise raised + + for _ in range(breaker.failure_threshold + 1): + with pytest.raises(Exception): + await _run_under_circuit_breaker(breaker, "op", failing_call) + + assert breaker.is_open() is opens_breaker diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py index 3169b9b08e0..2580197d6d2 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -5,6 +5,8 @@ Regression test for afile_retrieve called without credentials in async_post_call_success_hook when processing completed batch responses. """ +import json + import pytest from typing import Optional from unittest.mock import AsyncMock, MagicMock, patch @@ -385,3 +387,59 @@ async def test_afile_content_error_reports_unified_id_not_provider_uri(): message = str(exc_info.value) assert unified_file_id in message assert s3_uri not in message + + +def _make_real_managed_files_instance(): + """Create a _PROXY_LiteLLMManagedFiles with a real store_unified_file_id but + an AsyncMock prisma client, so the DB write path itself can be asserted.""" + from litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles, + ) + + mock_cache = MagicMock() + mock_cache.async_set_cache = AsyncMock() + + mock_prisma = MagicMock() + mock_prisma.db.litellm_managedfiletable.upsert = AsyncMock() + mock_prisma.db.litellm_managedfiletable.create = AsyncMock( + side_effect=AssertionError( + "store_unified_file_id must upsert, not create, on the retrieve path" + ) + ) + + return ( + _PROXY_LiteLLMManagedFiles( + internal_usage_cache=mock_cache, + prisma_client=mock_prisma, + ), + mock_prisma, + ) + + +@pytest.mark.asyncio +async def test_store_unified_file_id_is_idempotent_via_upsert(): + """Regression test for the managed-batch retrieve 500 (UniqueViolationError on + unified_file_id): re-registering an already-stored output file id must upsert on + unified_file_id, never do an unconditional create that raises on conflict.""" + managed_files, mock_prisma = _make_real_managed_files_instance() + file_id = "litellm_proxy_unified_output_id_abc" + model_mappings = {"model-deploy-xyz": "file-output-abc"} + + for _ in range(2): + await managed_files.store_unified_file_id( + file_id=file_id, + file_object=_make_file_object(), + litellm_parent_otel_span=None, + model_mappings=model_mappings, + user_api_key_dict=_make_user_api_key_dict(), + ) + + mock_prisma.db.litellm_managedfiletable.create.assert_not_awaited() + upsert_mock = mock_prisma.db.litellm_managedfiletable.upsert + assert upsert_mock.await_count == 2 + for upsert_call in upsert_mock.await_args_list: + assert upsert_call.kwargs["where"] == {"unified_file_id": file_id} + upsert_data = upsert_call.kwargs["data"] + assert upsert_data["create"]["unified_file_id"] == file_id + assert json.loads(upsert_data["create"]["model_mappings"]) == model_mappings + assert json.loads(upsert_data["update"]["model_mappings"]) == model_mappings diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_components.py b/tests/test_litellm/integrations/otel/test_otel_v2_components.py index 5191414edeb..d856d6871a3 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_components.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_components.py @@ -451,6 +451,30 @@ def test_parse_headers(): assert providers.parse_headers("no-equals") == {} +def test_parse_headers_percent_decodes_values(): + """A percent-encoded OTLP header value reaches the exporter decoded. + + ``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded, and Grafana Cloud + documents ``Authorization=Basic%20``. Forwarding the literal ``%20`` + makes the backend reject the export as a malformed credential. + """ + token = "MTMzNzc4MzpnbGNfZXlKdklqb2lNVEl6TkNJPQ==" + assert providers.parse_headers(f"Authorization=Basic%20{token}") == {"authorization": f"Basic {token}"} + assert providers.parse_headers("x-scope-orgid=team%20a") == {"x-scope-orgid": "team a"} + + +def test_parse_headers_keeps_unencoded_values_working(): + """Values that are not percent-encoded keep parsing unchanged. + + Vendors that document a bare space, and litellm's own presets, must survive + the switch to the spec-compliant parser. Base64 padding also means a value + can contain ``=``, so only the first one may split the pair. + """ + assert providers.parse_headers("Authorization=Bearer sk-123") == {"authorization": "Bearer sk-123"} + assert providers.parse_headers("api_key=abc,space_id=xyz") == {"api_key": "abc", "space_id": "xyz"} + assert providers.parse_headers("api_key=YWJjZA==") == {"api_key": "YWJjZA=="} + + def test_otlp_traces_endpoint_normalization(): norm = providers._otlp_traces_endpoint # A base endpoint gets the signal path appended (the common OTLP env shape). @@ -487,6 +511,24 @@ def test_build_span_exporter_variants(): assert "OTLPSpanExporter" in type(http_exporter).__name__ +def test_otlp_metric_exporter_uses_cumulative_histogram_temporality(): + """Histograms must export as cumulative, not delta. + + Prometheus-backed OTLP receivers (Grafana Cloud / Mimir) reject delta + histograms with ``invalid temporality and type combination`` and drop the + entire metric batch, so a delta default silently loses every GenAI metric. + """ + from opentelemetry.sdk.metrics import Histogram + from opentelemetry.sdk.metrics.export import AggregationTemporality + + reader = providers.build_metric_reader( + OpenTelemetryV2Config(exporter="otlp_http", endpoint="http://h:4318") + ) + temporality = reader._exporter._preferred_temporality # noqa: SLF001 # exporter exposes no public accessor + + assert temporality[Histogram] is AggregationTemporality.CUMULATIVE + + def test_otlp_logs_endpoint_normalization(): norm = providers._otlp_logs_endpoint # A base endpoint gets the signal path appended (the common OTLP env shape). diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_emitter.py b/tests/test_litellm/integrations/otel/test_otel_v2_emitter.py index 6b1da4c2952..b1b1b62c820 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_emitter.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_emitter.py @@ -17,6 +17,9 @@ from litellm.integrations.otel.plumbing import context as ctx_mod # noqa: E402 from litellm.integrations.otel.plumbing import providers # noqa: E402 from litellm.integrations.otel.emitter import SpanEmitter # noqa: E402 from litellm.integrations.otel.emitter import stamp_error # noqa: E402 +from litellm.integrations.otel.mappers.utils import ( # noqa: E402 + MAX_TOOL_DEFINITION_ATTRS_PER_SPAN, +) from litellm.integrations.otel.model.payloads import ( # noqa: E402 GuardrailSpanData, LLMCallSpanData, @@ -305,3 +308,135 @@ def test_guardrail_success_span_is_unset(): ) (span,) = exporter.get_finished_spans() assert span.status.status_code is StatusCode.UNSET + + +def _tools_payload(count): + """A request declaring ``count`` tools, in the chat-completion shape.""" + return _payload( + model_parameters={ + "temperature": 0.7, + "tools": [ + { + "type": "function", + "function": { + "name": f"tool_{i}", + "description": f"description for tool {i}", + "parameters": {"type": "object", "properties": {}}, + }, + } + for i in range(count) + ], + } + ) + + +def test_many_tools_do_not_evict_core_attributes(): + """Tool definitions must never crowd core telemetry off the span. + + An agentic client declares hundreds of tools. Spelling each one out as + per-index attributes overruns the OTel SDK's 128-attribute span limit, + which evicts oldest-first and so destroys the ``gen_ai.*`` attributes + written before it. Capping the tool family keeps the core intact. + """ + engine, exporter = _engine() + data = LLMCallSpanData.from_standard_logging_payload(_tools_payload(127)) + engine.emit(SpanRole.LLM_CALL, data) + (span,) = exporter.get_finished_spans() + a = span.attributes + + assert a[GenAI.REQUEST_MODEL] == "gpt-4o" + assert a[GenAI.PROVIDER_NAME] == "openai" + assert a[GenAI.USAGE_INPUT_TOKENS] == 10 + assert a[GenAI.USAGE_OUTPUT_TOKENS] == 5 + assert a[GenAI.RESPONSE_FINISH_REASONS] == ("stop",) + assert a[f"{LiteLLM.COST_PREFIX}total"] == 0.002 + assert a["gen_ai.usage.prompt_tokens"] == 10 + + assert span.dropped_attributes == 0 + assert a[LiteLLM.TOOLS_DECLARED] == 127 + assert a["gen_ai.tool.0.name"] == "tool_0" + assert "gen_ai.tool.126.name" not in a + assert "llm.request.functions.126.name" not in a + + +def test_tool_definitions_kept_in_full_below_the_cap(): + """A handful of tools keeps full per-index detail in both vocabularies.""" + engine, exporter = _engine() + data = LLMCallSpanData.from_standard_logging_payload(_tools_payload(3)) + engine.emit(SpanRole.LLM_CALL, data) + (span,) = exporter.get_finished_spans() + a = span.attributes + + assert a[LiteLLM.TOOLS_DECLARED] == 3 + for idx in range(3): + assert a[f"gen_ai.tool.{idx}.name"] == f"tool_{idx}" + assert a[f"gen_ai.tool.{idx}.description"] == f"description for tool {idx}" + assert a[f"gen_ai.tool.{idx}.parameters"] + assert a[f"llm.request.functions.{idx}.name"] == f"tool_{idx}" + + +def _tool_span(mapper_names, tool_count): + """The exported LLM-call span for ``mapper_names`` and ``tool_count`` tools.""" + cfg = OpenTelemetryV2Config( + exporter="in_memory", + legacy_compat=True, + mapper_names=list(mapper_names), + ) + provider, exporter = providers.in_memory_provider(cfg) + engine = SpanEmitter(providers.get_tracer(provider, "litellm-test"), cfg) + engine.emit( + SpanRole.LLM_CALL, + LLMCallSpanData.from_standard_logging_payload(_tools_payload(tool_count)), + ) + (span,) = exporter.get_finished_spans() + return span + + +def _tool_definition_keys(attributes): + return [ + key + for key in attributes + if key.startswith(("gen_ai.tool.", "llm.request.functions.", "llm.tools.")) + ] + + +@pytest.mark.parametrize( + "mapper_names", + [ + ["genai"], + ["genai", "openinference"], + ["genai", "openinference", "langfuse", "weave", "langtrace"], + ], +) +def test_tool_definitions_stay_within_one_span_wide_budget(mapper_names): + """Every supported composition has to leave core telemetry on the span. + + Each vocabulary spells the same tools out under its own keys, so an + allowance handed to each mapper separately multiplies by the number of + configured vocabularies and reaches the attribute limit again. Arize and + Phoenix already layer OpenInference on top of the default two, and every + vendor vocabulary can be listed at once. One budget shared across them all + is what keeps the total bounded. + """ + span = _tool_span(mapper_names, 127) + a = span.attributes + + assert span.dropped_attributes == 0 + assert a[GenAI.REQUEST_MODEL] == "gpt-4o" + assert a[GenAI.PROVIDER_NAME] == "openai" + assert a[GenAI.USAGE_INPUT_TOKENS] == 10 + assert a[GenAI.USAGE_OUTPUT_TOKENS] == 5 + assert a[f"{LiteLLM.COST_PREFIX}total"] == 0.002 + assert a[LiteLLM.TOOLS_DECLARED] == 127 + + emitted = _tool_definition_keys(a) + assert emitted, "some tool detail should survive in every composition" + assert len(emitted) <= MAX_TOOL_DEFINITION_ATTRS_PER_SPAN + + +def test_vendor_tool_definitions_are_truncated_not_dropped(): + """The OpenInference vocabulary keeps its leading tools and loses the tail.""" + a = _tool_span(["genai", "openinference"], 127).attributes + assert a["llm.tools.0.tool.name"] == "tool_0" + assert a["llm.tools.0.tool.json_schema"] + assert "llm.tools.126.tool.name" not in a diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py index 02954578644..41c02501acc 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py @@ -2120,9 +2120,9 @@ def test_valid_metric_filter_records_six_metrics(monkeypatch): assert _emitted_metric_names(reader) == { "gen_ai.client.operation.duration", "gen_ai.client.token.usage", - "gen_ai.client.token.cost", - "gen_ai.client.response.time_to_first_token", - "gen_ai.client.response.time_per_output_token", + "gen_ai.usage.cost", + "gen_ai.server.time_to_first_token", + "gen_ai.server.time_per_output_token", "gen_ai.client.response.duration", } diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py index 29067f91b5a..e1b8e4b5721 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py @@ -12,9 +12,17 @@ raises out of ``GenAIMetricRecorder.record`` -- asserted directly at the recorde layer -- and the logger turns that raise into a single ERROR ("metrics disabled") plus a quiet no-op for the rest of the process, asserted at the logger layer so the misconfig never breaks a request nor spams a log line per request. + +The failure path is driven the same way, through the real +``OpenTelemetryV2.async_log_failure_event``: a failed call records +``gen_ai.client.operation.duration`` and nothing else, tagged with ``error.type``, +and a success driven through the same reader keeps a datapoint whose attributes are +byte-for-byte what it had before the failure path existed -- the guard for every +dashboard already querying that histogram. """ import asyncio +import json from datetime import datetime, timedelta import pytest @@ -25,6 +33,9 @@ from opentelemetry.sdk.metrics import MeterProvider # noqa: E402 from opentelemetry.sdk.metrics.export import InMemoryMetricReader # noqa: E402 import litellm # noqa: E402 +from litellm.constants import ( # noqa: E402 + LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL, +) from litellm.integrations.otel.logger import OpenTelemetryV2 # noqa: E402 from litellm.integrations.otel.model.config import ( # noqa: E402 OpenTelemetryV2Config, @@ -39,9 +50,9 @@ from litellm.integrations.otel.plumbing.providers import ( # noqa: E402 OPERATION_DURATION = "gen_ai.client.operation.duration" TOKEN_USAGE = "gen_ai.client.token.usage" -TOKEN_COST = "gen_ai.client.token.cost" -TIME_TO_FIRST_TOKEN = "gen_ai.client.response.time_to_first_token" -TIME_PER_OUTPUT_TOKEN = "gen_ai.client.response.time_per_output_token" +TOKEN_COST = "gen_ai.usage.cost" +TIME_TO_FIRST_TOKEN = "gen_ai.server.time_to_first_token" +TIME_PER_OUTPUT_TOKEN = "gen_ai.server.time_per_output_token" RESPONSE_DURATION = "gen_ai.client.response.duration" ALL_METRICS = frozenset( @@ -57,15 +68,17 @@ ALL_METRICS = frozenset( TOKEN_TYPE = "gen_ai.token.type" MODEL_KEY = "gen_ai.request.model" +OPERATION_KEY = "gen_ai.operation.name" +PROVIDER_NAME_KEY = "gen_ai.provider.name" +SYSTEM_KEY = "gen_ai.system" -# Each is a member of VALID_METRIC_ATTRIBUTE_NAMES and is stamped on the metric -# by default (proven by the no-filter test below). -HIGH_CARDINALITY_KEYS = ( +# Keys inside the ceiling that an operator's filter must still be able to remove. +# Every one is bounded, so it survives the ceiling and only the operator's own +# exclude_list takes it off; that is what makes the filter tests non-vacuous. +FILTERABLE_KEYS = ( "hidden_params", "metadata.user_api_key_hash", - "metadata.requester_ip_address", - "metadata.requester_metadata", - "metadata.applied_guardrails", + "metadata.user_api_key_team_id", ) PROMPT_TOKENS = 137 @@ -73,18 +86,25 @@ COMPLETION_TOKENS = 89 RESPONSE_COST = 0.0023 -def _build_call(stream: bool = True): +def _build_call( + stream: bool = True, + provider: str | None = "openai", + call_type: str = "completion", +): """A captured success-call (kwargs, response_obj, start, end) that exercises every one of the six metrics: usage for token.usage, response_cost for cost, - streaming + timing for the response-time histograms.""" + streaming + timing for the response-time histograms. + + ``provider=None`` omits ``custom_llm_provider`` entirely, reproducing a call + litellm could not attribute to a provider.""" start = datetime(2026, 6, 12, 12, 0, 0) api_call_start = start + timedelta(seconds=0.1) completion_start = start + timedelta(seconds=0.5) end = start + timedelta(seconds=1.0) kwargs = { "model": "gpt-4o-mini", - "call_type": "completion", - "litellm_params": {"custom_llm_provider": "openai"}, + "call_type": call_type, + "litellm_params": ({"custom_llm_provider": provider} if provider is not None else {}), "optional_params": {"stream": stream}, "response_cost": RESPONSE_COST, "api_call_start_time": api_call_start, @@ -93,6 +113,7 @@ def _build_call(stream: bool = True): "standard_logging_object": { "metadata": { "user_api_key_hash": "hash-abc123", + "user_api_key_team_id": "team-1", "requester_ip_address": "10.0.0.7", "requester_metadata": {"team": "alpha", "tier": "gold"}, "applied_guardrails": ["pii", "toxicity"], @@ -131,7 +152,7 @@ def _metrics_by_name(reader): return out -def _drive_success(reader, callback_settings_attributes=None): +def _drive_success(reader, callback_settings_attributes=None, **call_overrides): """Construct a metrics-on logger, optionally populate callback_settings AFTER construction (mirroring the proxy ordering), run the real success hook.""" logger = _logger(reader, enable_metrics=True) @@ -141,7 +162,7 @@ def _drive_success(reader, callback_settings_attributes=None): "otel": {"attributes": callback_settings_attributes} } try: - kwargs, response_obj, start, end = _build_call() + kwargs, response_obj, start, end = _build_call(**call_overrides) asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) finally: litellm.callback_settings = previous @@ -205,14 +226,14 @@ def test_metrics_off_by_default_records_nothing(): def test_exclude_list_strips_high_cardinality_across_metrics(): - """exclude_list set AFTER construction (the proxy path) removes every - high-cardinality key from more than one metric while the low-cardinality - model attribute survives.""" + """exclude_list set AFTER construction (the proxy path) removes every listed + key from more than one metric while the low-cardinality model attribute + survives.""" metrics = _drive_success( InMemoryMetricReader(), - callback_settings_attributes={"exclude_list": list(HIGH_CARDINALITY_KEYS)}, + callback_settings_attributes={"exclude_list": list(FILTERABLE_KEYS)}, ) - excluded = set(HIGH_CARDINALITY_KEYS) + excluded = set(FILTERABLE_KEYS) for name in (OPERATION_DURATION, TOKEN_USAGE): points = metrics[name] @@ -241,18 +262,226 @@ def test_include_list_allows_only_listed_attributes(): assert set(dp.attributes.keys()) - {TOKEN_TYPE} == allowed -def test_no_filter_keeps_high_cardinality_keys(): - """Backward compatibility: without an attributes config every high-cardinality - key the call carries is still stamped, so the filter tests above prove a real - removal rather than a key that was never present.""" +def test_no_filter_still_keeps_the_filterable_keys(): + """Without an attributes config every key the filter tests remove is present, + so those tests prove a real removal rather than a key that was never there.""" metrics = _drive_success(InMemoryMetricReader()) - expected = set(HIGH_CARDINALITY_KEYS) + expected = set(FILTERABLE_KEYS) for name in (OPERATION_DURATION, TOKEN_USAGE): for dp in metrics[name]: assert expected.issubset(set(dp.attributes.keys())) +def test_a_metric_ineligible_filter_name_is_reported_not_silently_dropped(caplog): + """Naming a metric-ineligible attribute in a filter has to say so. + + The shared validator accepts every span attribute name, so an operator can put + one in an ``include_list``, get nothing for it, and have no way to tell that from + a value that happened to be absent. The ceiling is deliberate, but silent is what + makes it a support ticket. + """ + with caplog.at_level("WARNING"): + _drive_success( + InMemoryMetricReader(), + callback_settings_attributes={ + "include_list": [MODEL_KEY, "metadata.requester_ip_address"] + }, + ) + + reported = [ + r.getMessage().split(" cannot be a metric attribute")[0].removeprefix("OTel metrics: ") + for r in caplog.records + if r.levelname == "WARNING" and "cannot be a metric attribute" in r.getMessage() + ] + assert reported == ["metadata.requester_ip_address"], reported + + +def test_two_calls_differing_only_per_request_share_one_series(): + """The whole point of the ceiling: metric cardinality must not grow with traffic. + + Every field here moves on every real request -- the response cost, the call id, + the cache key, the provider's remaining-rate-limit headers -- and each one used + to reach the datapoint inside a single ``hidden_params`` label. A unique label + value is a new time series, so each of the six instruments minted one series per + request, which is both a Grafana Cloud bill proportional to traffic and a + histogram that cannot be aggregated. Identical attribute sets is what "one + series" means to the SDK. + """ + reader = InMemoryMetricReader() + logger = _logger(reader, enable_metrics=True) + + for index, cost in enumerate((RESPONSE_COST, RESPONSE_COST * 3)): + kwargs, response_obj, start, end = _build_call() + kwargs["response_cost"] = cost + kwargs["standard_logging_object"]["hidden_params"] = { + "model_id": "m-1", + # A documented per-call parameter, so it varies here on purpose: the same + # deployment reached under a caller-chosen base must not split the series. + "api_base": f"https://proxy-{index}.example.com/v1", + "litellm_call_id": f"call-{index}", + "cache_key": f"cache-{index}", + "response_cost": cost, + "litellm_overhead_time_ms": 1.5 + index, + "usage_object": {"prompt_tokens": index, "completion_tokens": index}, + "additional_headers": {"x_ratelimit_remaining_requests": 100 - index}, + } + asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) + + for name in ALL_METRICS: + attribute_sets = { + tuple(sorted((k, v) for k, v in dp.attributes.items() if k != TOKEN_TYPE)) + for dp in _metrics_by_name(reader)[name] + } + assert len(attribute_sets) == 1, f"{name} split into {len(attribute_sets)} series across 2 requests" + + +def test_hidden_params_label_carries_only_bounded_deployment_fields(): + """``hidden_params`` survives the ceiling, but only as the deployment identity. + + ``model_id`` is the router's deployment id, bounded by the deployment list, and is + what a per-deployment dashboard reads. Everything else in the object is + per-request or caller-chosen and belongs on the span, which already carries it. + ``api_base`` is excluded despite naming the same deployment: it is a documented + per-call parameter, so a caller varying it would restore the per-request + cardinality this cap exists to remove. + """ + kwargs, response_obj, start, end = _build_call() + kwargs["standard_logging_object"]["hidden_params"] = { + "model_id": "m-1", + "api_base": "https://api.openai.com/v1", + "litellm_call_id": "abc", + "cache_key": "ck-1", + "response_cost": RESPONSE_COST, + } + reader = InMemoryMetricReader() + logger = _logger(reader, enable_metrics=True) + asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) + + label = _metrics_by_name(reader)[OPERATION_DURATION][0].attributes["hidden_params"] + assert json.loads(label) == {"model_id": "m-1"} + + +def test_success_attributes_are_capped_at_the_ceiling(): + """The success path carries exactly the ceiling, no client-supplied attributes. + + The fixture deliberately sets every excluded key, so this asserts a real removal + rather than keys that were never present. + """ + kwargs, response_obj, start, end = _build_call() + metadata = kwargs["standard_logging_object"]["metadata"] + metadata.update( + { + "spend_logs_metadata": {"cost_center": "abc"}, + "user_api_key_end_user_id": "end-user-1", + "user_api_key_user_email": "someone@example.com", + } + ) + reader = InMemoryMetricReader() + logger = _logger(reader, enable_metrics=True) + asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) + metrics = _metrics_by_name(reader) + + for name in ALL_METRICS: + for dp in metrics[name]: + leaked = set(dp.attributes) - set(BOUNDED_KEYS) - {TOKEN_TYPE} + assert not leaked, f"{name} leaked {leaked}" +def test_provider_is_labelled_with_semconv_provider_name(): + """Every recorded point carries gen_ai.provider.name holding the semconv + provider value (bedrock -> aws.bedrock), the key the GenAI convention and the + dashboards built on it query. The deprecated gen_ai.system spelling alone is + unreadable to them.""" + metrics = _drive_success(InMemoryMetricReader(), provider="bedrock") + + for name in ALL_METRICS: + points = metrics[name] + assert points, f"{name} was not recorded" + for dp in points: + assert dp.attributes[PROVIDER_NAME_KEY] == "aws.bedrock" + + +def test_deprecated_gen_ai_system_is_dual_emitted_verbatim(): + """gen_ai.system keeps its raw litellm provider value alongside the new key + for one release, so a dashboard already filtering on it keeps matching. Its + value must not be swapped for the mapped one, which would break exactly the + queries the dual emission exists to protect.""" + metrics = _drive_success(InMemoryMetricReader(), provider="bedrock") + + for dp in metrics[OPERATION_DURATION]: + assert dp.attributes[SYSTEM_KEY] == "bedrock" + assert dp.attributes[PROVIDER_NAME_KEY] == "aws.bedrock" + + +def test_no_provider_attribute_when_provider_is_absent(): + """A call litellm could not attribute to a provider carries no provider label + at all. A placeholder value ("Unknown") would mint a permanent series that + aggregates every unattributable request and that no operator can act on.""" + metrics = _drive_success(InMemoryMetricReader(), provider=None) + + for name in ALL_METRICS: + points = metrics[name] + assert points, f"{name} was not recorded" + for dp in points: + keys = set(dp.attributes.keys()) + assert PROVIDER_NAME_KEY not in keys + assert SYSTEM_KEY not in keys + assert "Unknown" not in set(dp.attributes.values()) + + +def test_vector_store_search_is_not_labelled_as_chat(): + """A vector-store search records under gen_ai.operation.name=retrieval, so its + latency and cost stay out of the chat series.""" + metrics = _drive_success(InMemoryMetricReader(), call_type="avector_store_search") + + for name in (OPERATION_DURATION, TOKEN_COST): + for dp in metrics[name]: + assert dp.attributes[OPERATION_KEY] == "retrieval" + + +@pytest.mark.parametrize( + "call_type,expected", + [ + ("avector_store_create", "litellm.vector_store_management"), + ("avector_store_delete", "litellm.vector_store_management"), + ("avector_store_file_create", "litellm.vector_store_file_management"), + ("avector_store_file_list", "litellm.vector_store_file_management"), + ], +) +def test_vector_store_management_is_not_labelled_as_chat(call_type, expected): + """Store and file management reach the recorder through the same success hook as a + completion, so leaving them unmapped kept billing- and latency-relevant admin calls + inside the chat series.""" + metrics = _drive_success(InMemoryMetricReader(), call_type=call_type) + + for dp in metrics[OPERATION_DURATION]: + assert dp.attributes[OPERATION_KEY] == expected + + +@pytest.mark.parametrize("call_type", ["asend_message", "asend_message_streaming"]) +def test_agent_message_is_not_labelled_as_chat(call_type): + """An A2A agent send records under gen_ai.operation.name=invoke_agent, streamed or + not. The streaming iterator dispatches the same success handlers under its own + ``asend_message_streaming`` call type, so an unmapped streaming spelling puts every + streamed agent turn's latency and cost back into the chat series.""" + metrics = _drive_success(InMemoryMetricReader(), call_type=call_type) + + for name in (OPERATION_DURATION, TOKEN_COST): + for dp in metrics[name]: + assert dp.attributes[OPERATION_KEY] == "invoke_agent" + + +def test_provider_name_is_filterable(): + """gen_ai.provider.name is a member of the metric-attribute allowlist, so an + operator can include or exclude it; an unlisted name raises instead.""" + metrics = _drive_success( + InMemoryMetricReader(), + callback_settings_attributes={"include_list": [PROVIDER_NAME_KEY]}, + ) + + for dp in metrics[OPERATION_DURATION]: + assert set(dp.attributes.keys()) == {PROVIDER_NAME_KEY} + + def test_metrics_reach_operator_configured_global_provider(monkeypatch): """Regression: with no meter provider injected, the six gen_ai.client.* histograms must record through the operator's globally configured @@ -331,3 +560,262 @@ def test_token_type_rejected_from_either_list(attributes, monkeypatch): # the specific reason so dropping that guard (and falling through to "unknown # attribute name") is caught. assert "discriminator" in str(exc_info.value) + + +# --- failure path ------------------------------------------------------------ # + +ERROR_TYPE = "error.type" +ERROR_CLASS = "RateLimitError" +FAILURE_DURATION_S = 1.0 + +# Attributes a failure datapoint must never carry. Each is either supplied by the +# caller (so a caller could mint a fresh series per request, and a failure costs +# them no provider spend) or varies per request, or is PII duplicating an id that +# is already on the series. +UNBOUNDED_KEYS = ( + "metadata.requester_metadata", + "metadata.requester_ip_address", + "metadata.spend_logs_metadata", + "metadata.user_api_key_end_user_id", + "metadata.user_api_key_user_email", +) + +# The exact set a datapoint may carry on either path: the operation, the +# operator-provisioned identity, and the deployment that served it. +BOUNDED_KEYS = ( + "hidden_params", + "gen_ai.operation.name", + "gen_ai.provider.name", + "gen_ai.system", + "gen_ai.request.model", + "gen_ai.framework", + "metadata.user_api_key_hash", + "metadata.user_api_key_alias", + "metadata.user_api_key_team_id", + "metadata.user_api_key_team_alias", + "metadata.user_api_key_org_id", + "metadata.user_api_key_user_id", +) + + +def _build_failure( + *, + error_information=None, + exception=None, + no_upstream_call=False, +): + """A captured failure-call ``(kwargs, start, end)``. + + Mirrors what litellm actually hands ``async_log_failure_event``: no + ``response_obj`` at all, but the streaming timings and the recovered + ``response_cost`` a mid-stream failure still carries -- so routing the failure + path through the full success recorder would show up here as extra series + rather than passing unnoticed. The metadata carries both the bounded identity + keys and every caller-supplied / per-request key, so the allowlist test below + proves a real removal rather than a key that was never there. + """ + start = datetime(2026, 6, 12, 12, 0, 0) + api_call_start = start + timedelta(seconds=0.1) + completion_start = start + timedelta(seconds=0.5) + end = start + timedelta(seconds=FAILURE_DURATION_S) + standard_logging_object = { + "status": "failure", + "metadata": { + "user_api_key_hash": "hash-abc123", + "user_api_key_alias": "alias-abc", + "user_api_key_team_id": "team-1", + "user_api_key_team_alias": "team-alpha", + "user_api_key_org_id": "org-1", + "user_api_key_user_id": "user-1", + "user_api_key_user_email": "user@example.com", + "user_api_key_end_user_id": "end-user-42", + "requester_ip_address": "10.0.0.7", + "requester_metadata": {"trace": "caller-supplied-unique-value"}, + "spend_logs_metadata": {"ticket": "caller-supplied-unique-value"}, + }, + "hidden_params": { + "litellm_call_id": "abc", + "model_id": "m-1", + "api_base": "https://api.openai.com/v1", + }, + } + if error_information is not None: + standard_logging_object["error_information"] = error_information + kwargs = { + "model": "gpt-4o-mini", + "call_type": "completion", + "litellm_params": {"custom_llm_provider": "openai"}, + "optional_params": {"stream": True}, + "response_cost": RESPONSE_COST, + "api_call_start_time": api_call_start, + "completion_start_time": completion_start, + "end_time": end, + "standard_logging_object": standard_logging_object, + } + if exception is not None: + kwargs["exception"] = exception + if no_upstream_call: + kwargs[LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL] = True + return kwargs, start, end + + +def _drive_failure(reader, callback_settings_attributes=None, **failure_kwargs): + logger = _logger(reader, enable_metrics=True) + previous = litellm.callback_settings + if callback_settings_attributes is not None: + litellm.callback_settings = {"otel": {"attributes": callback_settings_attributes}} + try: + kwargs, start, end = _build_failure(**failure_kwargs) + asyncio.run(logger.async_log_failure_event(kwargs, None, start, end)) + finally: + litellm.callback_settings = previous + return _metrics_by_name(reader) + + +def test_failure_records_only_the_duration_histogram(): + """A failed call contributes to gen_ai.client.operation.duration -- before this + existed a failure recorded nothing at all, so the histogram measured only the + traffic that survived. It contributes to nothing else: the other five + instruments describe a completed generation, and the call carries a streaming + timing pair and a recovered response_cost that would light four of them up if + the failure were routed through the success recorder.""" + metrics = _drive_failure( + InMemoryMetricReader(), + error_information={"error_class": ERROR_CLASS, "error_code": "429"}, + ) + + assert set(metrics.keys()) == {OPERATION_DURATION} + points = metrics[OPERATION_DURATION] + assert len(points) == 1 + assert points[0].count == 1 + assert points[0].sum == pytest.approx(FAILURE_DURATION_S) + assert points[0].attributes[ERROR_TYPE] == ERROR_CLASS + + +def test_success_and_failure_are_separable_and_success_attributes_unchanged(): + """The pooled histogram stays queryable per outcome, and the existing + dashboards keep working. + + A success and a failure through one reader must land on two distinct series -- + one with error.type, one without -- so a failure-rate panel is expressible and + an operator can still get success-only latency by filtering error.type="". The + success datapoint's attribute map must be byte-for-byte the map a success-only + run produces, which is what stops the new attribute from leaking onto the + series every current query reads.""" + baseline_reader = InMemoryMetricReader() + baseline = _drive_success(baseline_reader) + baseline_points = baseline[OPERATION_DURATION] + assert len(baseline_points) == 1 + baseline_attributes = dict(baseline_points[0].attributes) + + reader = InMemoryMetricReader() + logger = _logger(reader, enable_metrics=True) + ok_kwargs, response_obj, ok_start, ok_end = _build_call() + asyncio.run(logger.async_log_success_event(ok_kwargs, response_obj, ok_start, ok_end)) + bad_kwargs, bad_start, bad_end = _build_failure(error_information={"error_class": ERROR_CLASS}) + asyncio.run(logger.async_log_failure_event(bad_kwargs, None, bad_start, bad_end)) + + points = metrics = _metrics_by_name(reader)[OPERATION_DURATION] + assert len(points) == 2, f"success and failure collapsed into {len(points)} series: {metrics}" + succeeded = [dp for dp in points if ERROR_TYPE not in dp.attributes] + failed = [dp for dp in points if dp.attributes.get(ERROR_TYPE) == ERROR_CLASS] + assert len(succeeded) == 1 and len(failed) == 1 + assert dict(succeeded[0].attributes) == baseline_attributes + + +def test_failure_attributes_are_a_bounded_allowlist(): + """A failure datapoint carries exactly the bounded allowlist plus error.type. + + A failed request needs no provider spend, so nothing rate-limits a caller who + puts a unique value into an attribute they control and mints one histogram + series per request. The same payload is driven through the success path first, + which does carry those keys, so this asserts a real removal on the failure path + rather than keys that were never present. The exact-set assertion is the guard + against the natural refactor of "just reuse _common_attributes".""" + reader = InMemoryMetricReader() + logger = _logger(reader, enable_metrics=True) + kwargs, start, end = _build_failure(error_information={"error_class": ERROR_CLASS}) + usage = {"usage": {"prompt_tokens": 1, "completion_tokens": 1}} + asyncio.run(logger.async_log_success_event(kwargs, usage, start, end)) + asyncio.run(logger.async_log_failure_event(kwargs, None, start, end)) + + points = _metrics_by_name(reader)[OPERATION_DURATION] + succeeded = next(dp for dp in points if ERROR_TYPE not in dp.attributes) + failed = next(dp for dp in points if ERROR_TYPE in dp.attributes) + + supplied = set(kwargs["standard_logging_object"]["metadata"]) + missing = {key for key in UNBOUNDED_KEYS if key.removeprefix("metadata.") not in supplied} + assert not missing, f"fixture never carried {missing}, so the exclusion below proves nothing" + leaked = set(UNBOUNDED_KEYS) & set(failed.attributes) + assert not leaked, f"failure datapoint leaked unbounded attributes: {leaked}" + assert set(failed.attributes) == set(BOUNDED_KEYS) | {ERROR_TYPE} + assert json.loads(failed.attributes["hidden_params"]) == {"model_id": "m-1"} + + +def test_operator_filter_can_still_narrow_the_failure_allowlist(): + """The allowlist is a ceiling, not a floor: an exclude_list an operator sets + still removes a listed key from the failure series.""" + metrics = _drive_failure( + InMemoryMetricReader(), + callback_settings_attributes={"exclude_list": ["metadata.user_api_key_hash"]}, + error_information={"error_class": ERROR_CLASS}, + ) + attributes = metrics[OPERATION_DURATION][0].attributes + assert "metadata.user_api_key_hash" not in attributes + assert attributes[ERROR_TYPE] == ERROR_CLASS + assert attributes[MODEL_KEY] == "gpt-4o-mini" + + +@pytest.mark.parametrize( + "failure_kwargs, expected", + [ + ({"error_information": {"error_class": ERROR_CLASS, "error_code": "429"}}, ERROR_CLASS), + ({"error_information": {"error_code": "429"}}, "429"), + ({"exception": ValueError("boom")}, "ValueError"), + ({}, "_OTHER"), + ], + ids=["error_class", "error_code_only", "exception_fallback", "unclassifiable"], +) +def test_error_type_is_bounded_and_falls_back(failure_kwargs, expected): + """error.type is always a bounded value: the mapped exception's class name, the + provider status code, the raw exception's class name, or the semconv _OTHER + fallback. Never the exception message, which is unbounded.""" + metrics = _drive_failure(InMemoryMetricReader(), **failure_kwargs) + assert metrics[OPERATION_DURATION][0].attributes[ERROR_TYPE] == expected + + +def test_include_list_cannot_strip_error_type(): + """error.type is a structural discriminator like gen_ai.token.type: an + include_list that does not mention it must not merge the failure series back + into the success series, so it is stamped after the filter runs.""" + metrics = _drive_failure( + InMemoryMetricReader(), + callback_settings_attributes={"include_list": [MODEL_KEY]}, + error_information={"error_class": ERROR_CLASS}, + ) + attributes = metrics[OPERATION_DURATION][0].attributes + assert dict(attributes) == {MODEL_KEY: "gpt-4o-mini", ERROR_TYPE: ERROR_CLASS} + + +def test_proxy_gate_rejection_records_no_duration(): + """A synthetic proxy-gate failure log (auth / rate-limit rejection) never made + an upstream call, so its wall time is not a GenAI operation's duration; it is + skipped for the same reason it gets no span. Recording it would pull the + histogram toward the proxy's own latency. + + Both failures go through one reader so the assertion is that exactly the + upstream one landed, rather than the vacuous "nothing was recorded" a + failure path that records nothing at all would also satisfy.""" + reader = InMemoryMetricReader() + logger = _logger(reader, enable_metrics=True) + gate_kwargs, gate_start, gate_end = _build_failure( + error_information={"error_class": "AuthenticationError"}, + no_upstream_call=True, + ) + asyncio.run(logger.async_log_failure_event(gate_kwargs, None, gate_start, gate_end)) + upstream_kwargs, upstream_start, upstream_end = _build_failure(error_information={"error_class": ERROR_CLASS}) + asyncio.run(logger.async_log_failure_event(upstream_kwargs, None, upstream_start, upstream_end)) + + points = _metrics_by_name(reader)[OPERATION_DURATION] + assert [dp.attributes[ERROR_TYPE] for dp in points] == [ERROR_CLASS] + assert points[0].count == 1 diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py index 71be28ea485..612ac1e5113 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py @@ -1,8 +1,13 @@ """Tests for the OTel v2 sources of truth: span registry, semconv keys, config, and the typed StandardLoggingPayload adapter. These need no OTel SDK.""" +import logging +import re +from pathlib import Path + import pytest +import litellm from litellm.integrations.otel import ( BAGGAGE_PROMOTED_KEYS, DB, @@ -208,6 +213,103 @@ def test_operation_resolution(): assert resolve_operation("call_mcp_tool") is GenAIOperation.EXECUTE_TOOL +@pytest.mark.parametrize("call_type", ["vector_store_search", "avector_store_search"]) +def test_vector_store_search_is_a_retrieval_operation(call_type): + """A vector-store search is a retrieval, so its duration and cost must not + land in the chat series that dashboards read latency off.""" + assert resolve_operation(call_type) is GenAIOperation.RETRIEVAL + assert resolve_operation(call_type).value == "retrieval" + + +@pytest.mark.parametrize("call_type", ["query", "aquery"]) +def test_rag_query_is_a_retrieval_operation(call_type): + """``/rag/query`` reaches the same recorder as a vector-store search and is the + same operation, so it must not be the one retrieval surface left reading as chat.""" + assert resolve_operation(call_type) is GenAIOperation.RETRIEVAL + + +@pytest.mark.parametrize( + "call_type", + [ + f"{prefix}vector_store_{verb}" + for verb in ("create", "retrieve", "list", "update", "delete") + for prefix in ("", "a") + ], +) +def test_vector_store_management_is_not_chat(call_type): + """The store lifecycle calls are not GenAI client operations and the convention + names nothing for them, so they take a vendor value rather than defaulting into + the chat series.""" + assert resolve_operation(call_type) is GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT + assert resolve_operation(call_type).value == "litellm.vector_store_management" + + +@pytest.mark.parametrize( + "call_type", + [ + f"{prefix}vector_store_file_{verb}" + for verb in ("create", "list", "retrieve", "content", "update", "delete") + for prefix in ("", "a") + ], +) +def test_vector_store_file_management_is_not_chat(call_type): + """The file operations are a distinct REST resource from the store lifecycle, so + they get their own vendor value instead of sharing one bucket.""" + assert resolve_operation(call_type) is GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT + assert resolve_operation(call_type).value == "litellm.vector_store_file_management" + + +def test_vendor_operation_values_are_namespaced(): + """A vendor value must stay under the ``litellm.`` prefix: an unprefixed invented + name could collide with a value the convention adds later, silently changing what + a conformant consumer thinks it is reading.""" + vendor = [op for op in GenAIOperation if op.name.startswith("LITELLM_")] + assert vendor, "no vendor operation values defined" + assert all(op.value.startswith("litellm.") for op in vendor) + + +@pytest.mark.parametrize("call_type", ["send_message", "asend_message", "asend_message_streaming"]) +def test_agent_message_is_an_invoke_agent_operation(call_type): + """An agent (A2A) message send is an agent invocation, not a chat completion. + + The streaming spelling counts: ``_build_streaming_logging_obj`` in + ``litellm/a2a_protocol/main.py`` stamps ``asend_message_streaming`` on the + logging object the streaming iterator dispatches success handlers with, so a + missing entry sends every streamed agent turn into the chat series. There is + no sync spelling because A2A streaming is async-only. + """ + assert resolve_operation(call_type) is GenAIOperation.INVOKE_AGENT + assert resolve_operation(call_type).value == "invoke_agent" + + +def test_every_call_type_the_a2a_package_stamps_is_an_agent_operation(): + """Pins the map to the call types the A2A code actually stamps on its logging + objects. A new spelling added there without a map entry fails here instead of + quietly landing in the chat series, which is how the streaming one was missed.""" + a2a_package = Path(litellm.__file__).parent / "a2a_protocol" + stamped = { + call_type + for source in a2a_package.rglob("*.py") + for call_type in re.findall(r'call_type="([^"]+)"', source.read_text()) + } + assert stamped, "no call_type literals found in litellm/a2a_protocol" + unmapped = { + call_type: resolve_operation(call_type).value + for call_type in stamped + if resolve_operation(call_type) is not GenAIOperation.INVOKE_AGENT + } + assert not unmapped, f"add these to _OPERATION_BY_CALL_TYPE: {unmapped}" + + +def test_unmapped_call_type_falls_back_to_chat_loudly(caplog): + """The fallback still labels the series ``chat`` so it is never unlabelled, + but it says so at debug: a silent default is how retrieval and agent calls + ended up in the chat charts in the first place.""" + with caplog.at_level(logging.DEBUG, logger="LiteLLM"): + assert resolve_operation("some_future_call_type") is GenAIOperation.CHAT + assert any("some_future_call_type" in record.getMessage() for record in caplog.records) + + # --- MCP tool-call (source of truth #1/#2/#3) ------------------------------- # diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index 7ffd09b931f..05205cb76f2 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -412,6 +412,40 @@ class TestOpenTelemetryProviderInitialization(unittest.TestCase): current_provider is existing_provider ), "Existing TracerProvider should be respected and not overridden" + @patch.dict( + os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True + ) + def test_init_metrics_creates_instruments_under_their_published_names(self): + """ + The v1 engine's instrument names are a public contract. + + Every name here is what a backend queries: four are GenAI semantic + conventions and gen_ai.usage.cost is the name backends query for spend. + A rename is breaking for anyone charting them, so it has to be a + deliberate edit to the shared Metric constants and to this list, never + a silent drift between the v1 and v2 engines. + """ + from opentelemetry import metrics + + metrics.set_meter_provider(MeterProvider(metric_readers=[InMemoryMetricReader()])) + otel_integration = OpenTelemetry(config=OpenTelemetryConfig.from_env()) + + assert { + otel_integration._operation_duration_histogram.name, + otel_integration._token_usage_histogram.name, + otel_integration._cost_histogram.name, + otel_integration._time_to_first_token_histogram.name, + otel_integration._time_per_output_token_histogram.name, + otel_integration._response_duration_histogram.name, + } == { + "gen_ai.client.operation.duration", + "gen_ai.client.token.usage", + "gen_ai.usage.cost", + "gen_ai.server.time_to_first_token", + "gen_ai.server.time_per_output_token", + "gen_ai.client.response.duration", + } + @patch.dict( os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True ) diff --git a/tests/test_litellm/integrations/test_prometheus_service_tier_label.py b/tests/test_litellm/integrations/test_prometheus_service_tier_label.py new file mode 100644 index 00000000000..9d702b19c6c --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_service_tier_label.py @@ -0,0 +1,232 @@ +""" +Unit tests for the service_tier Prometheus label on latency and spend metrics. + +Covers the label being declared on the metrics that carry it, the precedence +between the tier a provider served and the tier a caller requested, and the +end-to-end emit wiring through async_log_success_event. + +Run with: + uv run pytest tests/test_litellm/integrations/test_prometheus_service_tier_label.py -v +""" + +import datetime + +import pytest + +from litellm.integrations.prometheus import ( + KNOWN_REQUEST_SERVICE_TIERS, + PrometheusLogger, + get_service_tier_from_standard_logging_payload, +) +from litellm.types.integrations.prometheus import ( + PrometheusMetricLabels, + UserAPIKeyLabelNames, + UserAPIKeyLabelValues, +) + +SERVICE_TIER_METRICS = [ + "litellm_llm_api_latency_metric", + "litellm_llm_api_time_to_first_token_metric", + "litellm_request_total_latency_metric", + "litellm_spend_metric", +] + + +def _clear_prometheus_registry() -> None: + from prometheus_client import REGISTRY + + for collector in list(REGISTRY._collector_to_names.keys()): + try: + REGISTRY.unregister(collector) + except Exception: + pass + + +def _collected_samples(metric_name: str): + from prometheus_client import REGISTRY + + return [sample for metric in REGISTRY.collect() for sample in metric.samples if sample.name == metric_name] + + +def _standard_logging_payload( + response: object = None, + usage_object: object = None, + model_parameters: object = None, +) -> dict: + return { + "id": "t", + "call_type": "completion", + "response_cost": 0.001, + "status": "success", + "total_tokens": 30, + "prompt_tokens": 20, + "completion_tokens": 10, + "startTime": 1.0, + "endTime": 2.0, + "completionStartTime": 1.5, + "model": "gpt-4o-mini", + "model_id": "model-123", + "model_group": "gpt-4o-mini", + "api_base": "https://api.openai.com", + "custom_llm_provider": "openai", + "request_tags": [], + "end_user": None, + "cache_hit": False, + "stream": True, + "response": response, + "model_parameters": model_parameters, + "metadata": { + "user_api_key_hash": "h", + "user_api_key_alias": "a", + "user_api_key_team_id": "t", + "user_api_key_team_alias": "ta", + "user_api_key_user_id": "u", + "user_api_key_user_email": "e@x.com", + "user_api_key_org_id": None, + "user_api_key_org_alias": None, + "requester_metadata": None, + "user_api_key_end_user_id": None, + "usage_object": usage_object, + }, + "hidden_params": {"litellm_overhead_time_ms": None, "additional_headers": None}, + } + + +def test_service_tier_label_declared_on_latency_and_spend_metrics(): + for metric_name in SERVICE_TIER_METRICS: + labels = PrometheusMetricLabels.get_labels(metric_name) + assert UserAPIKeyLabelNames.SERVICE_TIER.value in labels, f"{metric_name} should carry the service_tier label" + + +def test_user_api_key_label_values_carries_service_tier(): + values = UserAPIKeyLabelValues(service_tier="flex") + + assert values.service_tier == "flex" + assert values.model_dump()["service_tier"] == "flex" + assert UserAPIKeyLabelValues().service_tier is None + + +def test_served_tier_wins_over_requested_tier(): + payload = _standard_logging_payload( + response={"service_tier": "default"}, + model_parameters={"service_tier": "auto"}, + ) + + assert get_service_tier_from_standard_logging_payload(payload) == "default" + + +def test_usage_object_tier_used_when_response_has_none(): + payload = _standard_logging_payload( + response={"id": "chatcmpl-1"}, + usage_object={"prompt_tokens": 1, "service_tier": "standard"}, + model_parameters={"service_tier": "auto"}, + ) + + assert get_service_tier_from_standard_logging_payload(payload) == "standard" + + +def test_requested_tier_used_when_no_served_tier(): + payload = _standard_logging_payload( + response={"id": "chatcmpl-1"}, + usage_object={"prompt_tokens": 1}, + model_parameters={"service_tier": "flex"}, + ) + + assert get_service_tier_from_standard_logging_payload(payload) == "flex" + + +def test_unrecognized_requested_tier_is_not_labelled(): + """ + A caller-supplied tier survives param mapping even where the provider then + ignores it (Bedrock and Groq drop an unrecognized tier and still answer), so + labelling it verbatim would let one caller mint a series per string. + """ + payload = _standard_logging_payload(model_parameters={"service_tier": "attacker-controlled-a1b2c3"}) + + assert get_service_tier_from_standard_logging_payload(payload) is None + + +def test_unrecognized_served_tier_is_labelled(): + """ + The tier a provider reports is not caller-controlled, so a tier added by a + provider after this release still gets labelled instead of being dropped. + """ + payload = _standard_logging_payload( + response={"service_tier": "tier-added-by-provider-later"}, + model_parameters={"service_tier": "auto"}, + ) + + assert get_service_tier_from_standard_logging_payload(payload) == "tier-added-by-provider-later" + + +@pytest.mark.parametrize("tier", sorted(KNOWN_REQUEST_SERVICE_TIERS)) +def test_every_known_requested_tier_is_labelled(tier): + payload = _standard_logging_payload(model_parameters={"service_tier": tier}) + + assert get_service_tier_from_standard_logging_payload(payload) == tier + + +@pytest.mark.parametrize( + "response, usage_object, model_parameters", + [ + (None, None, None), + ({"service_tier": None}, {"service_tier": ""}, {"service_tier": None}), + ("redacted-by-litellm", None, {}), + ({"service_tier": 1}, None, None), + ], +) +def test_no_tier_resolves_to_none(response, usage_object, model_parameters): + payload = _standard_logging_payload( + response=response, + usage_object=usage_object, + model_parameters=model_parameters, + ) + + assert get_service_tier_from_standard_logging_payload(payload) is None + + +@pytest.mark.asyncio +async def test_success_event_emits_service_tier_on_latency_and_spend_metrics(): + """ + End-to-end emit wiring. + + Drives the real logger with a payload whose response was served on the flex + tier and asserts every latency histogram and the spend counter carries + service_tier="flex". Fails if the label is dropped from a metric's label list + or if the value is not populated on the success path. + """ + payload = _standard_logging_payload( + response={"service_tier": "flex"}, + model_parameters={"service_tier": "auto"}, + ) + now = datetime.datetime.now() + kwargs = { + "model": "gpt-4o-mini", + "litellm_params": {"metadata": {}}, + "standard_logging_object": payload, + "stream": True, + "start_time": now - datetime.timedelta(seconds=3), + "api_call_start_time": now - datetime.timedelta(seconds=2), + "completion_start_time": now - datetime.timedelta(seconds=1), + "end_time": now, + } + + _clear_prometheus_registry() + try: + logger = PrometheusLogger() + await logger.async_log_success_event(kwargs, None, now, now) + + for metric_name in ( + "litellm_request_total_latency_metric_bucket", + "litellm_llm_api_latency_metric_bucket", + "litellm_llm_api_time_to_first_token_metric_bucket", + "litellm_spend_metric_total", + ): + samples = _collected_samples(metric_name) + assert samples, f"expected {metric_name} to be emitted" + assert all(sample.labels.get("service_tier") == "flex" for sample in samples), ( + f"{metric_name} must carry service_tier=flex, got " + f"{sorted({sample.labels.get('service_tier') for sample in samples})}" + ) + finally: + _clear_prometheus_registry() diff --git a/tests/test_litellm/interactions/test_openapi_compliance.py b/tests/test_litellm/interactions/test_openapi_compliance.py index 11b08fa45a8..1fe343ca6ee 100644 --- a/tests/test_litellm/interactions/test_openapi_compliance.py +++ b/tests/test_litellm/interactions/test_openapi_compliance.py @@ -37,6 +37,13 @@ def _load_openapi_spec_dict() -> Dict[str, Any]: ) +def _declared_type_value(variant_schema: Dict[str, Any]) -> Any: + """The single `type` value a union variant pins, whether spelled as a const or a 1-item enum.""" + type_property = variant_schema.get("properties", {}).get("type", {}) + enum_values = type_property.get("enum") or [] + return type_property.get("const") or (enum_values[0] if len(enum_values) == 1 else None) + + @pytest.fixture(scope="module") def spec_dict() -> Dict[str, Any]: """Load raw spec dict for manual validation.""" @@ -105,26 +112,51 @@ class TestRequestCompliance: assert "string" in input_types, "Input should support string" assert "array" in input_types, "Input should support array" - def test_content_schema_uses_discriminator(self, spec_dict): - """Verify Content uses type discriminator.""" + def test_content_variants_are_identified_by_their_type_field(self, spec_dict): + """Verify a Content part can be told apart by its `type`, however the spec spells that. + + Our transformation reads `type` off each content part to route it, so what has to hold is + that every variant of the union pins a distinct `type` value and that text is one of them. + A spec may express that with an OpenAPI `discriminator` on the union or with a `const` on + each member's own `type`; both are equivalent for us, so accepting only the first makes + this test fail on a stylistic change upstream that costs us nothing. + """ content_schema = spec_dict["components"]["schemas"]["Content"] - assert "discriminator" in content_schema - assert content_schema["discriminator"]["propertyName"] == "type" - - # Check TextContent is an option (via mapping if present, or via oneOf refs) - mapping = content_schema["discriminator"].get("mapping") - if mapping: - assert "text" in mapping - print(f"Content type discriminator mapping: {list(mapping.keys())}") - else: - # Discriminator without explicit mapping — verify via oneOf - one_of = content_schema.get("oneOf", []) - ref_names = [opt["$ref"].split("/")[-1] for opt in one_of if "$ref" in opt] + discriminator = content_schema.get("discriminator") + if discriminator is not None: assert ( - "TextContent" in ref_names - ), f"TextContent not found in oneOf refs: {ref_names}" - print(f"Content type discriminator (no mapping), oneOf refs: {ref_names}") + discriminator.get("propertyName") == "type" + ), f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'" + + variant_names = [ + option["$ref"].split("/")[-1] + for option in content_schema.get("oneOf", []) + if "$ref" in option + ] + assert variant_names, f"Content is not a union of named variants: {content_schema}" + + mapping = (discriminator or {}).get("mapping") or {} + type_values = { + variant: mapping_value + for mapping_value, ref in mapping.items() + for variant in [ref.split("/")[-1]] + } or { + variant: _declared_type_value(spec_dict["components"]["schemas"].get(variant, {})) + for variant in variant_names + } + + assert set(type_values) == set(variant_names) and all(type_values.values()), ( + f"every Content variant needs a discoverable type value, " + f"got {type_values} for variants {sorted(variant_names)}" + ) + assert len(set(type_values.values())) == len(type_values), ( + f"Content variants must pin DISTINCT type values, got {type_values}" + ) + assert type_values.get("TextContent") == "text", ( + f"TextContent must be reachable as type 'text', got {type_values}" + ) + print(f"Content variants by type: {type_values}") def test_text_content_schema(self, spec_dict): """Verify TextContent schema.""" diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index dfe7e0c3a51..c0c6e315b5b 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -1582,6 +1582,67 @@ def test_thinking_still_translated_to_reasoning_effort_for_non_claude_model(): assert new_kwargs["reasoning_effort"] == "low" +def test_thinking_disabled_translated_to_reasoning_effort_none_for_non_claude_model(): + adapter = LiteLLMAnthropicMessagesAdapter() + thinking = {"type": "disabled"} + + new_kwargs = {"model": CACHE_CONTROL_NON_ANTHROPIC_MODEL} + adapter._translate_thinking_to_openai(cast(Any, {"thinking": thinking}), cast(Any, new_kwargs)) + + assert "thinking" not in new_kwargs + assert new_kwargs["reasoning_effort"] == "none" + + +def test_thinking_disabled_stays_plain_string_when_auto_summary_enabled(): + import litellm + + adapter = LiteLLMAnthropicMessagesAdapter() + thinking = {"type": "disabled"} + + original = litellm.reasoning_auto_summary + try: + litellm.reasoning_auto_summary = True + new_kwargs = {"model": CACHE_CONTROL_NON_ANTHROPIC_MODEL} + adapter._translate_thinking_to_openai(cast(Any, {"thinking": thinking}), cast(Any, new_kwargs)) + finally: + litellm.reasoning_auto_summary = original + + assert new_kwargs["reasoning_effort"] == "none" + + +def test_stop_sequences_translated_to_stop_for_non_claude_model(): + from litellm.types.llms.anthropic import AnthropicMessagesRequest + + anthropic_request = AnthropicMessagesRequest( + model=CACHE_CONTROL_NON_ANTHROPIC_MODEL, + max_tokens=1024, + messages=[{"role": "user", "content": "hi"}], + stop_sequences=[""], + ) + + adapter = LiteLLMAnthropicMessagesAdapter() + openai_request, _ = adapter.translate_anthropic_to_openai(anthropic_message_request=anthropic_request) + + assert openai_request["stop"] == [""] + assert "stop_sequences" not in openai_request + + +def test_empty_stop_sequences_does_not_set_stop(): + from litellm.types.llms.anthropic import AnthropicMessagesRequest + + anthropic_request = AnthropicMessagesRequest( + model=CACHE_CONTROL_NON_ANTHROPIC_MODEL, + max_tokens=1024, + messages=[{"role": "user", "content": "hi"}], + stop_sequences=[], + ) + + adapter = LiteLLMAnthropicMessagesAdapter() + openai_request, _ = adapter.translate_anthropic_to_openai(anthropic_message_request=anthropic_request) + + assert "stop" not in openai_request + + def test_cache_control_preserved_in_image_content_for_claude(): """Cache control should be preserved in image content for Claude models.""" anthropic_messages = [ diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py index 19ec1a04b45..17a57d974de 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py @@ -506,3 +506,162 @@ def test_empty_content_chunk_mid_text_block_is_suppressed_sync(): assert _text_deltas(events) == ["Hi", " there"] _assert_deltas_match_their_block_type(events) + + +def _thinking_first_chunks() -> List[MagicMock]: + return [ + _thinking_chunk("Let me think"), + _thinking_chunk("about it."), + _make_chunk(Delta(content="42")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + + +def _assert_thinking_first_block_opens_at_index_zero(events: List[dict]) -> None: + starts = [ + (e["index"], e["content_block"]["type"]) + for e in events + if e.get("type") == "content_block_start" + ] + assert starts == [(0, "thinking"), (1, "text")], starts + assert "" not in _text_deltas(events) + assert _thinking_deltas(events) == ["Let me think", "about it."] + assert _text_deltas(events) == ["42"] + _assert_deltas_match_their_block_type(events) + + +def test_thinking_first_stream_opens_thinking_block_at_index_zero_sync(): + """Bug A regression: when the model's first output is reasoning the adapter + must open the first content block as ``thinking`` at index 0. The previous + code pre-emitted a hardcoded empty ``text`` block at index 0 before + inspecting any upstream chunk, then opened ``thinking`` at index 1; strict + Anthropic SDK clients with thinking enabled reject that stream with + "Content block is not a thinking block". + """ + wrapper = AnthropicStreamWrapper( + completion_stream=iter(_thinking_first_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper)) + + +@pytest.mark.asyncio +async def test_thinking_first_stream_opens_thinking_block_at_index_zero_async(): + """Async twin of the Bug A regression; the proxy serves the async iterator, + so the first block must be ``thinking`` at index 0 on this path too. + """ + wrapper = AnthropicStreamWrapper( + completion_stream=_AsyncStream(_thinking_first_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper)) + + +def _reasoning_content_chunk(reasoning: str) -> MagicMock: + return _make_chunk(Delta(content=None, reasoning_content=reasoning)) + + +def _reasoning_first_chunks() -> List[MagicMock]: + return [ + _reasoning_content_chunk("Let me think"), + _reasoning_content_chunk("about it."), + _make_chunk(Delta(content="42")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + + +def test_reasoning_content_first_stream_opens_thinking_block_at_index_zero_sync(): + """The reported backend (hosted_vllm; vLLM and SGLang reasoning parsers) + surfaces reasoning as OpenAI ``reasoning_content`` with no + ``thinking_blocks``. Such a stream must also open the first content block as + ``thinking`` at index 0, exercising the reasoning_content branch of the + chunk translator rather than the thinking_blocks branch the other twins use. + """ + wrapper = AnthropicStreamWrapper( + completion_stream=iter(_reasoning_first_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper)) + + +def _blank_lead_chunks() -> List[MagicMock]: + return [ + _make_chunk(Delta(content=None)), + _thinking_chunk("Let me think"), + _thinking_chunk("about it."), + _make_chunk(Delta(content="42")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + + +def _role_only_reasoning_content_lead_chunks() -> List[MagicMock]: + return [ + _make_chunk(Delta(role="assistant", content=None, tool_calls=[])), + _reasoning_content_chunk("Let me think"), + _reasoning_content_chunk("about it."), + _make_chunk(Delta(content="42")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + + +def test_contentless_lead_chunk_does_not_open_text_block_before_thinking_sync(): + """OpenAI-compatible streaming backends open the response with a contentless + priming chunk (an empty delta, e.g. the {role: assistant} lead-in) before the + first real token. Such a lead chunk must NOT commit index 0 to an empty text + block; the following thinking chunk must still open thinking at index 0, or + strict Anthropic SDK clients reject the stream. + """ + wrapper = AnthropicStreamWrapper( + completion_stream=iter(_blank_lead_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper)) + + +def test_role_only_lead_chunk_does_not_open_text_block_before_reasoning_content_sync(): + wrapper = AnthropicStreamWrapper( + completion_stream=iter(_role_only_reasoning_content_lead_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper)) + + +@pytest.mark.asyncio +async def test_contentless_lead_chunk_does_not_open_text_block_before_thinking_async(): + """Async twin; the proxy serves the async iterator, so the contentless lead + chunk must be skipped on this path too. + """ + wrapper = AnthropicStreamWrapper( + completion_stream=_AsyncStream(_blank_lead_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper)) + + +@pytest.mark.asyncio +async def test_role_only_lead_chunk_does_not_open_text_block_before_reasoning_content_async(): + wrapper = AnthropicStreamWrapper( + completion_stream=_AsyncStream(_role_only_reasoning_content_lead_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper)) + + +def test_finish_first_chunk_is_not_deferred_sync(): + """A stream whose first upstream chunk is already the finish event must not + be skipped by the blank-delta deferral. ``_is_blank_delta`` returns False + for a finish chunk so the message_delta still flows (with an empty text + block opened and closed first); without that guard the deferral would drop + the terminal event entirely. + """ + chunks = [_make_chunk(Delta(content=None), finish_reason="stop")] + wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x") + events = _drain_sync(wrapper) + + assert [e["type"] for e in events] == [ + "message_start", + "content_block_start", + "content_block_stop", + "message_delta", + "message_stop", + ] diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py index 8875a75e86f..df3db3d2c57 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py @@ -651,6 +651,26 @@ class TestThinkingSummaryPreservation: "reasoning_effort": {"effort": "high", "summary": "concise"} } + def test_translate_thinking_for_model_disabled_stays_plain_string_when_auto_summary_enabled(self): + """Disabled thinking must stay a plain string even when reasoning_auto_summary is on.""" + import litellm + from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import ( + LiteLLMAnthropicMessagesAdapter, + ) + + original = litellm.reasoning_auto_summary + try: + litellm.reasoning_auto_summary = True + thinking = {"type": "disabled"} + result = LiteLLMAnthropicMessagesAdapter.translate_thinking_for_model( + thinking=thinking, + model="openai/gpt-5.2", + ) + finally: + litellm.reasoning_auto_summary = original + + assert result == {"reasoning_effort": "none"} + # --------------------------------------------------------------------------- # Parity tests: redundant empty-text-block sanitization scan removal. diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py index e44413cf837..6d7cd2f88be 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py @@ -134,13 +134,6 @@ def test_anthropic_stream_wrapper_single_tool_call(): # Verify the expected sequence of chunk types expected_types = [ "message_start", # Initial message start - # TODO: for future contributors: if the initial content_block_start - # respects the upstream's starting chunk, the initial empty text block - # should be removed (and this test should be updated accordingly) - # --------------------------------------------------------------------- - "content_block_start", # Initial empty text block start - "content_block_stop", # End of empty text block - # --------------------------------------------------------------------- "content_block_start", # Start of first tool_use content block "content_block_delta", # {"city": "content_block_delta", # "NY"} @@ -196,13 +189,6 @@ def test_anthropic_stream_wrapper_back_to_back_tool_calls(): # Verify the expected sequence of chunk types expected_types = [ "message_start", # Initial message start - # TODO: for future contributors: if the initial content_block_start - # respects the upstream's starting chunk, the initial empty text block - # should be removed (and this test should be updated accordingly) - # --------------------------------------------------------------------- - "content_block_start", # Initial empty text block start - "content_block_stop", # End of empty text block - # --------------------------------------------------------------------- "content_block_start", # Start of first tool_use content block "content_block_delta", # {"city": "content_block_delta", # "NY"} @@ -267,13 +253,6 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text(): # Verify the expected sequence of chunk types expected_types = [ "message_start", # Initial message start - # TODO: for future contributors: if the initial content_block_start - # respects the upstream's starting chunk, the initial empty text block - # should be removed (and this test should be updated accordingly) - # --------------------------------------------------------------------- - "content_block_start", # Initial empty text block start - "content_block_stop", # End of empty text block - # --------------------------------------------------------------------- "content_block_start", # Start of first tool_use content block "content_block_delta", # {"city": "content_block_delta", # "NY"} diff --git a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py index 61987d25d9c..ee50b9db015 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py @@ -8,14 +8,11 @@ sys.path.insert( 0, os.path.abspath("../../../../..") ) # Adds the parent directory to the system path -import litellm from litellm.llms.bedrock.chat.invoke_handler import ( AWSEventStreamDecoder, - BedrockLLM, make_call, make_sync_call, ) -from litellm.llms.custom_httpx.http_handler import HTTPHandler def test_transform_thinking_blocks_with_redacted_content(): @@ -296,33 +293,3 @@ def test_make_sync_call_honors_explicit_stream_chunk_size(): response.iter_bytes.assert_called_once_with(chunk_size=2048) - -def test_legacy_bedrock_llm_streaming_does_not_rechunk_by_default(): - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.iter_bytes = MagicMock(return_value=iter([])) - client = HTTPHandler() - client.post = MagicMock(return_value=mock_response) - - BedrockLLM().completion( - model="cohere.command-text-v14", - messages=[{"role": "user", "content": "hi"}], - api_base=None, - custom_prompt_dict={}, - model_response=litellm.ModelResponse(), - print_verbose=lambda *args, **kwargs: None, - encoding=litellm.encoding, - logging_obj=MagicMock(), - optional_params={ - "stream": True, - "aws_access_key_id": "fake", - "aws_secret_access_key": "fake", - "aws_region_name": "us-east-1", - }, - acompletion=False, - timeout=None, - litellm_params={}, - client=client, - ) - - mock_response.iter_bytes.assert_called_once_with(chunk_size=None) diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py b/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py index 5515e6ce815..5a37e681c4a 100644 --- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py +++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py @@ -1,6 +1,9 @@ import socket from unittest.mock import MagicMock, patch +import aiohttp +import pytest + def _invoke_connector_factory(http_handler_module): """ @@ -159,3 +162,32 @@ def test_socket_factory_uses_tcp_keepalive_when_keepidle_unavailable(monkeypatch setsockopt_calls[(socket.IPPROTO_TCP, fake_socket_module.TCP_KEEPALIVE)] == 60 ) assert (socket.IPPROTO_TCP, getattr(socket, "TCP_KEEPIDLE", -1)) not in setsockopt_calls + + +@pytest.mark.asyncio +async def test_shared_session_transport_rebuilds_with_socket_factory(monkeypatch): + """ + The proxy hands _create_aiohttp_transport an already-built shared session. + When that session is rebuilt (closed session, or a session from another + event loop) the replacement must still carry the keep-alive socket factory + and the configured keepalive timeout, otherwise AIOHTTP_SO_KEEPALIVE stops + protecting every later request served by that transport. + """ + from litellm.llms.custom_httpx import http_handler as http_handler_module + + monkeypatch.setattr(http_handler_module, "AIOHTTP_SO_KEEPALIVE", True) + monkeypatch.setattr(http_handler_module, "_AIOHTTP_SUPPORTS_SOCKET_FACTORY", True) + + shared_session = aiohttp.ClientSession() + transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport(shared_session=shared_session) + await shared_session.close() + + rebuilt_session = MagicMock(name="rebuilt_session") + + with patch.object(http_handler_module, "TCPConnector", return_value=MagicMock(name="connector")) as mock_tcp_connector: + with patch.object(http_handler_module, "ClientSession", return_value=rebuilt_session): + assert transport._get_valid_client_session() is rebuilt_session + + assert mock_tcp_connector.call_count == 1 + assert callable(mock_tcp_connector.call_args.kwargs.get("socket_factory")) + assert mock_tcp_connector.call_args.kwargs["keepalive_timeout"] == http_handler_module.AIOHTTP_KEEPALIVE_TIMEOUT diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py index d1bc356662f..0550898c73d 100644 --- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py +++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py @@ -727,3 +727,103 @@ async def test_response_stream_closes_response_on_generator_exit(): await iterator.aclose() assert mock_response.closed is True + + +@pytest.mark.asyncio +async def test_closed_shared_session_rebuild_uses_injected_session_factory(): + """ + A transport handed an already-built session (the proxy's shared session) + must rebuild through the injected factory. Rebuilding with a bare + ClientSession drops the connector's keep-alive socket options, pool limits + and DNS cache for every later request on that transport. + """ + shared_session = aiohttp.ClientSession() + await shared_session.close() + + rebuilt = [] + + def session_factory(): + session = _make_mock_session() + rebuilt.append(session) + return session + + transport = LiteLLMAiohttpTransport( + client=shared_session, + owns_session=False, + session_factory=session_factory, # type: ignore + ) + + assert transport._get_valid_client_session() in rebuilt + + +def test_rebuild_without_running_loop_uses_injected_session_factory(): + """ + The loop-validity fallback must also go through the injected factory, so a + transport recovering outside a running event loop does not silently swap in + an unconfigured session. + """ + rebuilt = [] + + def session_factory(): + session = _make_mock_session() + rebuilt.append(session) + return session + + transport = LiteLLMAiohttpTransport( + client=object(), # type: ignore + session_factory=session_factory, # type: ignore + ) + + assert transport._get_valid_client_session() in rebuilt + + +@pytest.mark.asyncio +async def test_rebuilt_session_becomes_transport_owned(): + """ + A rebuilt session is reachable only from the transport, so aclose() must + close it even when the transport was handed a session it did not own. + """ + shared_session = aiohttp.ClientSession() + await shared_session.close() + + replacement = aiohttp.ClientSession() + transport = LiteLLMAiohttpTransport( + client=shared_session, + owns_session=False, + session_factory=lambda: replacement, + ) + + assert transport._get_valid_client_session() is replacement + + await transport.aclose() + + assert replacement.closed + + +@pytest.mark.asyncio +async def test_stale_loop_rebuild_does_not_close_unowned_session(): + """ + A session the transport does not own (the proxy's shared session) is used by + other transports too, so a rebuild must leave it open for them. + """ + shared_session = aiohttp.ClientSession() + running_loop = asyncio.get_running_loop() + other_loop = asyncio.new_event_loop() + + replacement = _make_mock_session() + transport = LiteLLMAiohttpTransport( + client=shared_session, + owns_session=False, + session_factory=lambda: replacement, # type: ignore + ) + + try: + shared_session._loop = other_loop + assert transport._get_valid_client_session() is replacement + shared_session._loop = running_loop + await asyncio.sleep(0.05) + assert not shared_session.closed + finally: + shared_session._loop = running_loop + other_loop.close() + await shared_session.close() diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index cf75964ddb7..ad890d0c7ea 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1452,6 +1452,157 @@ class TestContextCachingEndpoints: # Restart the patcher so teardown_method can stop it cleanly self._token_check_patcher.start() + def _model_turn_final_messages(self, final_cached_role): + tool_call = { + "id": "call_abc123", + "type": "function", + "function": {"name": "get_weather", "arguments": '{"location": "Boston"}'}, + } + cached_tail = { + "assistant": [], + "tool": [ + { + "role": "tool", + "tool_call_id": "call_abc123", + "content": "72F and sunny", + "cache_control": {"type": "ephemeral"}, + } + ], + "system": [ + { + "role": "system", + "content": "Tool results are authoritative.", + "cache_control": {"type": "ephemeral"}, + } + ], + }[final_cached_role] + return [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "Use the weather tool for every answer.", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + { + "role": "assistant", + "content": "", + "tool_calls": [tool_call], + "cache_control": {"type": "ephemeral"}, + }, + *cached_tail, + {"role": "user", "content": "What is the weather in Boston?"}, + ] + + @pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"]) + def test_check_and_create_cache_skips_when_cached_block_ends_on_model_turn( + self, final_cached_role + ): + """The cachedContents API rejects contents ending on an assistant or tool turn + with HTTP 400 "Requests ending with a model turn are not supported", so the + request must proceed uncached instead of failing. + """ + all_messages = self._model_turn_final_messages(final_cached_role) + optional_params = self.sample_optional_params.copy() + + result = self.context_caching.check_and_create_cache( + messages=all_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-3.6-flash", + client=self.mock_client, + timeout=30.0, + logging_obj=self.mock_logging, + cached_content=None, + custom_llm_provider="vertex_ai", + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + + messages, returned_params, returned_cache = result + assert messages == all_messages + assert returned_cache is None + assert "tools" in returned_params + self.mock_client.get.assert_not_called() + self.mock_client.post.assert_not_called() + + @pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"]) + @pytest.mark.asyncio + async def test_async_check_and_create_cache_skips_when_cached_block_ends_on_model_turn( + self, final_cached_role + ): + """Async variant: an unsupported terminal turn skips caching instead of failing.""" + all_messages = self._model_turn_final_messages(final_cached_role) + optional_params = self.sample_optional_params.copy() + + result = await self.context_caching.async_check_and_create_cache( + messages=all_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-3.6-flash", + client=self.mock_async_client, + timeout=30.0, + logging_obj=self.mock_logging, + cached_content=None, + custom_llm_provider="vertex_ai", + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + + messages, returned_params, returned_cache = result + assert messages == all_messages + assert returned_cache is None + assert "tools" in returned_params + self.mock_async_client.get.assert_not_called() + self.mock_async_client.post.assert_not_called() + + +def test_cached_messages_end_on_supported_turn(): + from litellm.llms.vertex_ai.context_caching.transformation import ( + cached_messages_end_on_supported_turn, + ) + + assert ( + cached_messages_end_on_supported_turn( + [{"role": "assistant", "content": "hi"}, {"role": "user", "content": "hello"}] + ) + is True + ) + assert cached_messages_end_on_supported_turn([{"role": "system", "content": "be brief"}]) is True + assert cached_messages_end_on_supported_turn([{"role": "assistant", "content": "hi"}]) is False + assert ( + cached_messages_end_on_supported_turn( + [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "hi"}, + {"role": "system", "content": "be brief"}, + ] + ) + is False + ) + assert ( + cached_messages_end_on_supported_turn( + [{"role": "system", "content": "be brief"}, {"role": "user", "content": "hello"}] + ) + is True + ) + assert ( + cached_messages_end_on_supported_turn([{"role": "tool", "tool_call_id": "x", "content": "y"}]) + is False + ) + assert ( + cached_messages_end_on_supported_turn([{"role": "function", "name": "f", "content": "y"}]) + is False + ) + assert cached_messages_end_on_supported_turn([]) is False + class TestCheckCachePagination: """Test pagination logic in check_cache and async_check_cache methods.""" diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py index 453a0c14bf9..5e854bbad70 100644 --- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py @@ -31,10 +31,7 @@ class TestVertexAIFilesHandler: def test_extract_bucket_and_object_from_file_id_standard_path(self): """Test extraction of bucket and object from URL-encoded file_id with standard path""" # Sample file_id with nested folder structure - file_id = ( - "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files" - "%2Ftest-folder%2Fsub-folder%2Ftest-file.txt" - ) + file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Ftest-folder%2Fsub-folder%2Ftest-file.txt" bucket_name, object_path = self.handler._extract_bucket_and_object_from_file_id( file_id=file_id, @@ -105,21 +102,14 @@ class TestVertexAIFilesHandler: async def test_afile_content_success(self): """Test successful async file content retrieval""" # Setup test data - file_id = ( - "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files" - "%2Fuploads%2Fabc-test-file.txt" - ) + file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt" expected_content = b"test file content" - file_content_request = FileContentRequest( - file_id=file_id, extra_headers=None, extra_body=None - ) + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) # Mock the download_gcs_object method with ( - patch.object( - self.handler, "download_gcs_object", new_callable=AsyncMock - ) as mock_download, + patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download, patch.object( self.handler, "get_gcs_logging_config", @@ -148,15 +138,9 @@ class TestVertexAIFilesHandler: # Verify the download was called with correct parameters mock_download.assert_called_once() call_args = mock_download.call_args - assert ( - call_args.kwargs["object_name"] - == "litellm-vertex-files/uploads/abc-test-file.txt" - ) + assert call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-test-file.txt" assert "standard_callback_dynamic_params" in call_args.kwargs - assert ( - call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"] - == "test-bucket" - ) + assert call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"] == "test-bucket" @pytest.mark.asyncio async def test_afile_content_missing_file_id(self): @@ -164,9 +148,7 @@ class TestVertexAIFilesHandler: file_content_request = FileContentRequest(extra_headers=None, extra_body=None) # Should raise ValueError for missing file_id - with pytest.raises( - ValueError, match="file_id is required in file_content_request" - ): + with pytest.raises(ValueError, match="file_id is required in file_content_request"): await self.handler.afile_content( file_content_request=file_content_request, vertex_credentials=None, @@ -179,20 +161,13 @@ class TestVertexAIFilesHandler: @pytest.mark.asyncio async def test_afile_content_download_failure(self): """Test async file content retrieval when download fails""" - file_id = ( - "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files" - "%2Fuploads%2Fabc-test-file.txt" - ) + file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt" - file_content_request = FileContentRequest( - file_id=file_id, extra_headers=None, extra_body=None - ) + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) # Mock download to return None (failure) with ( - patch.object( - self.handler, "download_gcs_object", new_callable=AsyncMock - ) as mock_download, + patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download, patch.object( self.handler, "get_gcs_logging_config", @@ -216,14 +191,130 @@ class TestVertexAIFilesHandler: max_retries=3, ) + def test_resolve_read_gcs_config_prefers_per_model_bucket(self, monkeypatch): + monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket") + monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json") + + bucket, service_account = self.handler._resolve_read_gcs_config( + litellm_params={ + "gcs_bucket_name": "my-model-bucket", + "vertex_credentials": "/model/sa.json", + }, + vertex_credentials=None, + ) + + assert bucket == "my-model-bucket" + assert service_account == "/model/sa.json" + + def test_resolve_read_gcs_config_falls_back_to_env(self, monkeypatch): + monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket") + monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json") + + bucket, service_account = self.handler._resolve_read_gcs_config(litellm_params={}, vertex_credentials=None) + + assert bucket == "env-default-bucket" + assert service_account == "/env/sa.json" + + def test_resolve_read_gcs_config_serializes_dict_credentials(self, monkeypatch): + monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) + + _, service_account = self.handler._resolve_read_gcs_config( + litellm_params={"gcs_bucket_name": "my-model-bucket"}, + vertex_credentials={"type": "service_account", "project_id": "p"}, + ) + + assert service_account == '{"type": "service_account", "project_id": "p"}' + + @pytest.mark.asyncio + async def test_afile_content_honors_per_model_bucket_over_env(self, monkeypatch): + """ + Regression for #32640: a batch output written to a per-model gcs_bucket_name must be + readable even when the global GCS_BUCKET_NAME points at a different bucket. Before the + fix the read path resolved the bucket from env only and raised + "file_id bucket does not match the configured storage bucket". + """ + monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket") + monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) + + file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl" + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) + + with ( + patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download, + patch.object( + self.handler, + "get_or_create_vertex_instance", + new_callable=AsyncMock, + return_value=object(), + ), + ): + mock_download.return_value = b"batch output" + + result = await self.handler.afile_content( + file_content_request=file_content_request, + vertex_credentials="/model/sa.json", + vertex_project="test-project", + vertex_location="us-central1", + timeout=60.0, + max_retries=0, + litellm_params={ + "gcs_bucket_name": "my-model-bucket", + "vertex_credentials": "/model/sa.json", + }, + ) + + assert isinstance(result, HttpxBinaryResponseContent) + assert result.response.content == b"batch output" + + dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"] + assert dynamic_params["gcs_bucket_name"] == "my-model-bucket" + assert dynamic_params["gcs_path_service_account"] == "/model/sa.json" + assert mock_download.call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-batch-output.jsonl" + + @pytest.mark.asyncio + async def test_afile_content_reads_without_global_env_bucket(self, monkeypatch): + """ + Regression for #32640: with no global GCS_BUCKET_NAME set, a model-group-level + deployment (per-model gcs_bucket_name) must still be readable. Before the fix the read + path raised "GCS_BUCKET_NAME is not set in the environment". + """ + monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) + monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) + + file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl" + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) + + with ( + patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download, + patch.object( + self.handler, + "get_or_create_vertex_instance", + new_callable=AsyncMock, + return_value=object(), + ), + ): + mock_download.return_value = b"batch output" + + result = await self.handler.afile_content( + file_content_request=file_content_request, + vertex_credentials="/model/sa.json", + vertex_project="test-project", + vertex_location="us-central1", + timeout=60.0, + max_retries=0, + litellm_params={"gcs_bucket_name": "my-model-bucket"}, + ) + + assert isinstance(result, HttpxBinaryResponseContent) + dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"] + assert dynamic_params["gcs_bucket_name"] == "my-model-bucket" + def test_file_content_sync_success(self): """Test successful sync file content retrieval""" file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt" expected_content = b"test file content" - file_content_request = FileContentRequest( - file_id=file_id, extra_headers=None, extra_body=None - ) + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) # Create expected response mock_response = httpx.Response( @@ -261,25 +352,17 @@ class TestVertexAIFilesHandler: file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt" expected_content = b"test file content" - file_content_request = FileContentRequest( - file_id=file_id, extra_headers=None, extra_body=None - ) + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) # Mock the afile_content method - with patch.object( - self.handler, "afile_content", new_callable=AsyncMock - ) as mock_afile_content: + with patch.object(self.handler, "afile_content", new_callable=AsyncMock) as mock_afile_content: mock_response = httpx.Response( status_code=200, content=expected_content, headers={"content-type": "application/octet-stream"}, - request=httpx.Request( - method="GET", url="gs://test-bucket/test-file.txt" - ), - ) - mock_afile_content.return_value = HttpxBinaryResponseContent( - response=mock_response + request=httpx.Request(method="GET", url="gs://test-bucket/test-file.txt"), ) + mock_afile_content.return_value = HttpxBinaryResponseContent(response=mock_response) # Call the method with _is_async=True result = self.handler.file_content( diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 5871644ca5f..51cc2857252 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -2276,82 +2276,8 @@ def test_is_gemini_3_or_newer(): assert VertexGeminiConfig._is_gemini_3_or_newer("") == False -def test_forward_gemini_function_call_id_vertex_vs_google_ai_studio(): - """Vertex AI rejects `id` on function_call/function_response; Google AI Studio accepts it on Gemini 3.5+.""" - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - - model = "gemini-3.5-flash" - assert ( - VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai") is False - ) - assert ( - VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai_beta") - is False - ) - assert VertexGeminiConfig._forward_gemini_function_call_id(model, "gemini") is True - assert VertexGeminiConfig._forward_gemini_function_call_id(model, None) is False - assert ( - VertexGeminiConfig._forward_gemini_function_call_id( - "gemini-2.5-flash", "gemini" - ) - is False - ) - - -def test_vertex_ai_gemini_35_tool_calls_omit_function_call_id(): - """Regression: Vertex must not send OpenAI tool_call id inside Gemini function_call parts.""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - messages = [ - {"role": "user", "content": "Explore this directory"}, - { - "role": "assistant", - "content": "", - "tool_calls": [ - { - "id": "call_50e7e0fe0989464a89f188eda443", - "type": "function", - "function": { - "name": "read", - "arguments": '{"filePath": "/tmp"}', - }, - } - ], - }, - { - "role": "tool", - "tool_call_id": "call_50e7e0fe0989464a89f188eda443", - "content": "ok", - }, - ] - - contents = _gemini_convert_messages_with_history( - messages=messages, - model="gemini-3.5-flash", - custom_llm_provider="vertex_ai", - ) - - for content in contents: - for part in content.get("parts", []): - fc = part.get("function_call") - if fc is not None: - assert "id" not in fc, f"Vertex payload must not include id: {fc}" - fr = part.get("function_response") - if fr is not None: - assert "id" not in fr, f"Vertex payload must not include id: {fr}" - - -def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id(): - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - tool_call_id = "call_50e7e0fe0989464a89f188eda443" - messages = [ +def _tool_call_messages(tool_call_id: str): + return [ {"role": "user", "content": "hi"}, { "role": "assistant", @@ -2374,12 +2300,8 @@ def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id(): }, ] - contents = _gemini_convert_messages_with_history( - messages=messages, - model="gemini-3.5-flash", - custom_llm_provider="gemini", - ) +def _collect_function_call_ids(contents): function_call_ids = [] function_response_ids = [] for content in contents: @@ -2390,9 +2312,120 @@ def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id(): fr = part.get("function_response") if fr is not None: function_response_ids.append(fr.get("id")) + return function_call_ids, function_response_ids - assert function_call_ids == [tool_call_id] - assert function_response_ids == [tool_call_id] + +def test_forward_gemini_function_call_id_is_gated_on_model_version_only(): + """Gemini 3+ takes `id` on Vertex AI and Google AI Studio alike; older models reject it.""" + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + + assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3.5-flash") is True + assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3-pro") is True + assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.5-flash") is False + assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.0-flash") is False + + +@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "vertex_ai_beta", "gemini"]) +def test_gemini_35_tool_calls_include_function_call_id(custom_llm_provider): + """Vertex AI accepts `id` on Gemini 3+, so it must be sent there and not just on AI Studio. + + Both parts are asserted together: Vertex pairs a result to its call by id, so emitting one + side without the other would break strict tool-call matching. + """ + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + tool_call_id = "call_50e7e0fe0989464a89f188eda443" + contents = _gemini_convert_messages_with_history( + messages=_tool_call_messages(tool_call_id), + model="gemini-3.5-flash", + custom_llm_provider=custom_llm_provider, + ) + + assert _collect_function_call_ids(contents) == ([tool_call_id], [tool_call_id]) + + +@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "gemini"]) +def test_gemini_25_tool_calls_omit_function_call_id(custom_llm_provider): + """Regression: models older than Gemini 3 reject `id`, so the key must be absent entirely.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + contents = _gemini_convert_messages_with_history( + messages=_tool_call_messages("call_50e7e0fe0989464a89f188eda443"), + model="gemini-2.5-flash", + custom_llm_provider=custom_llm_provider, + ) + + for content in contents: + for part in content.get("parts", []): + fc = part.get("function_call") + if fc is not None: + assert "id" not in fc, f"gemini-2.5 payload must not include id: {fc}" + fr = part.get("function_response") + if fr is not None: + assert "id" not in fr, f"gemini-2.5 payload must not include id: {fr}" + + +def test_vertex_ai_forwarded_function_call_id_strips_thought_signature_suffix(): + """The thought signature rides along on the OpenAI id but must not reach Vertex. + + Vertex now sees this code path for the first time, so the suffix has to be stripped here too. + """ + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + from litellm.litellm_core_utils.prompt_templates.factory import ( + THOUGHT_SIGNATURE_SEPARATOR, + ) + + bare_id = "call_50e7e0fe0989464a89f188eda443" + contents = _gemini_convert_messages_with_history( + messages=_tool_call_messages(f"{bare_id}{THOUGHT_SIGNATURE_SEPARATOR}sig123"), + model="gemini-3.5-flash", + custom_llm_provider="vertex_ai", + ) + + _, function_response_ids = _collect_function_call_ids(contents) + assert function_response_ids == [bare_id] + + +@pytest.mark.parametrize("model", ["gemini-3.5-flash", "gemini-2.5-flash"]) +def test_tool_response_without_matching_tool_call_is_rejected(model): + """An unpairable tool result must raise, not ship a functionResponse with no matching call.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + messages = [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_50e7e0fe0989464a89f188eda443", + "type": "function", + "function": { + "name": "read", + "arguments": '{"filePath": "/tmp"}', + }, + } + ], + }, + {"role": "tool", "content": "ok"}, + ] + + with pytest.raises(Exception, match="Missing corresponding tool call"): + _gemini_convert_messages_with_history( + messages=messages, + model=model, + custom_llm_provider="vertex_ai", + ) def test_reasoning_effort_maps_to_thinking_level_gemini_3(): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index b3c0dcd1681..4be2bb053ef 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -7645,3 +7645,346 @@ class TestSessionBearerEgressScrub: assert oauth2 is None assert "authorization" not in {k.lower() for k in raw} assert per_server == {"github": {"Authorization": "Bearer gh_injected_upstream"}} + + +# --------------------------------------------------------------------------- +# Internal-user (human) MCP entitlement tests +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +class TestUserMCPEntitlement: + """The entitlement attached to the HUMAN, read at both list time and tool-call time. + + A key's object_permission scopes the credential and a team's scopes the group; the user's own + scopes the person, so it must cap every key they hold and every tool those keys may invoke. + """ + + def _auth(self, user_id: str = "human-1", **kwargs) -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test", user_id=user_id, **kwargs) + + def _perm(self, *, servers=None, access_groups=None, tool_permissions=None) -> LiteLLM_ObjectPermissionTable: + return LiteLLM_ObjectPermissionTable( + object_permission_id="perm-human-1", + mcp_servers=servers if servers is not None else [], + mcp_access_groups=access_groups if access_groups is not None else [], + mcp_tool_permissions=tool_permissions, + ) + + @contextlib.contextmanager + def _entitled(self, perm): + """Patch the human's entitlement lookup. ``perm`` may be a permission row, None, or an + exception instance to raise (an entitlement that cannot be resolved).""" + side_effect = perm if isinstance(perm, Exception) else None + with patch.object( + MCPRequestHandler, + "_get_user_object_permission", + new_callable=AsyncMock, + return_value=None if side_effect else perm, + side_effect=side_effect, + ) as patched: + yield patched + + @contextlib.contextmanager + def _key_and_team_servers(self, key_servers, team_servers): + with ( + patch.object( + MCPRequestHandler, + "_get_allowed_mcp_servers_for_key", + new_callable=AsyncMock, + return_value=key_servers, + ), + patch.object( + MCPRequestHandler, + "_get_allowed_mcp_servers_for_team", + new_callable=AsyncMock, + return_value=team_servers, + ), + patch.object( + MCPRequestHandler, + "_get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ), + patch.object( + MCPRequestHandler, + "_get_key_access_group_mcp_server_extras", + new_callable=AsyncMock, + return_value=[], + ), + ): + yield + + async def test_entitlement_caps_the_servers_the_key_reaches(self): + """The key grants two servers; the human is entitled to one, so only that one resolves.""" + with self._key_and_team_servers(["srv-a", "srv-b"], []): + with self._entitled(self._perm(servers=["srv-a"])): + result = await MCPRequestHandler.get_allowed_mcp_servers(self._auth()) + assert result == ["srv-a"] + + async def test_entitlement_never_widens_the_key(self): + """A human entitled to a server their key does not grant still cannot reach it: the level is a + ceiling, so it intersects rather than unions.""" + with self._key_and_team_servers(["srv-a"], []): + with self._entitled(self._perm(servers=["srv-a", "srv-elsewhere"])): + result = await MCPRequestHandler.get_allowed_mcp_servers(self._auth()) + assert result == ["srv-a"] + + async def test_no_entitlement_places_no_ceiling(self): + """A human with no entitlement row leaves the key/team result untouched.""" + with self._key_and_team_servers(["srv-a", "srv-b"], []): + with self._entitled(None): + result = await MCPRequestHandler.get_allowed_mcp_servers(self._auth()) + assert sorted(result) == ["srv-a", "srv-b"] + + async def test_unresolvable_entitlement_denies_every_server(self): + """A KNOWN entitlement whose contents cannot be read must deny, not fall back to the key's + wider scope.""" + with self._key_and_team_servers(["srv-a", "srv-b"], []): + with self._entitled(ValueError("permission row unreadable")): + result = await MCPRequestHandler.get_allowed_mcp_servers(self._auth()) + assert result == [] + + async def test_entitlement_caps_the_tools_the_key_reaches(self): + """Tool-level: the key allows three tools on the server, the human is entitled to one.""" + key_perm = self._perm(tool_permissions={"srv-a": ["read", "write", "delete"]}) + with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_perm): + with patch.object( + MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock, return_value=None + ): + with self._entitled(self._perm(tool_permissions={"srv-a": ["read"]})): + result = await MCPRequestHandler.get_allowed_tools_for_server("srv-a", self._auth()) + assert result == ["read"] + + async def test_entitlement_alone_restricts_tools_on_an_otherwise_unrestricted_key(self): + """An unrestricted key (no tool permissions of its own) is still bound by the human's tools.""" + with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=None): + with patch.object( + MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock, return_value=None + ): + with self._entitled(self._perm(tool_permissions={"srv-a": ["read"]})): + result = await MCPRequestHandler.get_allowed_tools_for_server("srv-a", self._auth()) + assert result == ["read"] + + async def test_entitlement_on_another_server_does_not_restrict_this_one(self): + """Tool grants are per server: naming tools on srv-b places no bound on srv-a.""" + with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=None): + with patch.object( + MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock, return_value=None + ): + with self._entitled(self._perm(tool_permissions={"srv-b": ["read"]})): + result = await MCPRequestHandler.get_allowed_tools_for_server("srv-a", self._auth()) + assert result is None + + async def test_unresolvable_entitlement_denies_every_tool(self): + """Fail closed on the tool axis too. The caller's own except-handler treats a raise as + allow-all for key auth, so the ceiling must return the empty allowlist itself.""" + with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=None): + with patch.object( + MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock, return_value=None + ): + with self._entitled(ValueError("permission row unreadable")): + result = await MCPRequestHandler.get_allowed_tools_for_server("srv-a", self._auth()) + assert result == [] + + async def test_tool_call_is_rejected_at_call_time(self): + """The end-to-end contract: a tool the human is not entitled to is refused when INVOKED, not + merely hidden from the advertised list.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="srv-a", + name="srv-a", + server_name="srv-a", + url="https://srv-a.example.com", + transport=MCPTransport.http, + ) + with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=None): + with patch.object( + MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock, return_value=None + ): + with self._entitled(self._perm(tool_permissions={"srv-a": ["read"]})): + await global_mcp_server_manager.check_tool_permission_for_key_team( + tool_name="read", server=server, user_api_key_auth=self._auth() + ) + with pytest.raises(HTTPException) as exc: + await global_mcp_server_manager.check_tool_permission_for_key_team( + tool_name="delete", server=server, user_api_key_auth=self._auth() + ) + assert exc.value.status_code == 403 + + async def test_keyless_admitted_source_is_not_capped_by_the_user_level(self): + """A gateway-admitted human resolves as a UNION over their own grants plus their teams', and + their own grants ARE the user source there. Re-applying them as a ceiling per source would + make one team's narrower scope silently bound another's, so the level is skipped.""" + with self._key_and_team_servers(["srv-a", "srv-b"], []): + with self._entitled(self._perm(servers=["srv-a"])) as lookup: + result = await MCPRequestHandler.get_allowed_mcp_servers(self._auth(), keyless_source=True) + assert sorted(result) == ["srv-a", "srv-b"] + lookup.assert_not_awaited() + + async def test_keyless_admitted_source_tools_are_not_capped_by_the_user_level(self): + with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=None): + with patch.object( + MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock, return_value=None + ): + with self._entitled(self._perm(tool_permissions={"srv-a": ["read"]})) as lookup: + result = await MCPRequestHandler.get_allowed_tools_for_server( + "srv-a", self._auth(), keyless_source=True + ) + assert result is None + lookup.assert_not_awaited() + + async def test_servers_named_only_under_tool_permissions_are_entitled(self): + """Granting one tool on a server entitles the human to that server, so an admin never has to + name it twice.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + global_mcp_server_manager.registry["srv-a"] = MCPServer( + server_id="srv-a", + name="srv-a", + server_name="srv-a", + url="https://srv-a.example.com", + transport=MCPTransport.http, + ) + try: + with self._entitled(self._perm(tool_permissions={"srv-a": ["read"]})): + with patch.object( + MCPRequestHandler, + "_get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_user(self._auth()) + finally: + global_mcp_server_manager.registry.pop("srv-a", None) + assert result == ["srv-a"] + + async def test_places_ceiling_is_true_when_unresolvable(self): + """``_user_places_mcp_ceiling`` gates the admin shortcut that hands over the whole registry, so + an entitlement it cannot resolve must still count as a ceiling.""" + with self._entitled(ValueError("boom")): + assert await MCPRequestHandler._user_places_mcp_ceiling(self._auth()) is True + with self._entitled(None): + assert await MCPRequestHandler._user_places_mcp_ceiling(self._auth()) is False + + +@pytest.mark.asyncio +class TestGetUserObjectPermission: + """Resolution of the ``user_id -> object_permission_id -> grants`` chain.""" + + def _prisma_with_user(self, user_row): + prisma_client = MagicMock() + prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + return prisma_client + + async def test_resolves_through_the_shared_permission_cache(self): + from litellm.caching.dual_cache import DualCache + + user_row = MagicMock() + user_row.object_permission_id = "perm-1" + prisma_client = self._prisma_with_user(user_row) + auth = UserAPIKeyAuth(api_key="sk-test", user_id="human-shared") + expected = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + new_callable=AsyncMock, + return_value=expected, + ) as mock_get_perm, + ): + assert await MCPRequestHandler._get_user_object_permission(auth) is expected + assert mock_get_perm.await_args.kwargs["object_permission_id"] == "perm-1" + + # The user_id -> object_permission_id link is cached, so the user row is read once. + prisma_client.db.litellm_usertable.find_unique.reset_mock() + await MCPRequestHandler._get_user_object_permission(auth) + prisma_client.db.litellm_usertable.find_unique.assert_not_called() + + async def test_caches_a_sentinel_for_a_human_with_no_entitlement(self): + """A human without an entitlement is the common case and must cost no DB read per request.""" + from litellm.caching.dual_cache import DualCache + + user_row = MagicMock() + user_row.object_permission_id = None + prisma_client = self._prisma_with_user(user_row) + auth = UserAPIKeyAuth(api_key="sk-test", user_id="human-no-perm") + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock) as mock_get_perm, + ): + assert await MCPRequestHandler._get_user_object_permission(auth) is None + assert await MCPRequestHandler._get_user_object_permission(auth) is None + mock_get_perm.assert_not_awaited() + prisma_client.db.litellm_usertable.find_unique.assert_awaited_once() + + async def test_missing_user_row_places_no_ceiling(self): + """Whether this human is entitled at all is unknown when their row is absent, which is the + state before the level existed, so it must not deny.""" + from litellm.caching.dual_cache import DualCache + + prisma_client = self._prisma_with_user(None) + auth = UserAPIKeyAuth(api_key="sk-test", user_id="ghost") + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + ): + assert await MCPRequestHandler._get_user_object_permission(auth) is None + + async def test_unreadable_user_row_places_no_ceiling(self): + from litellm.caching.dual_cache import DualCache + + prisma_client = MagicMock() + prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=Exception("db down")) + auth = UserAPIKeyAuth(api_key="sk-test", user_id="human-db-down") + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + ): + assert await MCPRequestHandler._get_user_object_permission(auth) is None + + async def test_named_but_unreadable_permission_raises(self): + """A KNOWN entitlement with unknown contents is indeterminate: it must surface so the callers + can deny rather than serve the wider key scope.""" + from litellm.caching.dual_cache import DualCache + + user_row = MagicMock() + user_row.object_permission_id = "perm-gone" + prisma_client = self._prisma_with_user(user_row) + auth = UserAPIKeyAuth(api_key="sk-test", user_id="human-dangling") + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + new_callable=AsyncMock, + return_value=None, + ), + ): + with pytest.raises(ValueError): + await MCPRequestHandler._get_user_object_permission(auth) + + async def test_no_user_id_places_no_ceiling(self): + assert await MCPRequestHandler._get_user_object_permission(UserAPIKeyAuth(api_key="sk-test")) is None + assert await MCPRequestHandler._get_user_object_permission(None) is None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py b/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py index 5dbad53948b..2e286a237c4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py @@ -1,11 +1,14 @@ """Tests for the MCP guardrail translation handler.""" import pytest +from mcp.types import CallToolResult, ImageContent, TextContent +from litellm.exceptions import BlockedPiiEntityError from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import ( MCPGuardrailTranslationHandler, ) +from litellm.types.utils import GenericGuardrailAPIInputs class MockGuardrail(CustomGuardrail): @@ -80,3 +83,304 @@ async def test_process_input_messages_handles_minimal_data(): tools = guardrail.last_inputs.get("tools", []) assert len(tools) == 1 assert tools[0]["function"]["name"] == "simple_tool" + + +class MaskingGuardrail(CustomGuardrail): + """Guardrail that rewrites every scanned text, recording what it saw.""" + + def __init__(self, masked_texts=None, raises=None): + super().__init__(guardrail_name="masking-mcp-guardrail") + self.masked_texts = masked_texts + self.raises = raises + self.call_count = 0 + self.last_inputs = None + self.last_input_type = None + self.last_request_data = None + + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + self.call_count += 1 + self.last_inputs = inputs + self.last_input_type = input_type + self.last_request_data = request_data + if self.raises is not None: + raise self.raises + if self.masked_texts is None: + return inputs + return GenericGuardrailAPIInputs(texts=list(self.masked_texts)) + + +@pytest.mark.asyncio +async def test_process_output_response_masks_text_content(): + """Masked text returned by the guardrail must land in the tool result.""" + handler = MCPGuardrailTranslationHandler() + guardrail = MaskingGuardrail(masked_texts=["email ", "call "]) + result = CallToolResult( + content=[ + TextContent(type="text", text="email jane@example.com"), + TextContent(type="text", text="call 415-555-0132"), + ], + isError=False, + ) + + returned = await handler.process_output_response( + response=result, + guardrail_to_apply=guardrail, + request_data={"mcp_tool_name": "echo"}, + ) + + assert guardrail.call_count == 1 + assert guardrail.last_input_type == "response" + assert guardrail.last_inputs["texts"] == ["email jane@example.com", "call 415-555-0132"] + assert [item.text for item in returned.content] == ["email ", "call "] + assert [item.text for item in result.content] == ["email ", "call "] + + +@pytest.mark.asyncio +async def test_process_output_response_masks_dict_shaped_result(): + """A dict-shaped tool result (REST/JSON-RPC payload) must be masked too.""" + handler = MCPGuardrailTranslationHandler() + guardrail = MaskingGuardrail(masked_texts=[""]) + result = {"content": [{"type": "text", "text": "jane@example.com"}], "isError": False} + + returned = await handler.process_output_response(response=result, guardrail_to_apply=guardrail) + + assert returned["content"][0]["text"] == "" + assert returned["content"][0]["type"] == "text" + + +@pytest.mark.asyncio +async def test_process_output_response_propagates_block(): + """A guardrail rejecting the tool result must not be swallowed.""" + handler = MCPGuardrailTranslationHandler() + guardrail = MaskingGuardrail( + raises=BlockedPiiEntityError(entity_type="EMAIL_ADDRESS", guardrail_name="masking-mcp-guardrail") + ) + result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False) + + with pytest.raises(BlockedPiiEntityError): + await handler.process_output_response(response=result, guardrail_to_apply=guardrail) + + +@pytest.mark.asyncio +async def test_process_output_response_skips_non_text_content(): + """A result carrying no text content must not be sent to the guardrail.""" + handler = MCPGuardrailTranslationHandler() + guardrail = MaskingGuardrail(masked_texts=["should not be used"]) + result = CallToolResult( + content=[ImageContent(type="image", data="aGk=", mimeType="image/png")], + isError=False, + ) + + returned = await handler.process_output_response(response=result, guardrail_to_apply=guardrail) + + assert guardrail.call_count == 0 + assert returned is result + + +@pytest.mark.asyncio +async def test_process_output_response_handles_result_without_content(): + """An unexpected result shape must be passed through, not crash the tool call.""" + handler = MCPGuardrailTranslationHandler() + guardrail = MaskingGuardrail(masked_texts=["should not be used"]) + + returned = await handler.process_output_response(response={"error": "boom"}, guardrail_to_apply=guardrail) + + assert guardrail.call_count == 0 + assert returned == {"error": "boom"} + + +@pytest.mark.asyncio +async def test_process_output_response_leaves_result_unmasked_on_text_count_mismatch(): + """A guardrail returning the wrong number of texts must not shuffle content.""" + handler = MCPGuardrailTranslationHandler() + guardrail = MaskingGuardrail(masked_texts=[""]) + result = CallToolResult( + content=[ + TextContent(type="text", text="jane@example.com"), + TextContent(type="text", text="415-555-0132"), + ], + isError=False, + ) + + returned = await handler.process_output_response(response=result, guardrail_to_apply=guardrail) + + assert [item.text for item in returned.content] == ["jane@example.com", "415-555-0132"] + + +class SubstitutingGuardrail(CustomGuardrail): + """Masks one substring wherever it appears, across however many texts it is given.""" + + def __init__(self, needle: str, replacement: str): + super().__init__(guardrail_name="substituting-mcp-guardrail") + self.needle = needle + self.replacement = replacement + self.seen_texts: list = [] + + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + self.seen_texts = list(inputs.get("texts") or []) + return GenericGuardrailAPIInputs( + texts=[text.replace(self.needle, self.replacement) for text in self.seen_texts] + ) + + +@pytest.mark.asyncio +async def test_structured_content_is_masked_alongside_content(): + """structuredContent goes to the client too, so it must be masked, not just content.""" + handler = MCPGuardrailTranslationHandler() + guardrail = SubstitutingGuardrail("jane@example.com", "") + response = CallToolResult( + content=[TextContent(type="text", text="email jane@example.com")], + structuredContent={"contact": {"email": "jane@example.com"}, "balance": 42.0}, + isError=False, + ) + + returned = await handler.process_output_response(response=response, guardrail_to_apply=guardrail) + + assert returned.content[0].text == "email " + assert returned.structuredContent == {"contact": {"email": ""}, "balance": 42.0} + + +@pytest.mark.asyncio +async def test_value_present_only_in_structured_content_is_masked(): + """The gap this closes: a sensitive value that never appears in the text content. + + Scanning only content would hand it to the guardrail never, so it would reach + the client unscanned behind a result that looks inspected. + """ + handler = MCPGuardrailTranslationHandler() + guardrail = SubstitutingGuardrail("jane@example.com", "") + response = CallToolResult( + content=[TextContent(type="text", text="lookup complete")], + structuredContent={"records": [{"email": "jane@example.com"}]}, + isError=False, + ) + + returned = await handler.process_output_response(response=response, guardrail_to_apply=guardrail) + + assert "jane@example.com" in guardrail.seen_texts + assert returned.structuredContent == {"records": [{"email": ""}]} + assert returned.content[0].text == "lookup complete" + + +@pytest.mark.asyncio +async def test_structured_content_without_a_match_is_untouched(): + """Unrelated structured data keeps its values and its types.""" + handler = MCPGuardrailTranslationHandler() + guardrail = SubstitutingGuardrail("jane@example.com", "") + response = CallToolResult( + content=[TextContent(type="text", text="lookup complete")], + structuredContent={"record_id": "C-1001", "balance": 42.0, "active": True, "note": None}, + isError=False, + ) + + returned = await handler.process_output_response(response=response, guardrail_to_apply=guardrail) + + assert returned.structuredContent == {"record_id": "C-1001", "balance": 42.0, "active": True, "note": None} + + +@pytest.mark.asyncio +async def test_structured_content_nested_too_deeply_is_blocked(): + """Too deep to walk must block rather than pass the deeper values unscanned.""" + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server.utils import MAX_STRUCTURED_CONTENT_SCAN_DEPTH + + handler = MCPGuardrailTranslationHandler() + guardrail = SubstitutingGuardrail("jane@example.com", "") + nested: dict = {"leaf": "jane@example.com"} + for _ in range(MAX_STRUCTURED_CONTENT_SCAN_DEPTH + 1): + nested = {"next": nested} + response = CallToolResult( + content=[TextContent(type="text", text="lookup complete")], + structuredContent=nested, + isError=False, + ) + + with pytest.raises(HTTPException) as exc_info: + await handler.process_output_response(response=response, guardrail_to_apply=guardrail) + + assert exc_info.value.status_code == 400 + + +def test_too_deep_json_returns_a_sentinel_rather_than_raising(): + """The too-deep signal must be a return value, not a custom exception. + + mcp_server/utils.py is reloaded by tests that override its environment-backed + constants, which gives any exception class defined there a fresh identity and + lets it escape a caller's except clause; under xdist that surfaced as a failure + in an unrelated shard. A sentinel has no identity to lose. Asserted directly on + the helper so this pins the contract without reloading the module and leaking + that reload into other tests. + """ + from litellm.proxy._experimental.mcp_server.utils import ( + MAX_STRUCTURED_CONTENT_SCAN_DEPTH, + json_string_leaves, + ) + + nested: dict = {"leaf": "jane@example.com"} + for _ in range(MAX_STRUCTURED_CONTENT_SCAN_DEPTH + 1): + nested = {"next": nested} + + assert json_string_leaves(nested) is None + assert json_string_leaves({"a": "b"}) == ((("a",), "b"),) + + +@pytest.mark.asyncio +async def test_sensitive_structured_content_key_is_blocked(): + """A dict key is client-visible but not rewritable, so a match must block. + + Maps keyed by an identifier are a common API shape, and renaming the key would + change the payload contract rather than redact a value; the content filter takes + the same position on MCP tool call arguments. + """ + from fastapi import HTTPException + + handler = MCPGuardrailTranslationHandler() + guardrail = SubstitutingGuardrail("jane@example.com", "") + response = CallToolResult( + content=[TextContent(type="text", text="lookup complete")], + structuredContent={"jane@example.com": {"balance": 42.0}}, + isError=False, + ) + + with pytest.raises(HTTPException) as exc_info: + await handler.process_output_response(response=response, guardrail_to_apply=guardrail) + + assert exc_info.value.status_code == 400 + assert "non-rewritable" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_sensitive_structured_content_numeric_value_is_blocked(): + """A numeric value cannot be masked in place either, so a match must block.""" + from fastapi import HTTPException + + handler = MCPGuardrailTranslationHandler() + guardrail = SubstitutingGuardrail("4155550199", "") + response = CallToolResult( + content=[TextContent(type="text", text="lookup complete")], + structuredContent={"phone": 4155550199}, + isError=False, + ) + + with pytest.raises(HTTPException) as exc_info: + await handler.process_output_response(response=response, guardrail_to_apply=guardrail) + + assert exc_info.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_clean_structured_content_keys_do_not_block(): + """Ordinary keys and numbers must pass through untouched.""" + handler = MCPGuardrailTranslationHandler() + guardrail = SubstitutingGuardrail("jane@example.com", "") + response = CallToolResult( + content=[TextContent(type="text", text="email jane@example.com")], + structuredContent={"record_id": "C-1001", "balance": 42.0, "count": 3}, + isError=False, + ) + + returned = await handler.process_output_response(response=response, guardrail_to_apply=guardrail) + + assert returned.content[0].text == "email " + assert returned.structuredContent == {"record_id": "C-1001", "balance": 42.0, "count": 3} diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py index 375ec022115..85a19331777 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -1,6 +1,7 @@ """Tests for the aggregate gateway DCR flow (register, authorize, complete, token).""" import hashlib +import html import json from base64 import urlsafe_b64encode from datetime import datetime, timedelta, timezone @@ -16,7 +17,10 @@ from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( GATEWAY_AUTH_CODE_PREFIX, GATEWAY_AUTH_CODE_TTL_SECONDS, GATEWAY_DCR_CLIENT_ID_PREFIX, + MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS, + _AUTH_CODE_DEBUG_KEY, _GatewayAuthCode, + _open_sealed, _seal, aggregate_authorize, aggregate_token, @@ -588,3 +592,206 @@ async def test_single_use_guard_fails_closed_when_redis_errors(): guard = _SingleUseGuard(cache) assert await guard.claim("jti-fault", 60) is False # fail closed, not a fallback count of 1 + + +LOOPBACK_REDIRECT_URI = "http://localhost:3118/callback" + + +async def _complete(redirect_uri: str, delivery, cookies=None, handle=None, session_user_id="u1"): + client_id = (await _register([redirect_uri]))["client_id"] + if cookies is None: + handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1", redirect_uri=redirect_uri)) + response = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id=session_user_id, + cache=DualCache(), + delivery=delivery, + ) + return client_id, response + + +def _callback_url_from_page(response) -> str: + import html as html_lib + import re + + match = re.search(r'value="([^"]+)"', response.body.decode()) + assert match is not None + return html_lib.unescape(match.group(1)) + + +@pytest.mark.asyncio +async def test_manual_delivery_renders_pasteable_callback_url_for_loopback_client(): + """The LIT-4863 headless path: a loopback client on another machine gets the callback + URL on a page instead of a dead 303, and the code on that page is a full-fidelity + authorization code (PKCE-bound, single-use, redeemable at /token).""" + client_id, response = await _complete(LOOPBACK_REDIRECT_URI, delivery="manual") + assert response.status_code == 200 + assert response.headers["content-type"].startswith("text/html") + assert response.headers["cache-control"] == "no-store" + assert f"{CONNECT_FLOW_COOKIE_PREFIX}" in response.headers["set-cookie"] + + callback_url = _callback_url_from_page(response) + parsed = urlparse(callback_url) + assert f"{parsed.scheme}://{parsed.netloc}{parsed.path}" == LOOPBACK_REDIRECT_URI + params = parse_qs(parsed.query) + assert params["state"] == ["client-state-123"] + code = params["code"][0] + assert code.startswith(GATEWAY_AUTH_CODE_PREFIX) + + cache = DualCache() + token_response = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code=code, + redirect_uri=LOOPBACK_REDIRECT_URI, + client_id=client_id, + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=cache, + ) + assert token_response.status_code == 200 + + replay = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code=code, + redirect_uri=LOOPBACK_REDIRECT_URI, + client_id=client_id, + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=cache, + ) + assert json.loads(replay.body)["error"] == "invalid_grant" + + +@pytest.mark.asyncio +async def test_manual_delivery_code_gets_the_longer_ttl_and_redirect_code_does_not(): + _, manual = await _complete(LOOPBACK_REDIRECT_URI, delivery="manual") + manual_code = parse_qs(urlparse(_callback_url_from_page(manual)).query)["code"][0] + opened_manual = _open_sealed(manual_code, GATEWAY_AUTH_CODE_PREFIX, _GatewayAuthCode, _AUTH_CODE_DEBUG_KEY) + assert opened_manual is not None + assert opened_manual.exp - opened_manual.iat == MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS + + _, redirected = await _complete(LOOPBACK_REDIRECT_URI, delivery=None) + redirect_code = parse_qs(urlparse(redirected.headers["location"]).query)["code"][0] + opened_redirect = _open_sealed(redirect_code, GATEWAY_AUTH_CODE_PREFIX, _GatewayAuthCode, _AUTH_CODE_DEBUG_KEY) + assert opened_redirect is not None + assert opened_redirect.exp - opened_redirect.iat == GATEWAY_AUTH_CODE_TTL_SECONDS + + +@pytest.mark.asyncio +@pytest.mark.parametrize("delivery", [None, "redirect"]) +async def test_loopback_client_still_redirects_when_manual_not_requested(delivery): + _, response = await _complete(LOOPBACK_REDIRECT_URI, delivery=delivery) + assert response.status_code == 303 + assert response.headers["location"].startswith(LOOPBACK_REDIRECT_URI) + + +@pytest.mark.asyncio +async def test_manual_delivery_is_ignored_for_routable_redirect_uri(): + """A routable redirect URI works from any browser by construction, so manual is a + no-op there and the flow keeps its normal shape.""" + _, response = await _complete(REDIRECT_URI, delivery="manual") + assert response.status_code == 303 + assert response.headers["location"].startswith(REDIRECT_URI) + + +@pytest.mark.asyncio +async def test_unknown_delivery_value_is_rejected_before_the_flow_is_consumed(): + """A typo'd delivery must not burn the single-use flow: the user fixes the form and + finishes normally.""" + client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] + handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1", redirect_uri=LOOPBACK_REDIRECT_URI)) + + rejected = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="u1", + cache=DualCache(), + delivery="carrier-pigeon", + ) + assert rejected.status_code == 400 + assert json.loads(rejected.body)["error"] == "invalid_request" + + retried = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="u1", + cache=DualCache(), + delivery="manual", + ) + assert retried.status_code == 200 + + +@pytest.mark.asyncio +async def test_manual_delivery_page_escapes_client_influenced_values(): + """redirect_uri (and everything else on the page) is client-registered input; a quote + or tag in its path must render inert.""" + hostile_uri = 'http://127.0.0.1:9/cb">' + _, response = await _complete(hostile_uri, delivery="manual") + assert response.status_code == 200 + body = response.body.decode() + assert "" not in body + assert "<script>" in body + + +class _TtlRecordingCache(DualCache): + """Captures the TTL of every single-use claim recorded through the in-memory arm.""" + + def __init__(self): + super().__init__() + self.claim_ttls: dict = {} + + async def async_increment_cache(self, key, value, ttl=None, **kwargs): + self.claim_ttls[key] = ttl + return await super().async_increment_cache(key, value, ttl=ttl, **kwargs) + + +@pytest.mark.asyncio +async def test_used_code_marker_outlives_the_manually_delivered_code(): + """Veria review finding on the LIT-4863 change: a manual code lives 300s, but the + used-code marker was retained for the 120s redirect lifetime plus buffer, so a client + could redeem, wait out the marker, and redeem the still-valid code again. The marker's + TTL must cover the code's own remaining lifetime plus the claim buffer.""" + client_id, response = await _complete(LOOPBACK_REDIRECT_URI, delivery="manual") + code = parse_qs(urlparse(_callback_url_from_page(response)).query)["code"][0] + + cache = _TtlRecordingCache() + token_response = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code=code, + redirect_uri=LOOPBACK_REDIRECT_URI, + client_id=client_id, + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=cache, + ) + assert token_response.status_code == 200 + + marker_ttls = [ttl for key, ttl in cache.claim_ttls.items() if key.startswith("mcp_gateway_dcr_code_used:")] + assert len(marker_ttls) == 1 + assert marker_ttls[0] >= MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS + + +@pytest.mark.asyncio +@pytest.mark.parametrize("redirect_uri", [LOOPBACK_REDIRECT_URI, "http://127.0.0.1:9/cb$(whoami)&calc& rem x"]) +async def test_manual_delivery_page_renders_the_url_as_data_never_as_a_shell_command(redirect_uri): + """Two review rounds proved no single command string is safe across POSIX shells, + cmd.exe, and PowerShell (single quotes are not quoting in cmd.exe; percent expands + there even inside double quotes), so the page must render the callback URL as data + only and never as a ready-to-paste command.""" + _, response = await _complete(redirect_uri, delivery="manual") + assert response.status_code == 200 + body = response.body.decode() + assert "" not in body + assert 'curl "' not in body + assert "curl '" not in body + assert 'value="' in body diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py index c063915e2e8..f6bd79c5d2d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py @@ -240,10 +240,10 @@ async def test_explicit_null_clears_upstream_resource_and_keeps_the_rest_of_the_ @pytest.mark.asyncio -async def test_url_change_clears_stale_discovered_oauth_fields(): - """Re-pointing the server url at a potentially different upstream must clear the discovered or - trust-on-first-use OAuth issuer and endpoints, so the new upstream re-discovers instead of - anchoring on the previous upstream's issuer (RFC 8414 §3.3 against a stale anchor).""" +async def test_url_change_clears_stale_oauth_fields(): + """Re-pointing the server url at a potentially different upstream must clear the OAuth issuer and + endpoints, so the new upstream re-discovers instead of anchoring on the previous upstream's issuer + (RFC 8414 §3.3 against a stale anchor).""" mock_prisma = _mock_prisma() existing = MagicMock() existing.auth_type = "oauth2" @@ -350,11 +350,13 @@ async def test_repointing_pinned_issuer_clears_stale_endpoints_keeps_new_issuer( @pytest.mark.asyncio -async def test_establishing_issuer_first_time_preserves_discovered_fields(): - """Establishing an issuer for the first time (None -> X), which is exactly what the trust-on-first-use - discovery write-back does, must NOT clear the endpoints or oauth2_flow it discovered in the same - write. Only an issuer that was already pinned and is now changed or cleared invalidates its - endpoints, so the discovery persist cannot wipe the fields it just resolved.""" +async def test_establishing_issuer_first_time_preserves_endpoints_set_in_the_same_write(): + """Establishing an issuer for the first time (None -> X) must NOT clear endpoints or oauth2_flow + submitted in the same write. Only an issuer that was already pinned and is now changed or cleared + invalidates its endpoints, so an admin configuring an issuer and its endpoints together keeps + both. The write-back this once guarded (trust-on-first-use discovery stamping the issuer it had + just resolved) no longer exists; the db.py rule it relies on still governs admin writes, which is + what this now covers.""" mock_prisma = _mock_prisma() existing = MagicMock() existing.auth_type = "oauth2" @@ -370,7 +372,7 @@ async def test_establishing_issuer_first_time_preserves_discovered_fields(): token_url="https://discovered-idp.example.com/token", oauth2_flow="authorization_code", ) - await update_mcp_server(mock_prisma, data, "mcp_oauth_discovery") + await update_mcp_server(mock_prisma, data, "some-admin@example.com") data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"] assert data_dict["issuer"] == "https://discovered-idp.example.com" @@ -380,9 +382,9 @@ async def test_establishing_issuer_first_time_preserves_discovered_fields(): @pytest.mark.asyncio -async def test_unchanged_url_does_not_clear_discovered_oauth_fields(): - """A partial update that resends the same url (or omits it) must not clear the discovered OAuth - fields, so a routine save does not force needless re-discovery.""" +async def test_unchanged_url_does_not_clear_oauth_fields(): + """A partial update that resends the same url (or omits it) must not clear the OAuth fields, so a + routine save does not force needless re-discovery.""" mock_prisma = _mock_prisma() existing = MagicMock() existing.auth_type = "oauth2" 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 1753b0d92a8..c0affdf46b3 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 @@ -7130,6 +7130,14 @@ def _mock_mcp_logging_obj() -> MagicMock: return logging_obj +def _mock_mcp_proxy_logging() -> MagicMock: + """ProxyLogging stand-in whose post_mcp_call_hook passes the result through.""" + proxy_logging_mock = MagicMock() + proxy_logging_mock.post_call_failure_hook = AsyncMock() + proxy_logging_mock.post_mcp_call_hook = AsyncMock(side_effect=lambda response, **_: response) + return proxy_logging_mock + + def test_extract_mcp_tool_result_error_message(): from litellm.proxy._experimental.mcp_server.utils import ( extract_mcp_tool_result_error_message, @@ -7160,8 +7168,7 @@ async def test_fire_mcp_tool_call_logging_iserror_logs_failure(): from litellm.proxy._experimental.mcp_server.exceptions import MCPToolResultError logging_obj = _mock_mcp_logging_obj() - proxy_logging_mock = MagicMock() - proxy_logging_mock.post_call_failure_hook = AsyncMock() + proxy_logging_mock = _mock_mcp_proxy_logging() user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock): @@ -7199,8 +7206,7 @@ async def test_fire_mcp_tool_call_logging_success_path_unchanged(): ) logging_obj = _mock_mcp_logging_obj() - proxy_logging_mock = MagicMock() - proxy_logging_mock.post_call_failure_hook = AsyncMock() + proxy_logging_mock = _mock_mcp_proxy_logging() result = _call_tool_result(False, "all good") with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock): @@ -7229,8 +7235,7 @@ async def test_fire_mcp_tool_call_logging_iserror_without_auth_skips_failure_hoo ) logging_obj = _mock_mcp_logging_obj() - proxy_logging_mock = MagicMock() - proxy_logging_mock.post_call_failure_hook = AsyncMock() + proxy_logging_mock = _mock_mcp_proxy_logging() with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock): await _fire_mcp_tool_call_logging( @@ -7256,8 +7261,7 @@ async def test_fire_mcp_tool_call_logging_strips_credentials_from_failure_hook() ) logging_obj = _mock_mcp_logging_obj() - proxy_logging_mock = MagicMock() - proxy_logging_mock.post_call_failure_hook = AsyncMock() + proxy_logging_mock = _mock_mcp_proxy_logging() user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") request_data = { "name": "explode", @@ -7528,8 +7532,7 @@ async def test_call_mcp_tool_skips_failure_hook_for_upstream_auth_error(): transport=MCPTransport.http, mcp_info={"server_name": "test_server"}, ) - proxy_logging_mock = MagicMock() - proxy_logging_mock.post_call_failure_hook = AsyncMock() + proxy_logging_mock = _mock_mcp_proxy_logging() user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") with ( @@ -7841,3 +7844,83 @@ class TestPreemptive401ModeAware: await self._run(delegate, None, has_stored_token=False) assert exc.value.status_code == 401 await self._run(delegate, self.LITELLM_KEY_HEADERS, has_stored_token=False) + + +@pytest.mark.asyncio +async def test_post_mcp_call_guardrails_return_the_rewritten_result(): + """The result a post_mcp_call guardrail rewrote must be what the caller sends back.""" + from litellm.proxy._experimental.mcp_server.server import ( + _run_post_mcp_call_guardrails, + ) + + logging_obj = _mock_mcp_logging_obj() + raw_result = _call_tool_result(False, "jane@example.com") + masked_result = _call_tool_result(False, "") + proxy_logging_mock = _mock_mcp_proxy_logging() + proxy_logging_mock.post_mcp_call_hook = AsyncMock(return_value=masked_result) + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock): + returned = await _run_post_mcp_call_guardrails( + result=raw_result, + litellm_logging_obj=logging_obj, + user_api_key_auth=UserAPIKeyAuth(api_key="test-key", user_id="test-user"), + request_data={}, + ) + + assert returned is masked_result + hook_kwargs = proxy_logging_mock.post_mcp_call_hook.await_args.kwargs + assert hook_kwargs["response"] is raw_result + assert hook_kwargs["request_data"] is logging_obj.model_call_details + + +@pytest.mark.asyncio +async def test_post_mcp_call_guardrails_run_without_a_logging_object(): + """Enforcement must not depend on logging being configured. + + A tool call dispatched without a litellm_logging_obj (tool search, and any + caller that omits it) would otherwise skip the guardrail entirely and return + the unscanned tool output to the client. + """ + from litellm.proxy._experimental.mcp_server.server import ( + _run_post_mcp_call_guardrails, + ) + + raw_result = _call_tool_result(False, "jane@example.com") + masked_result = _call_tool_result(False, "") + proxy_logging_mock = _mock_mcp_proxy_logging() + proxy_logging_mock.post_mcp_call_hook = AsyncMock(return_value=masked_result) + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock): + returned = await _run_post_mcp_call_guardrails( + result=raw_result, + litellm_logging_obj=None, + user_api_key_auth=UserAPIKeyAuth(api_key="test-key", user_id="test-user"), + request_data={"name": "fetch_record"}, + ) + + assert returned is masked_result + proxy_logging_mock.post_mcp_call_hook.assert_awaited_once() + assert proxy_logging_mock.post_mcp_call_hook.await_args.kwargs["request_data"] == {"name": "fetch_record"} + + +@pytest.mark.asyncio +async def test_post_mcp_call_guardrails_propagate_a_block(): + """A post_mcp_call guardrail rejection must propagate instead of returning the result.""" + from litellm.exceptions import BlockedPiiEntityError + from litellm.proxy._experimental.mcp_server.server import ( + _run_post_mcp_call_guardrails, + ) + + proxy_logging_mock = _mock_mcp_proxy_logging() + proxy_logging_mock.post_mcp_call_hook = AsyncMock( + side_effect=BlockedPiiEntityError(entity_type="EMAIL_ADDRESS", guardrail_name="presidio-mcp") + ) + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock): + with pytest.raises(BlockedPiiEntityError): + await _run_post_mcp_call_guardrails( + result=_call_tool_result(False, "jane@example.com"), + litellm_logging_obj=_mock_mcp_logging_obj(), + user_api_key_auth=None, + request_data={}, + ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 5f7f2267fc7..8a8dea0ba28 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -2,6 +2,7 @@ import importlib import asyncio import json import logging +import time import os import sys from datetime import datetime @@ -35,6 +36,8 @@ from mcp.types import Tool as MCPTool from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, _deserialize_json_dict, + _flow_endpoints_missing, + _oauth_endpoints_unresolved, _deserialize_json_list, _normalize_mcp_server_cost_info, _should_strip_caller_authorization, @@ -1594,21 +1597,15 @@ class TestMCPServerManager: token_url="https://idp.example.com/token", scopes=["read"], ) - with ( - patch.object( - manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=issuer_resolved) - ) as anchored, - patch.object(manager, "_persist_discovered_oauth_endpoints", new=AsyncMock()) as mock_persist, - ): + with patch.object( + manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=issuer_resolved) + ) as anchored: built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) anchored.assert_awaited_once_with("https://idp.example.com", "https://up.example.com/mcp") assert built.authorization_url == "https://idp.example.com/authorize" assert built.token_url == "https://idp.example.com/token" assert built.token_url != "https://attacker.example.com/steal" - # The issuer-anchored endpoints are never persisted into the endpoint columns, so a later - # build cannot treat them as authoritative stored values. - assert mock_persist.await_args.kwargs["is_issuer_anchored"] is True @pytest.mark.asyncio @pytest.mark.parametrize( @@ -1624,8 +1621,8 @@ class TestMCPServerManager: and PKCE verifier to the attacker (config-time RFC 9700 mix-up). The resource-driven scopes are kept, because scope selection is resource-driven (MCP Scope Selection Strategy) and scope inflation is bounded by the authorization server at consent (RFC 6749 §3.3), not by dropping - scopes on an endpoint mismatch. Both the in-memory merge and the persisted metadata drop only - the uncorroborated endpoints.""" + scopes on an endpoint mismatch. The gateway persists nothing, so the in-memory merge is the + entire behavior.""" manager = MCPServerManager() row = LiteLLM_MCPServerTable( server_id="manual-auth-url-3", @@ -1645,20 +1642,13 @@ class TestMCPServerManager: registration_url="https://attacker.example.com/register", scopes=["read", "admin"], ) - with ( - patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)), - patch.object(manager, "_persist_discovered_oauth_endpoints", new=AsyncMock()) as mock_persist, - ): + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)): built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) assert built.authorization_url == "https://idp.example.com/authorize" assert built.token_url is None assert built.registration_url is None assert built.scopes == ["read", "admin"] - persisted_metadata = mock_persist.await_args.kwargs["metadata"] - assert persisted_metadata.token_url is None - assert persisted_metadata.registration_url is None - assert persisted_metadata.scopes == ["read", "admin"] @pytest.mark.asyncio async def test_build_from_table_skips_discovery_when_all_upstream_oauth_fields_present(self): @@ -5586,388 +5576,300 @@ class TestMCPServerTimestamps: assert server.token_exchange_endpoint == "https://idp.example.com/token" @pytest.mark.asyncio - async def test_build_mcp_server_from_table_persists_discovered_obo_token_url(self): - """A DB-backed OBO server with no configured endpoint discovers token_url and must write it - back to the row, so the next rebuild skips discovery instead of re-running it every time.""" + async def test_discovery_never_writes_the_database(self): + """The #34985 regression, stated as the design invariant that fixes it: the gateway never + writes discovery results to the row. The OAuth columns and credentials.scopes carry admin + intent alone, so nothing the gateway learns can read back as an admin pin on a later build + (which is what anchored stamped servers fail-closed and 400ed /authorize). Discovery output + lives on the in-memory registry entry only, for oauth2 and OBO alike.""" manager = MCPServerManager() async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): - assert server_url == "https://example.com/mcp" - assert allow_origin_fallback is False # OBO never guesses the origin return MCPOAuthMetadata( - scopes=None, - authorization_url=None, - token_url="https://discovered.example.com/token", - registration_url=None, - ) - - manager._descovery_metadata = fake_discovery # type: ignore[attr-defined] - - record = LiteLLM_MCPServerTable( - server_id="obo-persist-1", - server_name="obo_persist", - url="https://example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2_token_exchange, - credentials={"client_id": "cid", "client_secret": "csec", "audience": "aud"}, - ) - - update_mock = AsyncMock() - repo_instance = MagicMock() - repo_instance.table.update = update_mock - with ( - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", - return_value=repo_instance, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - server = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) - - assert server.token_url == "https://discovered.example.com/token" - update_mock.assert_awaited_once() - assert update_mock.call_args.kwargs["where"] == {"server_id": "obo-persist-1"} - assert update_mock.call_args.kwargs["data"] == {"token_url": "https://discovered.example.com/token"} - - @pytest.mark.asyncio - async def test_persist_discovered_obo_token_url_skips_when_not_needed(self): - """The write-back fires only for an OBO server that discovered a new endpoint: a row that - already has token_url, a non-OBO auth_type, or a discovery that found nothing all no-op.""" - manager = MCPServerManager() - update_mock = AsyncMock() - repo_instance = MagicMock() - repo_instance.table.update = update_mock - - with ( - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", - return_value=repo_instance, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - # already populated -> no write - await manager._persist_discovered_obo_token_url( - server_id="s", - auth_type=MCPAuth.oauth2_token_exchange, - existing_token_url="https://already.example.com/token", - discovered_token_url="https://new.example.com/token", - ) - # not an OBO server -> no write - await manager._persist_discovered_obo_token_url( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_token_url=None, - discovered_token_url="https://new.example.com/token", - ) - # discovery found nothing -> no write - await manager._persist_discovered_obo_token_url( - server_id="s", - auth_type=MCPAuth.oauth2_token_exchange, - existing_token_url=None, - discovered_token_url=None, - ) - - update_mock.assert_not_awaited() - - @pytest.mark.asyncio - async def test_persist_discovered_obo_token_url_is_best_effort(self): - """A write-back failure must not propagate; discovery just re-runs on the next build.""" - manager = MCPServerManager() - update_mock = AsyncMock(side_effect=Exception("db unavailable")) - repo_instance = MagicMock() - repo_instance.table.update = update_mock - - with ( - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", - return_value=repo_instance, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - await manager._persist_discovered_obo_token_url( - server_id="s", - auth_type=MCPAuth.oauth2_token_exchange, - existing_token_url=None, - discovered_token_url="https://new.example.com/token", - ) - - update_mock.assert_awaited_once() - - @pytest.mark.asyncio - async def test_build_mcp_server_from_table_persists_discovered_oauth_endpoints(self): - """A DB-backed oauth2 server with no configured endpoints discovers them and must write - authorization_url, token_url, and scopes back to the row; otherwise the resolved values - live only in memory and one failed re-discovery serves the 400 "authorization url is not configured" - from /authorize. registration_url must never be persisted because - _dcr_bridge_relays_client_registration keys off that column.""" - manager = MCPServerManager() - - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): - assert allow_origin_fallback is True - return MCPOAuthMetadata( - scopes=["mcp.read", "mcp.write"], + scopes=["mcp.read"], authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", - ) - - manager._descovery_metadata = fake_discovery # type: ignore[attr-defined] - - record = LiteLLM_MCPServerTable( - server_id="oauth-persist-1", - server_name="oauth_persist", - url="https://example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - oauth2_flow="authorization_code", - credentials={"client_id": "cid", "client_secret": "csec"}, - ) - - update_mcp_server_mock = AsyncMock() - with ( - patch( - "litellm.proxy._experimental.mcp_server.db.update_mcp_server", - new=update_mcp_server_mock, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - server = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) - - assert server.authorization_url == "https://idp.example.com/authorize" - update_mcp_server_mock.assert_awaited_once() - persisted = update_mcp_server_mock.call_args.kwargs["data"] - assert persisted.server_id == "oauth-persist-1" - assert persisted.authorization_url == "https://idp.example.com/authorize" - assert persisted.token_url == "https://idp.example.com/token" - assert persisted.credentials == {"scopes": ["mcp.read", "mcp.write"]} - assert "registration_url" not in persisted.fields_set() - assert update_mcp_server_mock.call_args.kwargs["touched_by"] == "mcp_oauth_discovery" - - @pytest.mark.asyncio - async def test_persist_discovered_oauth_endpoints_guards(self): - """The write-back must no-op for non-discovery auth types, empty discovery, origin-fallback - guesses (never harden an inferred authorization server into configuration), and rows whose - fields are all already populated.""" - manager = MCPServerManager() - advertised = MCPOAuthMetadata( - scopes=["s1"], - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", - ) - - update_mcp_server_mock = AsyncMock() - with ( - patch( - "litellm.proxy._experimental.mcp_server.db.update_mcp_server", - new=update_mcp_server_mock, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.api_key, - existing_issuer=None, - existing_authorization_url=None, - existing_token_url=None, - existing_scopes=None, - metadata=advertised, - ) - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer=None, - existing_authorization_url=None, - existing_token_url=None, - existing_scopes=None, - metadata=None, - ) - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer=None, - existing_authorization_url=None, - existing_token_url=None, - existing_scopes=None, - metadata=advertised.model_copy(update={"from_origin_fallback": True}), - ) - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer=None, - existing_authorization_url="https://configured.example.com/authorize", - existing_token_url="https://configured.example.com/token", - existing_scopes=["configured"], - metadata=advertised, - ) - - update_mcp_server_mock.assert_not_awaited() - - @pytest.mark.asyncio - async def test_persist_discovered_oauth_endpoints_only_fills_empty_fields(self): - """A row that already has token_url keeps it; only the missing authorization_url and - scopes are written, so admin-typed values always win over discovery.""" - manager = MCPServerManager() - - update_mcp_server_mock = AsyncMock() - with ( - patch( - "litellm.proxy._experimental.mcp_server.db.update_mcp_server", - new=update_mcp_server_mock, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer=None, - existing_authorization_url=None, - existing_token_url="https://configured.example.com/token", - existing_scopes=None, - metadata=MCPOAuthMetadata( - scopes=["s1"], - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", - ), - ) - - update_mcp_server_mock.assert_awaited_once() - persisted = update_mcp_server_mock.call_args.kwargs["data"] - assert persisted.authorization_url == "https://idp.example.com/authorize" - assert persisted.credentials == {"scopes": ["s1"]} - assert "token_url" not in persisted.fields_set() - - @pytest.mark.asyncio - async def test_persist_discovered_oauth_endpoints_writes_discovered_issuer_trust_on_first_use(self): - """A server with no configured issuer records the discovered issuer trust-on-first-use, so the - next rebuild anchors discovery on it (RFC 8414 §3.3) instead of re-trusting the resource. When - an issuer is already set (admin-typed or a prior discovery), it is never overwritten.""" - manager = MCPServerManager() - metadata = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", - discovered_issuer="https://idp.example.com", - ) - - update_mcp_server_mock = AsyncMock() - with ( - patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer=None, - existing_authorization_url=None, - existing_token_url=None, - existing_scopes=None, - metadata=metadata, - ) - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer="https://admin-configured.example.com", - existing_authorization_url="https://admin-configured.example.com/authorize", - existing_token_url="https://admin-configured.example.com/token", - existing_scopes=["cfg"], - metadata=metadata, - ) - - assert update_mcp_server_mock.await_count == 1 - persisted = update_mcp_server_mock.call_args.kwargs["data"] - assert persisted.issuer == "https://idp.example.com" - - @pytest.mark.asyncio - async def test_persist_discovered_oauth_endpoints_does_not_persist_endpoints_for_issuer_anchored(self): - """For an issuer-anchored server the endpoints are re-derived from the §3.3-validated issuer - document every build, so they must NOT be written into the endpoint columns: persisting them - would make the next build see populated endpoints and treat them as authoritative stored - values, defeating the issuer-only invariant. Only the resource-driven scopes are persisted.""" - manager = MCPServerManager() - metadata = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", - scopes=["read"], - ) - - update_mcp_server_mock = AsyncMock() - with ( - patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer="https://idp.example.com", - existing_authorization_url=None, - existing_token_url=None, - existing_scopes=None, - metadata=metadata, - is_issuer_anchored=True, - ) - - update_mcp_server_mock.assert_awaited_once() - persisted = update_mcp_server_mock.call_args.kwargs["data"] - assert "authorization_url" not in persisted.fields_set() - assert "token_url" not in persisted.fields_set() - assert persisted.credentials == {"scopes": ["read"]} - - @pytest.mark.asyncio - async def test_build_mcp_server_from_table_skips_persistence_for_temporary_servers(self): - """The session endpoint builds temporary servers whose server_id has no DB row; with - persist_discovered_endpoints=False neither the oauth2 nor the OBO write-back may fire.""" - manager = MCPServerManager() - - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): - return MCPOAuthMetadata( - scopes=["s1"], - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", + discovered_issuer="https://idp.example.com", ) manager._descovery_metadata = fake_discovery # type: ignore[attr-defined] update_mcp_server_mock = AsyncMock() - obo_update_mock = AsyncMock() repo_instance = MagicMock() - repo_instance.table.update = obo_update_mock + repo_instance.table.update = AsyncMock() with ( - patch( - "litellm.proxy._experimental.mcp_server.db.update_mcp_server", - new=update_mcp_server_mock, - ), + patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock), patch( "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", return_value=repo_instance, ), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), ): - oauth2_record = LiteLLM_MCPServerTable( - server_id="temp-oauth-1", - server_name="temp_oauth", - url="https://example.com/mcp", + for auth_type, flow in ((MCPAuth.oauth2, "authorization_code"), (MCPAuth.oauth2_token_exchange, None)): + record = LiteLLM_MCPServerTable( + server_id=f"no-write-{auth_type}", + server_name=f"no_write_{auth_type}", + url="https://example.com/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + oauth2_flow=flow, + credentials={"client_id": "cid", "client_secret": "csec", "audience": "aud"}, + ) + built = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) + assert built.token_url == "https://idp.example.com/token" + + update_mcp_server_mock.assert_not_awaited() + repo_instance.table.update.assert_not_awaited() + + @pytest.mark.asyncio + async def test_declared_endpoints_survive_a_failed_discovery(self): + """The reporter's configuration: explicit authorization_url/token_url/registration_url, + issuer left empty. With the gateway never stamping the issuer column, the server never turns + anchored, so the declared endpoints resolve on every build, including one whose discovery + fails entirely; /authorize keeps redirecting instead of serving the 400.""" + manager = MCPServerManager() + record = LiteLLM_MCPServerTable( + server_id="declared-1", + alias="declared", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + registration_url="https://idp.example.com/register", + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)): + built = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) + + assert built.issuer_is_anchored is False + assert built.authorization_url == "https://idp.example.com/authorize" + assert built.token_url == "https://idp.example.com/token" + assert built.registration_url == "https://idp.example.com/register" + + def test_flow_endpoints_missing_arms(self): + """The reload fast-path exemption's completeness rule. Interactive needs authorize+token, + client_credentials and OBO need token only, an OBO server with a configured exchange + endpoint never discovers and must not be sent into a rebuild loop, and non-OAuth auth types + are never unresolved.""" + assert _flow_endpoints_missing(MCPAuth.oauth2, "authorization_code", "https://idp/auth", None) is True + assert _flow_endpoints_missing(MCPAuth.oauth2, "authorization_code", None, "https://idp/token") is True + assert ( + _flow_endpoints_missing(MCPAuth.oauth2, "authorization_code", "https://idp/auth", "https://idp/token") + is False + ) + assert _flow_endpoints_missing(MCPAuth.oauth2, "client_credentials", None, "https://idp/token") is False + assert _flow_endpoints_missing(MCPAuth.oauth2, "client_credentials", None, None) is True + assert _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, None) is True + assert _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, "https://idp/token") is False + assert ( + _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, None, "https://idp/exchange") is False + ) + assert _flow_endpoints_missing(MCPAuth.api_key, None, None, None) is False + + def test_unresolved_check_uses_the_flow_judge_not_the_raw_column(self): + """A legacy row the startup backfill deliberately left unstamped (token_url plus client + credentials, no authorization_url: the ambiguous M2M shape) serves client_credentials at + request time via effective_oauth2_flow. The reload check must reach the same verdict, or the + row is classified as interactive-missing-endpoints and re-runs discovery on every reload + forever. A null-flow row without the M2M shape stays interactive and genuinely unresolved.""" + m2m_shaped = MCPServer( + server_id="null-flow-m2m", + name="null_flow_m2m", + server_name="null_flow_m2m", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow=None, + token_url="https://idp.example.com/token", + client_id="cid", + client_secret="csec", + ) + assert _oauth_endpoints_unresolved(m2m_shaped) is False + + interactive_unresolved = m2m_shaped.model_copy(update={"client_id": None, "client_secret": None}) + assert _oauth_endpoints_unresolved(interactive_unresolved) is True + + def test_dcr_bridge_relay_arm_needs_its_registration_endpoint(self): + """A dcr_bridge server with no admin-configured client can only register callers through the + upstream registration endpoint, so a partial discovery that resolved authorize and token but + not registration_endpoint leaves it silently degraded to the short-circuit arm. That counts as + unresolved so it keeps retrying. A bridge with a configured client_id uses the short-circuit + arm by design and is unaffected.""" + relay_arm = MCPServer( + server_id="bridge-partial", + name="bridge_partial", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + # dcr_bridge is only valid on the client-forwarded modes (see MCPServer.is_dcr_bridge) + auth_type=MCPAuth.oauth_delegate, + dcr_bridge=True, + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + registration_url=None, + ) + assert _oauth_endpoints_unresolved(relay_arm) is True + assert _oauth_endpoints_unresolved(relay_arm.model_copy(update={"registration_url": "https://idp/reg"})) is False + assert _oauth_endpoints_unresolved(relay_arm.model_copy(update={"client_id": "admin-client"})) is False + + def test_entra_obo_without_scopes_is_unresolved(self): + """entra_obo token exchange fails closed without a scope, and scopes can come from resource + discovery, so an entra_obo server that resolved its token endpoint but no scopes is still + unresolved for its flow. The default rfc8693 profile has no such requirement.""" + entra = MCPServer( + server_id="entra-noscope", + name="entra_noscope", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + token_exchange_profile="entra_obo", + token_url="https://idp.example.com/token", + scopes=None, + ) + assert _oauth_endpoints_unresolved(entra) is True + assert _oauth_endpoints_unresolved(entra.model_copy(update={"scopes": ["api://app/.default"]})) is False + assert _oauth_endpoints_unresolved(entra.model_copy(update={"token_exchange_profile": "rfc8693"})) is False + + def test_oauth_discovery_retry_backs_off_per_server(self): + """Without a cooldown the fast-path exemption re-runs the full discovery chain, and re-emits + the unresolved warning, on every reload forever for a server that can never resolve. Delay + doubles per consecutive failure up to the cap, a success clears the state so the next failure + starts from the base delay again, and the cooldown is per server.""" + manager = MCPServerManager() + + def unresolved(server_id): + return MCPServer( + server_id=server_id, + name=server_id, + url="https://up.example.com/mcp", transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", - credentials={"client_id": "cid", "client_secret": "csec"}, - ) - obo_record = LiteLLM_MCPServerTable( - server_id="temp-obo-1", - server_name="temp_obo", - url="https://example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2_token_exchange, - credentials={"client_id": "cid", "client_secret": "csec"}, - ) - built_oauth2 = await manager.build_mcp_server_from_table( - oauth2_record, credentials_are_encrypted=False, persist_discovered_endpoints=False - ) - await manager.build_mcp_server_from_table( - obo_record, credentials_are_encrypted=False, persist_discovered_endpoints=False ) - assert built_oauth2.authorization_url == "https://idp.example.com/authorize" - update_mcp_server_mock.assert_not_awaited() - obo_update_mock.assert_not_awaited() + assert manager._oauth_discovery_retry_due("a") is True + + manager._record_oauth_discovery_outcome(unresolved("a")) + assert manager._oauth_discovery_retry_due("a") is False + assert manager._oauth_discovery_retry_due("b") is True, "cooldown must be per server" + + failures_before, _ = manager._oauth_discovery_retry_state["a"] + manager._record_oauth_discovery_outcome(unresolved("a")) + failures_after, _ = manager._oauth_discovery_retry_state["a"] + assert failures_after == failures_before + 1 + + # An elapsed cooldown lets the retry through, and the delay grows with the failure count + manager._oauth_discovery_retry_state["a"] = (1, time.monotonic() - 31.0) + assert manager._oauth_discovery_retry_due("a") is True + manager._oauth_discovery_retry_state["a"] = (5, time.monotonic() - 31.0) + assert manager._oauth_discovery_retry_due("a") is False + + resolved = unresolved("a").model_copy( + update={ + "authorization_url": "https://idp.example.com/authorize", + "token_url": "https://idp.example.com/token", + } + ) + manager._record_oauth_discovery_outcome(resolved) + assert "a" not in manager._oauth_discovery_retry_state + assert manager._oauth_discovery_retry_due("a") is True + + @pytest.mark.asyncio + async def test_reload_fast_path_retries_unresolved_oauth_servers(self): + """A server whose discovery failed must not be pinned broken by the updated_at fast path: + the next reload rebuilds it, retrying discovery on the normal cadence instead of waiting for + an unrelated config write. A resolved server with an unchanged row still takes the fast path, + so the exemption costs nothing in the steady state.""" + manager = MCPServerManager() + stamp = datetime.now() + row = LiteLLM_MCPServerTable( + server_id="retry-1", + server_name="retry_server", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + created_at=stamp, + updated_at=stamp, + ) + + def entry(authorization_url, token_url): + return MCPServer( + server_id="retry-1", + name="retry_server", + server_name="retry_server", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + authorization_url=authorization_url, + token_url=token_url, + updated_at=stamp, + ) + + raw_row = MagicMock() + raw_row.model_dump.return_value = row.model_dump() + repo_instance = MagicMock() + repo_instance.table.find_many = AsyncMock(return_value=[raw_row]) + + async def run_reload(previous_entry): + manager.registry = {"retry-1": previous_entry} + build_mock = AsyncMock(return_value=previous_entry) + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + return_value=repo_instance, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch.object(manager, "build_mcp_server_from_table", new=build_mock), + ): + await manager.reload_servers_from_database() + return build_mock + + unresolved_build = await run_reload(entry(None, None)) + unresolved_build.assert_awaited_once() + + resolved_build = await run_reload(entry("https://idp.example.com/authorize", "https://idp.example.com/token")) + resolved_build.assert_not_awaited() + + @pytest.mark.asyncio + async def test_anchored_issuer_discarding_stored_endpoints_warns(self, caplog): + """An anchored server ignoring stored endpoint columns must say so: that state is exactly + what a row stamped by an earlier release looks like after upgrade, and the warning names the + remedy (clear the Issuer field) instead of leaving the 400 undiagnosable.""" + manager = MCPServerManager() + record = LiteLLM_MCPServerTable( + server_id="stamped-1", + alias="stamped_row", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + issuer="https://idp.example.com", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + with ( + patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=None)), + caplog.at_level(logging.WARNING, logger="LiteLLM"), + ): + built = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) + + assert built.issuer_is_anchored is True + assert built.authorization_url is None + assert "stamped_row" in caplog.text + assert "authorization_url, token_url" in caplog.text + assert "clear the Issuer" in caplog.text @pytest.mark.asyncio async def test_update_server_carries_forward_last_known_good_oauth_endpoints(self): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py new file mode 100644 index 00000000000..b6c946b95fa --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py @@ -0,0 +1,129 @@ +"""Tests for the one-time heal of issuer values a released version's discovery write-back stamped.""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy._experimental.mcp_server.oauth_issuer_stamp_backfill import ( + backfill_discovery_stamped_issuers, +) + + +def _row(**overrides): + fields = { + "server_id": "srv-1", + "alias": "srv_one", + "server_name": "srv_one", + "auth_type": "oauth2", + "issuer": "https://idp.example.com", + "authorization_url": "https://idp.example.com/authorize", + "token_url": "https://idp.example.com/token", + "registration_url": None, + "updated_by": "mcp_oauth_discovery", + } + fields.update(overrides) + return SimpleNamespace(**fields) + + +def _prisma(rows): + prisma_client = MagicMock() + prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=rows) + prisma_client.db.litellm_mcpservertable.update = AsyncMock() + return prisma_client + + +@pytest.mark.asyncio +async def test_clears_the_stamp_and_records_its_own_actor(): + """The GH #34985 row: discovery wrote the issuer, so the server reads as issuer-anchored and its + configured endpoints are ignored. Clearing the stamp makes them apply again. The heal records its + own actor, which is also what makes it idempotent: the row no longer matches the discovery-actor + filter, so it is never reconsidered on a later boot.""" + prisma_client = _prisma([_row()]) + + assert await backfill_discovery_stamped_issuers(prisma_client) == 1 + + call = prisma_client.db.litellm_mcpservertable.update.call_args + assert call.kwargs["where"] == {"server_id": "srv-1"} + assert call.kwargs["data"]["issuer"] is None + assert call.kwargs["data"]["updated_by"] == "mcp_oauth_issuer_stamp_backfill" + + where = prisma_client.db.litellm_mcpservertable.find_many.call_args.kwargs["where"] + assert where["updated_by"] == "mcp_oauth_discovery" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "overrides, reason", + [ + ({"updated_by": "some-admin@example.com"}, "an admin was the last writer, so the pin is theirs"), + ({"issuer": None}, "nothing to heal"), + ({"issuer": " "}, "blank issuer is not a pin"), + ( + {"authorization_url": None, "token_url": None, "registration_url": None}, + "issuer set with no configured endpoints is the canonical shape of a deliberate pin, and " + "there is nothing configured for anchoring to discard anyway", + ), + ( + {"authorization_url": "https://other-idp.example.com/authorize", "token_url": None}, + "endpoints addressing a different authority than the issuer are an intent a clear would " + "discard, so the row is warned about rather than healed", + ), + ( + {"issuer": "https://pinned.example.com"}, + "same shape from the other side: a pinned issuer whose origin differs from the configured " + "endpoints cannot have been derived from them by discovery", + ), + ], +) +async def test_leaves_rows_alone_that_do_not_carry_the_defect_signature(overrides, reason): + """updated_by records only the most recent writer and no audit trail says which field it touched, + so the heal is deliberately narrow: it fires only on the full signature of the defect. Every + exclusion here protects a row whose issuer may be a deliberate admin pin.""" + prisma_client = _prisma([_row(**overrides)]) + + assert await backfill_discovery_stamped_issuers(prisma_client) == 0, reason + prisma_client.db.litellm_mcpservertable.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_heals_across_url_forms_that_denote_the_same_origin(): + """Origin comparison runs through the shared canonicalizer, so a default port or host casing + difference between the stamped issuer and the endpoints an admin typed does not make a #34985 row + look like a deliberate pin at a different authority.""" + prisma_client = _prisma( + [ + _row( + issuer="https://IDP.example.com:443", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + ) + ] + ) + + assert await backfill_discovery_stamped_issuers(prisma_client) == 1 + + +@pytest.mark.asyncio +async def test_query_is_scoped_to_auth_types_where_an_issuer_anchors(): + """Only the discovery auth types read an issuer as a trust anchor; clearing it elsewhere would be + an unrelated mutation.""" + prisma_client = _prisma([]) + + await backfill_discovery_stamped_issuers(prisma_client) + + where = prisma_client.db.litellm_mcpservertable.find_many.call_args.kwargs["where"] + assert set(where["auth_type"]["in"]) == {"oauth2", "true_passthrough", "oauth_delegate"} + + +@pytest.mark.asyncio +async def test_a_failed_row_does_not_abort_the_rest(): + """Per-row best effort: one write failure must not leave later rows unhealed, and the next boot + retries the failed one since its updated_by is unchanged.""" + prisma_client = _prisma([_row(server_id="bad"), _row(server_id="good")]) + prisma_client.db.litellm_mcpservertable.update = AsyncMock( + side_effect=[Exception("write failed"), MagicMock()] + ) + + assert await backfill_discovery_stamped_issuers(prisma_client) == 1 + assert prisma_client.db.litellm_mcpservertable.update.await_count == 2 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 5c9612a055e..cec62f79e33 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -1710,6 +1710,89 @@ class TestCallToolRestAPI: assert captured["allowed_mcp_servers"] == [stub_server] fire_logging.assert_awaited_once() + async def test_returns_guardrail_rewritten_tool_result(self, monkeypatch): + """A post_mcp_call guardrail rewrite of the tool result must reach the REST caller, + not the raw result the upstream server returned.""" + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + server_id = "server-1" + alias = "server-1" + server_name = "server-1" + name = "stub" + allowed_tools = None + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + auth_type = None + + stub_server = StubServer() + + async def fake_add_litellm_data_to_request(**kwargs): + return kwargs.get("data", {}) + + async def fake_execute_mcp_tool(**kwargs): + return {"content": [{"type": "text", "text": "jane@example.com"}]} + + monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.add_litellm_data_to_request", + fake_add_litellm_data_to_request, + raising=False, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", {}, raising=False) + monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", fake_execute_mcp_tool, raising=False) + masked_result = {"content": [{"type": "text", "text": ""}]} + monkeypatch.setattr( + rest_endpoints, + "_fire_mcp_tool_call_logging", + AsyncMock(return_value=masked_result), + raising=False, + ) + + request = _build_request( + path="/mcp-rest/tools/call", + method="POST", + json_body={"server_id": "server-1", "name": "demo-tool", "arguments": {}}, + ) + + result = await rest_endpoints.call_tool_rest_api(request, user_api_key_dict=UserAPIKeyAuth()) + + assert result == masked_result + + async def test_success_logging_guardrail_rejection_propagates(self, monkeypatch): + """A guardrail rejecting the tool result must not be swallowed as a logging failure, + otherwise the unguarded result would still be returned to the caller.""" + from litellm.exceptions import BlockedPiiEntityError + + fire_logging = AsyncMock( + side_effect=BlockedPiiEntityError(entity_type="EMAIL_ADDRESS", guardrail_name="presidio-mcp") + ) + monkeypatch.setattr(rest_endpoints, "_fire_mcp_tool_call_logging", fire_logging, raising=False) + + with pytest.raises(BlockedPiiEntityError): + await rest_endpoints._safe_fire_mcp_tool_call_logging( + object(), {"result": "ok"}, datetime.now(), datetime.now() + ) + + fire_logging.assert_awaited_once() + @pytest.mark.parametrize("upstream_status", [401, 403]) async def test_call_tool_rest_relays_upstream_auth_failure(self, monkeypatch, upstream_status): """A pass-through call that hits an upstream 401/403 (surfaced by the manager as diff --git a/tests/test_litellm/proxy/analytics_endpoints/__init__.py b/tests/test_litellm/proxy/analytics_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py b/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py new file mode 100644 index 00000000000..c48b8cfd5a5 --- /dev/null +++ b/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py @@ -0,0 +1,133 @@ +""" +The cache dashboard chart is fed by /global/activity/cache_hits. Aggregation +lives server-side: the SQL groups per call_type (splitting cache hits vs +successful vs failed requests; failed spend logs have call_type '' today and +must surface as 'Unknown'), and the endpoint returns chart-ready groups, +totals for the stat cards, and the filter options for the UI dropdowns. +""" + +import json +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy.analytics_endpoints.analytics_endpoints import get_global_activity +from litellm.proxy.analytics_endpoints.cache_activity import ( + GROUPS_SQL, + CacheActivityGroup, + compute_totals, +) + +GROUP_ROWS = [ + { + "call_type": "acompletion", + "api_requests": 1000, + "cache_hits": 300, + "failed_requests": 200, + "cached_completion_tokens": 12000, + "generated_completion_tokens": 48000, + }, + { + "call_type": "Unknown", + "api_requests": 0, + "cache_hits": 0, + "failed_requests": 110, + "cached_completion_tokens": 0, + "generated_completion_tokens": 0, + }, +] +KEY_ALIAS_ROWS = [{"key_alias": "Unnamed Key"}, {"key_alias": "my-key"}] +MODEL_ROWS = [{"model": "gpt-5.1"}] + + +def build_prisma(query_raw: AsyncMock) -> MagicMock: + prisma = MagicMock() + prisma.db.query_raw = query_raw + return prisma + + +def dispatching_query_raw() -> AsyncMock: + async def dispatch(sql: str, *params: object) -> list[dict[str, object]]: + if "GROUP BY" in sql: + return GROUP_ROWS + if "key_alias" in sql: + return KEY_ALIAS_ROWS + return MODEL_ROWS + + return AsyncMock(side_effect=dispatch) + + +@pytest.fixture +def mock_prisma(monkeypatch: pytest.MonkeyPatch) -> MagicMock: + prisma = build_prisma(dispatching_query_raw()) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) + return prisma + + +@pytest.mark.asyncio +async def test_returns_groups_totals_and_filter_options(mock_prisma: MagicMock): + response = await get_global_activity( + start_date="2026-07-01", end_date="2026-07-27", key_aliases=[], models=[] + ) + + assert [group.call_type for group in response.groups] == ["acompletion", "Unknown"] + assert response.groups[0].api_requests == 1000 + assert response.groups[0].failed_requests == 200 + assert response.totals.api_requests == 1000 + assert response.totals.cache_hits == 300 + assert response.totals.failed_requests == 310 + assert response.totals.cached_completion_tokens == 12000 + assert response.totals.cache_hit_ratio == pytest.approx((300 / 1610) * 100) + assert response.filter_options.key_aliases == ["Unnamed Key", "my-key"] + assert response.filter_options.models == ["gpt-5.1"] + + +@pytest.mark.asyncio +async def test_filters_are_passed_to_sql_as_json_arrays(mock_prisma: MagicMock): + await get_global_activity( + start_date="2026-07-01", + end_date="2026-07-27", + key_aliases=["my-key"], + models=["gpt-5.1", "claude-opus-4-8"], + ) + + groups_call = next( + call for call in mock_prisma.db.query_raw.call_args_list if "GROUP BY" in call.args[0] + ) + assert groups_call.args[3] == json.dumps(["my-key"]) + assert groups_call.args[4] == json.dumps(["gpt-5.1", "claude-opus-4-8"]) + + +@pytest.mark.asyncio +async def test_rejects_malformed_dates_with_400(mock_prisma: MagicMock): + with pytest.raises(HTTPException) as exc_info: + await get_global_activity(start_date="07/01/2026", end_date="2026-07-27", key_aliases=[], models=[]) + + assert exc_info.value.status_code == 400 + mock_prisma.db.query_raw.assert_not_called() + + +def test_totals_ratio_is_zero_without_requests(): + totals = compute_totals([]) + + assert totals.cache_hit_ratio == 0.0 + assert totals.api_requests == 0 + + +def test_totals_denominator_includes_failed_requests(): + group = CacheActivityGroup( + call_type="acompletion", + api_requests=60, + cache_hits=20, + failed_requests=20, + cached_completion_tokens=0, + generated_completion_tokens=0, + ) + + assert compute_totals([group]).cache_hit_ratio == pytest.approx(20.0) + + +def test_groups_sql_splits_failures_and_labels_empty_call_type_unknown(): + assert "SUM(CASE WHEN sl.\"status\" = 'failure' THEN 1 ELSE 0 END)" in GROUPS_SQL + assert "CASE WHEN sl.\"call_type\" = '' THEN 'Unknown' ELSE sl.\"call_type\" END" in GROUPS_SQL diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index ccb20976df9..5f3b0f36b95 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -50,6 +50,7 @@ from litellm.proxy.auth.auth_checks import ( vector_store_access_check, ) from litellm.caching.in_memory_cache import InMemoryCache +from litellm.caching.redis_cache import RedisCache from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -4211,6 +4212,11 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): len(teams)==1 before populating the cache. 3. When team_alias is None, NO alias-key operation happens (no delete of an empty-keyed entry, no spurious write). + 4. DELETES the team_id-keyed entry from the internal usage cache + BEFORE the fresh write (LIT-4391). `_get_team_object_from_cache` + consults the internal usage cache first, so a leftover copy there + (backfilled from a Redis shared with `user_api_key_cache`) would + keep serving the pre-update team allowlist. """ from unittest.mock import AsyncMock, MagicMock @@ -4257,9 +4263,14 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): # (2) team_alias-keyed entry is deleted in BOTH the in-memory cache # and the Redis dual cache (mirrors _delete_cache_key_object pattern). cache.delete_cache.assert_called_once_with(key="team_alias:H-Capacity") - logging_obj.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with( - key="team_alias:H-Capacity" - ) + + # (4) internal usage cache: team_id entry deleted BEFORE the fresh + # write, alias entry deleted as before. + internal_deleted_keys = [ + (c.kwargs.get("key") or c.args[0]) + for c in logging_obj.internal_usage_cache.dual_cache.async_delete_cache.await_args_list + ] + assert internal_deleted_keys == ["team_id:team-1234", "team_alias:H-Capacity"] # ===== team_alias is None: no alias-key operation ===== aliasless = LiteLLM_TeamTableCachedObj(**{**base_team_row, "team_alias": None}) @@ -4277,7 +4288,9 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): ) cache2.delete_cache.assert_not_called() - logging_obj2.internal_usage_cache.dual_cache.async_delete_cache.assert_not_awaited() + logging_obj2.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with( + key="team_id:team-no-alias" + ) written_keys_aliasless = [ (c.kwargs.get("key") or c.args[0]) for c in cache2.async_set_cache.await_args_list @@ -4285,6 +4298,145 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): assert written_keys_aliasless == ["team_id:team-no-alias"] +class _SharedFakeRedis(RedisCache): + """Dict-backed stand-in for the single Redis that both + ``user_api_key_cache`` (enable_redis_auth_cache) and + ``proxy_logging_obj.internal_usage_cache.dual_cache`` share in the + LIT-4391 deployment topology. Only the methods DualCache calls are + implemented; ``super().__init__`` is skipped intentionally.""" + + def __init__(self): + self._store: dict = {} + + async def async_set_cache(self, key, value, **kwargs): + self._store[key] = json.dumps(value) + + async def async_get_cache(self, key, **kwargs): + raw = self._store.get(key) + return json.loads(raw) if raw is not None else None + + async def async_delete_cache(self, key): + self._store.pop(key, None) + + +@pytest.mark.asyncio +async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): + """ + Regression test for LIT-4391: keys with models=["all-team-models"] kept + getting 403 team_model_access_denied for models added via /team/update. + + `_get_team_object_from_cache` consults the internal usage cache BEFORE + `user_api_key_cache`. When both share one Redis (enable_redis_auth_cache), + any team read backfills the internal cache's in-memory tier with the team + object. `_cache_team_object` (the /team/update refresh) only wrote + `user_api_key_cache`, so that backfilled copy kept shadowing the update + until its TTL expired — and the auth-time write-back then pushed the stale + copy back into the shared Redis, making the staleness self-sustaining. + + Pins: + 1. After `_cache_team_object` writes an updated team, `get_team_object` + returns the UPDATED model list even though the internal usage cache's + in-memory tier was backfilled with the pre-update team. + 2. The shared Redis still holds the updated team afterwards — the + internal-cache invalidation must happen BEFORE the fresh write, or it + would wipe the value it just wrote. + """ + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.auth.auth_checks import _cache_team_object, get_team_object + + team_id = "team-lit-4391" + shared_redis = _SharedFakeRedis() + user_api_key_cache = UserApiKeyCache(redis_cache=shared_redis) + proxy_logging_obj = MagicMock() + proxy_logging_obj.internal_usage_cache.dual_cache = DualCache( + redis_cache=shared_redis, + default_in_memory_ttl=300, + ) + prisma_client = MagicMock() + + await _cache_team_object( + team_id=team_id, + team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a"]), + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + primed = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert primed is not None and primed.models == ["model-a"] + + await _cache_team_object( + team_id=team_id, + team_table=LiteLLM_TeamTableCachedObj( + team_id=team_id, models=["model-a", "model-b"] + ), + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + refreshed = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert refreshed is not None and refreshed.models == ["model-a", "model-b"], ( + "get_team_object served a stale team allowlist after _cache_team_object " + f"refreshed it. Got models={refreshed.models if refreshed else None}" + ) + + redis_copy = await shared_redis.async_get_cache(f"team_id:{team_id}") + assert redis_copy is not None and redis_copy["models"] == ["model-a", "model-b"], ( + "The shared Redis lost the refreshed team object — the internal-cache " + "invalidation must run BEFORE the fresh write, not after. " + f"Got: {redis_copy}" + ) + + +@pytest.mark.asyncio +async def test_cache_team_object_tolerates_cache_invalidation_failures(): + """ + Greptile review on the LIT-4391 fix: `_cache_team_object` runs after a + successful DB fetch (inside `get_team_object`) and after every team + mutation's DB write. A cache-backend error during the best-effort + invalidations must NOT fail those operations — otherwise a Redis blip + turns a healthy team lookup into a 404 and a committed /team/update into + a 500. The authoritative team_id-keyed write must still happen. + """ + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.auth.auth_checks import _cache_team_object + + cache = MagicMock() + cache.async_set_cache = AsyncMock() + cache.delete_cache = MagicMock(side_effect=Exception("redis down")) + logging_obj = MagicMock() + logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock( + side_effect=Exception("redis down") + ) + + await _cache_team_object( + team_id="team-cache-outage", + team_table=LiteLLM_TeamTableCachedObj( + team_id="team-cache-outage", + team_alias="cache-outage-alias", + models=["model-a"], + ), + user_api_key_cache=cache, + proxy_logging_obj=logging_obj, + ) + + written_keys = [ + (c.kwargs.get("key") or c.args[0]) + for c in cache.async_set_cache.await_args_list + ] + assert written_keys == ["team_id:team-cache-outage"] + + MODEL_DISCOVERY_ROUTES = [ "/v1/models", "/models", @@ -4683,66 +4835,26 @@ async def test_common_checks_personal_user_budget_blocks_in_gather(): @pytest.mark.asyncio -async def test_user_budget_enforced_on_team_key(): - """User budget must be enforced even when the key belongs to a team. +async def test_common_checks_personal_user_budget_skipped_for_team_key(): + """A user's personal max_budget does not apply to a team-scoped key. - Previously _user_max_budget_check skipped enforcement for team keys, - letting a user with a $100 personal budget spend unlimited through a - team key. This regression test ensures that is no longer the case. + Team keys are governed by the team (and team-member) budgets only; the key + owner's personal budget is deliberately out of scope. This asserts the read + path lets a team key through even when the user is far over their personal + budget, and fails if personal enforcement is reintroduced for team keys. """ from fastapi import Request from litellm.proxy.auth.auth_checks import common_checks user = LiteLLM_UserTable(user_id="u1", spend=0.0, max_budget=100.0) - team = LiteLLM_TeamTable(team_id="t1", max_budget=2100.0) + team = LiteLLM_TeamTable(team_id="t1", spend=0.0, max_budget=1000.0) token = UserAPIKeyAuth(token="k1", user_id="u1", team_id="t1") async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs): return 999.0 if counter_key == "spend:user:u1" else 0.0 - async def _no_membership(*a, **kw): - return None - - with patch("litellm.proxy.proxy_server.prisma_client", None), patch( - "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter - ), patch("litellm.proxy.auth.auth_checks.get_team_membership", _no_membership): - with pytest.raises(litellm.BudgetExceededError) as over: - await common_checks( - request_body={"messages": [{"role": "user", "content": "hi"}]}, - team_object=team, - user_object=user, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route="/chat/completions", - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=MagicMock(spec=Request), - ) - assert "User=u1" in str(over.value) - - -@pytest.mark.asyncio -async def test_skip_user_budget_on_team_key_flag_restores_old_behavior(): - """Setting skip_user_budget_on_team_key=True skips user budget for team keys. - - This is the opt-in escape hatch that restores the legacy behavior where - user budgets were not enforced when the key belonged to a team. - """ - from fastapi import Request - - from litellm.proxy.auth.auth_checks import common_checks - - user = LiteLLM_UserTable(user_id="u1", spend=0.0, max_budget=100.0) - team = LiteLLM_TeamTable(team_id="t1", max_budget=2100.0) - token = UserAPIKeyAuth(token="k1", user_id="u1", team_id="t1") - - async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs): - return 999.0 if counter_key == "spend:user:u1" else 0.0 - - async def _no_membership(*a, **kw): + async def _no_membership(*args, **kwargs): return None with patch("litellm.proxy.proxy_server.prisma_client", None), patch( @@ -4754,7 +4866,7 @@ async def test_skip_user_budget_on_team_key_flag_restores_old_behavior(): user_object=user, end_user_object=None, global_proxy_spend=None, - general_settings={"skip_user_budget_on_team_key": True}, + general_settings={}, route="/chat/completions", llm_router=None, proxy_logging_obj=MagicMock(), @@ -4762,3 +4874,362 @@ async def test_skip_user_budget_on_team_key_flag_restores_old_behavior(): request=MagicMock(spec=Request), ) assert result is True + + +@pytest.mark.parametrize( + "scope, route, expect_blocked", + [ + ("user", "/chat/completions", True), + ("user", "/key/list", False), + ("team", "/chat/completions", True), + ("team", "/key/list", False), + ("org", "/chat/completions", True), + ("org", "/key/list", False), + ], +) +@pytest.mark.asyncio +async def test_budget_checks_only_run_on_llm_api_routes(scope, route, expect_blocked): + """Budgets cap spend, so they must only gate routes that can spend. + + Enforcing them on management routes locked an over-budget caller out of the + Admin UI, which authenticates with a normal virtual key, leaving no way to + reach the page that raises the limit. + """ + from fastapi import Request + + from litellm.proxy.auth.auth_checks import common_checks + + over_budget_counter = {"user": "spend:user:u1", "team": "spend:team:t1", "org": "spend:org:o1"}[scope] + + async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs): + return 999.0 if counter_key == over_budget_counter else 0.0 + + async def _no_membership(*a, **kw): + return None + + org_table = MagicMock() + org_table.spend = 999.0 + org_table.litellm_budget_table = MagicMock() + org_table.litellm_budget_table.max_budget = 10.0 + + async def _get_org(*a, **kw): + return org_table + + user = LiteLLM_UserTable(user_id="u1", spend=0.0, max_budget=10.0 if scope == "user" else None) + team = LiteLLM_TeamTable(team_id="t1", max_budget=10.0) if scope == "team" else None + token = UserAPIKeyAuth( + token="k1", + user_id="u1", + team_id="t1" if scope == "team" else None, + org_id="o1" if scope == "org" else None, + ) + proxy_logging_obj = MagicMock() + proxy_logging_obj.budget_alerts = AsyncMock() + + async def _run(): + return await common_checks( + request_body={"messages": [{"role": "user", "content": "hi"}]}, + team_object=team, + user_object=user, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=proxy_logging_obj, + valid_token=token, + request=MagicMock(spec=Request), + ) + + with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( + "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter + ), patch("litellm.proxy.auth.auth_checks.get_team_membership", _no_membership), patch( + "litellm.proxy.auth.auth_checks.get_org_object", _get_org + ): + if expect_blocked: + with pytest.raises(litellm.BudgetExceededError): + await _run() + else: + assert await _run() is True + + +@pytest.mark.parametrize("route", ["/health", "/health/services", "/health/test_connection"]) +@pytest.mark.asyncio +async def test_spend_capable_non_llm_routes_still_enforce_budget(route): + """These routes are not LLM API routes but still reach a provider or an + external service: /health and /health/test_connection run litellm.ahealth_check + against real deployments, and /health/services fires Slack/email/webhook sends. + Exempting them with the other management routes would let an exhausted budget + keep spending. + """ + from fastapi import Request + + from litellm.proxy.auth.auth_checks import common_checks + + team = LiteLLM_TeamTable(team_id="t1", spend=150.0, max_budget=100.0) + + with pytest.raises(litellm.BudgetExceededError): + await common_checks( + request_body={}, + team_object=team, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=AsyncMock(), + valid_token=UserAPIKeyAuth(token="k1", team_id="t1"), + request=MagicMock(spec=Request), + ) + + +@pytest.mark.asyncio +async def test_get_default_end_user_budget_db_fetch_returns_validated_budget(monkeypatch): + from litellm.proxy.auth.auth_checks import get_default_end_user_budget + + monkeypatch.setattr(litellm, "max_end_user_budget_id", "budget-default-1") + + budget_row = MagicMock() + budget_row.dict = lambda: {"budget_id": "budget-default-1", "max_budget": 12.5, "tpm_limit": 100} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_default_end_user_budget( + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_BudgetTable) + assert result.max_budget == 12.5 + assert result.tpm_limit == 100 + mock_cache.async_set_cache.assert_awaited_once() + assert mock_cache.async_set_cache.call_args.kwargs["value"] is result + + +@pytest.mark.asyncio +async def test_get_end_user_object_db_fetch_returns_validated_end_user(): + from litellm.proxy.auth.auth_checks import get_end_user_object + + end_user_row = MagicMock() + end_user_row.dict = lambda: {"user_id": "eu-1", "blocked": False, "spend": 3.0} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock(return_value=end_user_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_end_user_object( + end_user_id="eu-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_EndUserTable) + assert result.user_id == "eu-1" + assert result.blocked is False + assert result.spend == 3.0 + + +@pytest.mark.asyncio +async def test_get_team_membership_db_fetch_returns_validated_membership(): + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.auth.auth_checks import get_team_membership + + membership_row = MagicMock() + membership_row.dict = lambda: {"user_id": "u-1", "team_id": "t-1", "spend": 1.5} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_team_membership( + user_id="u-1", + team_id="t-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_TeamMembership) + assert result.user_id == "u-1" + assert result.team_id == "t-1" + assert result.spend == 1.5 + + +@pytest.mark.asyncio +async def test_get_access_object_db_fetch_returns_validated_access_group(): + from litellm.proxy._types import LiteLLM_AccessGroupTable + from litellm.proxy.auth.auth_checks import get_access_object + + access_row = MagicMock() + access_row.dict = lambda: { + "access_group_id": "ag-1", + "access_group_name": "group one", + "access_model_names": ["gpt-4"], + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=access_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_access_object( + access_group_id="ag-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + proxy_logging_obj=None, + ) + + assert isinstance(result, LiteLLM_AccessGroupTable) + assert result.access_group_id == "ag-1" + assert result.access_model_names == ["gpt-4"] + + +@pytest.mark.asyncio +async def test_get_team_object_by_alias_db_fetch_returns_cached_obj(): + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.auth.auth_checks import get_team_object_by_alias + + team_row = MagicMock() + team_row.model_dump = lambda: {"team_id": "t-9", "team_alias": "alias-9", "models": ["gpt-4"]} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_team_object_by_alias( + team_alias="alias-9", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_TeamTableCachedObj) + assert result.team_id == "t-9" + assert result.team_alias == "alias-9" + assert result.models == ["gpt-4"] + + +@pytest.mark.asyncio +async def test_get_org_object_by_alias_db_fetch_returns_validated_org(): + from litellm.proxy._types import LiteLLM_OrganizationTable + from litellm.proxy.auth.auth_checks import get_org_object_by_alias + + org_row = MagicMock() + org_row.model_dump = lambda: { + "organization_id": "org-1", + "organization_alias": "org-alias", + "budget_id": "b-1", + "created_by": "admin", + "updated_by": "admin", + "models": [], + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[org_row]) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_org_object_by_alias( + org_alias="org-alias", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_OrganizationTable) + assert result.organization_id == "org-1" + assert result.budget_id == "b-1" + + +@pytest.mark.asyncio +async def test_get_object_permission_db_fetch_returns_validated_permission(): + from litellm.proxy.auth.auth_checks import get_object_permission + + perm_row = MagicMock() + perm_row.dict = lambda: {"object_permission_id": "op-1", "vector_stores": ["vs-1"]} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=perm_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_object_permission( + object_permission_id="op-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_ObjectPermissionTable) + assert result.object_permission_id == "op-1" + assert result.vector_stores == ["vs-1"] + + +@pytest.mark.asyncio +async def test_get_managed_vector_store_rows_by_uuids_db_fetch_validates_rows(): + from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable + from litellm.proxy.auth.auth_checks import get_managed_vector_store_rows_by_uuids + + vs_row = MagicMock() + vs_row.model_dump = lambda: {"vector_store_id": "vs-7", "custom_llm_provider": "openai"} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[vs_row]) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_managed_vector_store_rows_by_uuids( + uuids=["vs-7"], + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert len(result) == 1 + assert isinstance(result[0], LiteLLM_ManagedVectorStoresTable) + assert result[0].vector_store_id == "vs-7" + assert result[0].custom_llm_provider == "openai" + + +@pytest.mark.asyncio +async def test_get_project_object_db_fetch_returns_cached_obj(): + from litellm.proxy._types import LiteLLM_ProjectTableCachedObj + from litellm.proxy.auth.auth_checks import get_project_object + + project_row = MagicMock() + project_row.model_dump = lambda: {"project_id": "p-1", "project_alias": "proj", "team_id": "t-1"} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock(return_value=project_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_project_object( + project_id="p-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_ProjectTableCachedObj) + assert result.project_id == "p-1" + assert result.project_alias == "proj" diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 9f24c662581..1610d76efb7 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -569,6 +569,103 @@ def test_get_model_from_request_resolves_video_id_model_with_router(): ) +_BATCH_DEPLOYMENT_ID = "8d0eaa7e6c6f54a425dfd0062cb6b0dc" + + +def _managed_batch_router(): + from litellm.router import Router + + return Router( + model_list=[ + { + "model_name": "bedrock-batch-model", + "litellm_params": { + "model": "bedrock/global.anthropic.claude-haiku-4-5-20251001-v1:0", + }, + "model_info": {"id": _BATCH_DEPLOYMENT_ID}, + }, + { + "model_name": "some-other-model", + "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"}, + "model_info": {"id": "a-different-deployment-id"}, + }, + ] + ) + + +def _encode_managed_id(decoded: str) -> str: + return base64.urlsafe_b64encode(decoded.encode()).decode().rstrip("=") + + +_MANAGED_BATCH_ID = _encode_managed_id( + f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123" +) +_MANAGED_BATCH_OUTPUT_FILE_ID = _encode_managed_id( + f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123;" + "llm_output_file_id:provider-file-456" +) + + +@pytest.mark.parametrize( + "route, request_data", + [ + ("/v1/batches/{batch_id}", {"batch_id": _MANAGED_BATCH_ID}), + ("/v1/batches/{batch_id}/cancel", {"batch_id": _MANAGED_BATCH_ID}), + ("/v1/files/{file_id}", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}), + ("/v1/files/{file_id}/content", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}), + ], +) +def test_get_model_from_request_resolves_batch_id_deployment_to_model_name(route, request_data): + """Regression for #32580: managed batch retrieve/cancel and managed batch output + file reads encode the deployment model_id into the resource id. The auth layer must + resolve that id back to the public model group name so model-access checks compare + against the model group, not the raw deployment id.""" + assert ( + get_model_from_request( + request_data=request_data, + route=route, + llm_router=_managed_batch_router(), + ) + == "bedrock-batch-model" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "route, request_data", + [ + ("/v1/batches/{batch_id}", {"batch_id": _MANAGED_BATCH_ID}), + ("/v1/batches/{batch_id}/cancel", {"batch_id": _MANAGED_BATCH_ID}), + ("/v1/files/{file_id}/content", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}), + ], +) +async def test_managed_batch_routes_pass_team_model_access_check(route, request_data): + """End-to-end regression for #32580: a team scoped to the batch model group got + ``team_model_access_denied`` on retrieve/cancel because the deployment id, not the + model group, was authorized. Fails pre-fix with the deployment id in the message.""" + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.proxy.auth.auth_checks import can_team_access_model + + llm_router = _managed_batch_router() + model = get_model_from_request(request_data=request_data, route=route, llm_router=llm_router) + + assert ( + await can_team_access_model( + model=model, + team_object=LiteLLM_TeamTable(team_id="team-batch", models=["bedrock-batch-model"]), + llm_router=llm_router, + ) + is True + ) + + with pytest.raises(Exception, match="team not allowed to access model"): + await can_team_access_model( + model=model, + team_object=LiteLLM_TeamTable(team_id="team-other", models=["some-other-model"]), + llm_router=llm_router, + ) + + def test_get_model_from_request_resolves_character_id_model_with_router(): from litellm.types.videos.utils import encode_character_id_with_provider diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index ffc5241d027..3840c90d691 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -1240,6 +1240,87 @@ async def test_find_team_with_model_access_model_group(monkeypatch): assert team_obj.team_id == "team-1" +@pytest.mark.asyncio +async def test_find_team_with_model_access_v1_messages_default_routes(monkeypatch): + """Regression for #31189: a single-team JWT that grants the requested model + through an access group must resolve on /v1/messages without an explicit + x-litellm-team-id header. /v1/messages lives in `anthropic_routes`, so when a + team has no `team_allowed_routes` configured the default allowlist must cover + it just like /chat/completions and /v1/responses; otherwise the internal route + check fails and surfaces a misleading "No team has access to the requested + model" 403.""" + from litellm.caching import DualCache + from litellm.proxy.utils import ProxyLogging + from litellm.router import Router + + router = Router( + model_list=[ + { + "model_name": "claude-sonnet-4-6", + "litellm_params": {"model": "claude-sonnet-4-6"}, + "model_info": {"access_groups": ["coding_only_models"]}, + } + ] + ) + import sys + import types + + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + team = LiteLLM_TeamTable(team_id="coding-team", models=["coding_only_models"]) + + async def mock_get_team_object(*args, **kwargs): # type: ignore + return team + + monkeypatch.setattr( + "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object + ) + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() + + user_api_key_cache = DualCache() + proxy_logging_obj = ProxyLogging(user_api_key_cache=user_api_key_cache) + + team_id, team_obj = await JWTAuthManager.find_team_with_model_access( + team_ids={"coding-team"}, + requested_model="claude-sonnet-4-6", + route="/v1/messages", + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + + assert team_id == "coding-team" + assert team_obj.team_id == "coding-team" + + +@pytest.mark.parametrize( + "route,expected", + [ + ("/v1/messages", True), + ("/v1/messages/count_tokens", True), + ("/v1/skills", False), + ("/v1/skills/skill_abc123", False), + ], +) +def test_default_team_allowed_routes_cover_messages_but_not_skills(route, expected): + from litellm.proxy.auth.auth_checks import allowed_routes_check + + assert ( + allowed_routes_check( + user_role=LitellmUserRoles.TEAM, + user_route=route, + litellm_proxy_roles=LiteLLM_JWTAuth(), + ) + is expected + ) + + @pytest.mark.asyncio async def test_auth_builder_returns_team_membership_object(): """ diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index ca2e282a119..affaaa3fbf4 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -2918,6 +2918,123 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): setattr(_proxy_server_mod, k, v) +@pytest.mark.asyncio +async def test_auth_flow_never_persists_fallback_team_object_lit_4391(): + """ + Regression test for LIT-4391 (stale team allowlist poisoning). + + When `get_team_object` fails at the "Check 6" team-auth step (cache miss + inside the DB-throttle window, DB blip, ...), the builder falls back to a + team object reconstructed from the CACHED token's team_* snapshot — which + can be arbitrarily stale (e.g. pre-/team/update models). + + The builder used to write that team object back into `user_api_key_cache` + under "team_id:" after Check 6. Writing a cache-read (or worse, a + token-snapshot) value back into the shared cache re-poisons it — with + enable_redis_auth_cache it clobbered the fresh team `/team/update` had + just written to Redis, making the stale allowlist self-sustaining across + requests. Only authoritative writers (`_cache_team_object` on DB reads and + team mutations) may populate the team cache. + + Pins: the auth flow completes on the fallback path WITHOUT writing any + "team_id:*" cache entry. + """ + from starlette.datastructures import URL + from starlette.requests import Request + from fastapi import HTTPException + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + api_key = "sk-test-lit-4391-no-team-writeback" + valid_token = UserAPIKeyAuth( + api_key=api_key, + token=api_key, + user_role=LitellmUserRoles.INTERNAL_USER, + team_id="team-lit-4391", + team_models=["model-a"], + models=["all-team-models"], + ) + + mock_cache = AsyncMock() + mock_cache.async_get_cache = AsyncMock(return_value=valid_token) + mock_cache.async_set_cache = AsyncMock(return_value=None) + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs = { + "prisma_client": MagicMock(), + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _originals = {k: getattr(_proxy_server_mod, k, None) for k in _attrs} + + try: + for k, v in _attrs.items(): + setattr(_proxy_server_mod, k, v) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + with ( + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new_callable=AsyncMock, + return_value=valid_token, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + side_effect=HTTPException( + status_code=404, + detail={"error": "Team doesn't exist in db."}, + ), + ), + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + assert result.team_id == "team-lit-4391" + + team_cache_writes = [ + key + for c in mock_cache.async_set_cache.await_args_list + if isinstance(key := (c.kwargs.get("key") if "key" in c.kwargs else c.args[0]), str) + and key.startswith("team_id:") + ] + assert team_cache_writes == [], ( + "The auth flow wrote a team object into the cache. Fallback/" + "cache-read team objects must never be persisted — only " + "_cache_team_object (DB reads and team mutations) may write " + f"'team_id:*' entries. Got writes: {team_cache_writes}" + ) + + finally: + for k, v in _originals.items(): + setattr(_proxy_server_mod, k, v) + + # --------------------------------------------------------------------------- # _run_centralized_common_checks — centralized authz gate @@ -4517,101 +4634,6 @@ async def test_real_jwt_still_requires_license_when_jwt_auth_enabled(monkeypatch assert "enterprise only feature" in message -@pytest.mark.asyncio -async def test_auth_path_caches_team_object_under_canonical_team_id_key(): - """Regression for LIT-4000: the auth builder must cache the team object under - the canonical ``team_id:{id}`` key that ``get_team_object`` and - ``_update_team_cache`` read, never under the raw ``team_id`` (and never under - a ``None`` key, which Redis rejects with a NoneType key error). A raw or None - key is silently dropped by Redis / never served back, so every request - re-hits Postgres for the team object instead of the L2 cache. - - Drives the real builder for a team-scoped key against a real in-memory - ``UserApiKeyCache`` and reads the team object back. Mutating the cache key at - the write site to the raw ``valid_token.team_id`` (or ``None``) makes the - canonical-key read miss and fails this test. - """ - from fastapi import Request - from starlette.datastructures import URL - - import litellm.proxy.proxy_server as _proxy_server_mod - from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder - from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache - from litellm.proxy.proxy_server import hash_token - - team_id = "team-lit-4000" - api_key = "sk-lit-4000-team-key" - cache = UserApiKeyCache() - - team_token = UserAPIKeyAuth(token=hash_token(api_key), team_id=team_id) - team_obj = LiteLLM_TeamTableCachedObj(team_id=team_id) - - proxy_logging_obj = MagicMock() - proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) - attrs = { - "prisma_client": MagicMock(), - "user_api_key_cache": cache, - "proxy_logging_obj": proxy_logging_obj, - "master_key": "sk-test-master", - "general_settings": {"allow_requests_on_db_unavailable": False}, - "llm_model_list": [], - "llm_router": None, - "open_telemetry_logger": None, - "model_max_budget_limiter": MagicMock(), - "user_custom_auth": None, - "jwt_handler": None, - "litellm_proxy_admin_name": "admin", - } - originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} - try: - for k, v in attrs.items(): - setattr(_proxy_server_mod, k, v) - request = Request(scope={"type": "http"}) - request._url = URL(url="/chat/completions") - with ( - patch( - "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", - AsyncMock(return_value=team_token), - ), - patch( - "litellm.proxy.auth.user_api_key_auth.get_team_object", - AsyncMock(return_value=team_obj), - ), - patch( - "litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access", - new_callable=AsyncMock, - ), - patch( - "litellm.proxy.auth.user_api_key_auth._return_user_api_key_auth_obj", - new_callable=AsyncMock, - return_value=team_token, - ), - patch( - "litellm.proxy.auth.auth_exception_handler.seed_request_identity", - ), - ): - await _user_api_key_auth_builder( - request=request, - api_key=f"Bearer {api_key}", - azure_api_key_header="", - anthropic_api_key_header=None, - google_ai_studio_api_key_header=None, - azure_apim_header=None, - request_data={}, - ) - finally: - for k, v in originals.items(): - setattr(_proxy_server_mod, k, v) - - served = cache.get_cache( - key=f"team_id:{team_id}", model_type=LiteLLM_TeamTableCachedObj - ) - assert served is not None and served.team_id == team_id - assert cache.get_cache(key=team_id) is None - assert cache.get_cache(key=None) is None - - @pytest.mark.asyncio async def test_auth_does_not_rewrite_cached_key_object_back_into_cache(): """A cache-hit auth must not write the token back into the cache. diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 9573fddd435..6a185988c9b 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -51,7 +51,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( from litellm.proxy.utils import ProxyLogging from litellm.router import Router from litellm.types.llms.openai import BatchJobStatus -from litellm.types.utils import LiteLLMBatch +from litellm.types.utils import CredentialItem, LiteLLMBatch from fastapi import Response @@ -2091,3 +2091,154 @@ async def test_retrieve__unified_no_router_500(retrieve_harness): assert exc.value.code == "500" retrieve_harness.router_aretrieve.assert_not_called() retrieve_harness.litellm_aretrieve.assert_not_called() + + +# =========================================================================== # +# SCENARIO 3 + configured deployments: a provider-only call (custom-llm-provider +# header, no model anywhere) must resolve the gateway/team deployment's named +# credential for that provider and attach it to the provider call kwargs, +# instead of silently falling through to the host environment's default +# credentials (regression: vertex batch jobs landing in the hosting env's GCP +# project because litellm_credential_name never reached the call). +# =========================================================================== # + +VERTEX_NAMED_CREDENTIAL = CredentialItem( + credential_name="vertex-named-cred", + credential_info={}, + credential_values={ + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + }, +) + + +def vertex_named_credential_router() -> Router: + return Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "litellm_credential_name": "vertex-named-cred", + }, + } + ] + ) + + +@pytest.mark.asyncio +async def test_create__provider_only_resolves_named_vertex_credentials(harness): + """Provider-only create must attach the configured named credential, and must + NOT turn the call into a model-routed one (no model kwarg injected).""" + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + harness.provider_from_headers.return_value = "vertex_ai" + + with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]): + with patch.object(proxy_server, "llm_router", vertex_named_credential_router()): + await call_create(harness) + + assert harness.acreate_kwargs() == { + "custom_llm_provider": "vertex_ai", + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "metadata": None, + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + } + + +@pytest.mark.asyncio +async def test_create__provider_only_ignores_other_provider_deployments(harness): + """A provider-only vertex call must not pick up credentials from deployments + of a different provider; with no vertex deployment the payload is exactly the + pre-fix env-var fallback.""" + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + harness.provider_from_headers.return_value = "vertex_ai" + openai_only_router = Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-global-openai"}, + } + ] + ) + + with patch.object(proxy_server, "llm_router", openai_only_router): + await call_create(harness) + + assert harness.acreate_kwargs() == { + "custom_llm_provider": "vertex_ai", + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "metadata": None, + } + + +@pytest.mark.asyncio +async def test_retrieve__provider_only_resolves_named_vertex_credentials(retrieve_harness): + retrieve_harness.provider_from_headers.return_value = "vertex_ai" + + with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]): + with patch.object(proxy_server, "llm_router", vertex_named_credential_router()): + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.aretrieve_kwargs() == { + "custom_llm_provider": "vertex_ai", + "batch_id": "batch-raw-xyz", + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + } + + +@pytest.mark.asyncio +async def test_list__provider_only_resolves_named_vertex_credentials(list_harness): + list_harness.provider_from_headers.return_value = "vertex_ai" + + with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]): + with patch.object(proxy_server, "llm_router", vertex_named_credential_router()): + await call_list(list_harness) + + assert list_harness.alist_kwargs() == { + "custom_llm_provider": "vertex_ai", + "after": None, + "limit": None, + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + } + + +@pytest.mark.asyncio +async def test_cancel__provider_only_resolves_named_vertex_credentials(cancel_harness): + cancel_harness.provider_from_headers.return_value = "vertex_ai" + + with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]): + with patch.object(proxy_server, "llm_router", vertex_named_credential_router()): + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.acancel_kwargs() == { + "custom_llm_provider": "vertex_ai", + "batch_id": "batch-raw-xyz", + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + } diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py index 2fbc9c5c82f..f0aa49ff123 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -1,5 +1,6 @@ import json import os +import stat import sys import time from pathlib import Path @@ -12,6 +13,7 @@ import pytest from click.testing import CliRunner from litellm.constants import CLI_JWT_EXPIRATION_HOURS +from litellm.proxy.client.cli import cli from litellm.proxy.client.cli.commands.auth import ( clear_token, get_stored_api_key, @@ -201,31 +203,22 @@ class TestTokenUtilities: mock_mkdir.assert_called_once_with(exist_ok=True) - def test_save_token(self): + def test_save_token(self, tmp_path): """Test saving token data to file""" token_data = { "key": "test-key", "user_id": "test-user", "timestamp": 1234567890, } + token_file = tmp_path / "token.json" - with ( - patch("builtins.open", mock_open()) as mock_file, - patch("litellm.proxy.client.cli.commands.auth.get_token_file_path") as mock_path, - patch("os.chmod") as mock_chmod, - ): - mock_path.return_value = "/test/path/token.json" + with patch("litellm.proxy.client.cli.commands.auth.get_token_file_path") as mock_path: + mock_path.return_value = str(token_file) save_token(token_data) - mock_file.assert_called_once_with("/test/path/token.json", "w") - mock_file().write.assert_called() - mock_chmod.assert_called_once_with("/test/path/token.json", 0o600) - - # Verify JSON content was written correctly - written_content = "".join(call[0][0] for call in mock_file().write.call_args_list) - parsed_content = json.loads(written_content) - assert parsed_content == token_data + assert json.loads(token_file.read_text()) == token_data + assert stat.S_IMODE(token_file.stat().st_mode) == 0o600 def test_load_token_success(self): """Test loading token data from file successfully""" @@ -808,7 +801,8 @@ class TestPrintTokenCommand: since there is no explicit target to check it against. `--base-url`/ `LITELLM_PROXY_URL` only enforces the match when a caller explicitly passes it (tracked via ctx.obj["base_url_explicit"], set by the `cli` - group from click's ParameterSource). + group from click's ParameterSource); a base_url saved via + `lite config set` counts as explicit too. """ def setup_method(self): @@ -928,3 +922,110 @@ class TestPrintTokenCommand: assert "sk-stale-key" not in result.output assert "lite login" in result.output mock_post.assert_not_called() + + +def _write_home_json(home: Path, filename: str, payload: dict[str, object]) -> None: + litellm_dir = home / ".litellm" + litellm_dir.mkdir(exist_ok=True) + (litellm_dir / filename).write_text(json.dumps(payload)) + + +class TestPrintTokenWithConfigFile: + """A config-file base_url is a drop-in replacement for exporting + LITELLM_PROXY_URL, so print-token must treat it as an explicit server + choice: a token minted for a different proxy is never handed out.""" + + @pytest.fixture + def isolated_home(self, monkeypatch, tmp_path): + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + monkeypatch.delenv("LITELLM_PROXY_URL", raising=False) + monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False) + return tmp_path + + def test_config_base_url_mismatch_fails_closed(self, isolated_home): + _write_home_json( + isolated_home, + "token.json", + {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()}, + ) + _write_home_json(isolated_home, "config.json", {"base_url": "https://server-b.example.com"}) + + result = CliRunner().invoke(cli, ["auth", "print-token"]) + + assert result.exit_code == 1 + assert "sk-issued-for-a" not in result.output + assert "Not authenticated for this server" in result.output + + def test_config_base_url_match_prints_token(self, isolated_home): + _write_home_json( + isolated_home, + "token.json", + {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()}, + ) + _write_home_json(isolated_home, "config.json", {"base_url": "https://server-a.example.com"}) + + result = CliRunner().invoke(cli, ["auth", "print-token"]) + + assert result.exit_code == 0 + assert result.stdout.strip() == "sk-issued-for-a" + + def test_empty_config_base_url_treated_as_unset(self, isolated_home): + """A hand-edited config.json with base_url "" must behave like no config at all: + base_url falls back to the default AND explicitness stays False.""" + _write_home_json( + isolated_home, + "token.json", + {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()}, + ) + _write_home_json(isolated_home, "config.json", {"base_url": ""}) + + result = CliRunner().invoke(cli, ["auth", "print-token"]) + + assert result.exit_code == 0 + assert result.stdout.strip() == "sk-issued-for-a" + + def test_bare_invocation_without_config_file_unchanged(self, isolated_home): + """No config file means base_url_explicit stays False, so the stored + token's own server is trusted (pre-config behavior must not regress).""" + _write_home_json( + isolated_home, + "token.json", + {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()}, + ) + + result = CliRunner().invoke(cli, ["auth", "print-token"]) + + assert result.exit_code == 0 + assert result.stdout.strip() == "sk-issued-for-a" + + +class TestSaveTokenPrivateWrite: + """token.json holds the real API key: it must never be world-readable at any + instant, and a failed write must not destroy the previously stored token.""" + + @pytest.fixture + def isolated_home(self, monkeypatch, tmp_path): + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + monkeypatch.delenv("LITELLM_PROXY_URL", raising=False) + monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False) + return tmp_path + + def test_save_token_owner_only_permissions_and_no_temp_leftovers(self, isolated_home): + save_token({"key": "sk-secret", "user_id": "u-1", "timestamp": 1234567890}) + + token_file = isolated_home / ".litellm" / "token.json" + assert json.loads(token_file.read_text()) == {"key": "sk-secret", "user_id": "u-1", "timestamp": 1234567890} + assert stat.S_IMODE(token_file.stat().st_mode) == 0o600 + assert list(token_file.parent.glob(".tmp-*")) == [] + + def test_save_token_failure_mid_write_preserves_existing_token(self, isolated_home): + _write_home_json(isolated_home, "token.json", {"key": "sk-original", "timestamp": 1234567890}) + token_file = isolated_home / ".litellm" / "token.json" + + with pytest.raises(TypeError): + save_token({"key": object()}) + + assert json.loads(token_file.read_text()) == {"key": "sk-original", "timestamp": 1234567890} + assert list(token_file.parent.glob(".tmp-*")) == [] diff --git a/tests/test_litellm/proxy/client/cli/test_config_commands.py b/tests/test_litellm/proxy/client/cli/test_config_commands.py new file mode 100644 index 00000000000..698d6188768 --- /dev/null +++ b/tests/test_litellm/proxy/client/cli/test_config_commands.py @@ -0,0 +1,284 @@ +import json +import os +import stat +import sys +from pathlib import Path + +import pytest +from click.testing import CliRunner + +sys.path.insert(0, os.path.abspath("../../..")) + + +from litellm.proxy.client.cli import cli +from litellm.proxy.client.cli.commands.config import ( + get_config_file_path, + get_config_value, + load_config, + save_config, +) +from litellm.proxy.client.cli.commands.private_json import write_private_json + + +@pytest.fixture +def cli_runner(): + return CliRunner() + + +@pytest.fixture +def isolated_home(monkeypatch, tmp_path): + """Point HOME at tmp_path so tests never touch the developer's real ~/.litellm.""" + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + monkeypatch.delenv("LITELLM_PROXY_URL", raising=False) + monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False) + return tmp_path + + +def _config_path(home: Path) -> Path: + return home / ".litellm" / "config.json" + + +def _raise_home_unresolvable() -> str: + raise RuntimeError("Could not determine home directory.") + + +class TestConfigSet: + @pytest.mark.parametrize( + "value", + ["https://your-proxy.example.com", "http://your-proxy.example.com:8080"], + ) + def test_set_stores_value_with_owner_only_permissions(self, cli_runner, isolated_home, value): + result = cli_runner.invoke(cli, ["config", "set", "base_url", value]) + + assert result.exit_code == 0 + config_file = _config_path(isolated_home) + assert json.loads(config_file.read_text()) == {"base_url": value} + assert stat.S_IMODE(config_file.stat().st_mode) == 0o600 + assert str(config_file) in result.output + + def test_set_strips_trailing_slash(self, cli_runner, isolated_home): + """Downstream commands join paths onto base_url; a stored trailing + slash would produce double slashes in every request URL.""" + result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com/"]) + + assert result.exit_code == 0 + assert json.loads(_config_path(isolated_home).read_text()) == {"base_url": "https://your-proxy.example.com"} + + def test_set_unknown_key_rejected_and_names_allowed_keys(self, cli_runner, isolated_home): + result = cli_runner.invoke(cli, ["config", "set", "api_key", "sk-secret"]) + + assert result.exit_code != 0 + assert "base_url" in result.output + assert not _config_path(isolated_home).exists() + + @pytest.mark.parametrize("value", ["your-proxy.example.com", "ftp://your-proxy.example.com"]) + def test_set_base_url_without_http_scheme_rejected(self, cli_runner, isolated_home, value): + result = cli_runner.invoke(cli, ["config", "set", "base_url", value]) + + assert result.exit_code != 0 + assert "http" in result.output + assert not _config_path(isolated_home).exists() + + @pytest.mark.parametrize("value", ["https://", "http://", "https:///some-path"]) + def test_set_base_url_without_host_rejected(self, cli_runner, isolated_home, value): + """rstrip("/") would otherwise persist a bare "https:" that breaks every later request.""" + result = cli_runner.invoke(cli, ["config", "set", "base_url", value]) + + assert result.exit_code != 0 + assert not _config_path(isolated_home).exists() + + @pytest.mark.parametrize( + "value", + [ + "https://proxy.example.com?env=prod", + "https://proxy.example.com#prod", + "https://proxy.example.com/?", + "https://proxy.example.com/#", + ], + ) + def test_set_base_url_with_query_or_fragment_rejected(self, cli_runner, isolated_home, value): + """Downstream commands join paths onto base_url; a stored query string or + fragment would silently corrupt every request URL built from it. Bare + trailing '?' / '#' parse as EMPTY query/fragment yet still break every + joined path, so rejection must key off the raw characters.""" + result = cli_runner.invoke(cli, ["config", "set", "base_url", value]) + + assert result.exit_code != 0 + assert "query" in result.output or "fragment" in result.output + assert not _config_path(isolated_home).exists() + + def test_set_base_url_with_path_prefix_accepted(self, cli_runner, isolated_home): + """Proxies are commonly served under a path prefix; the query/fragment + rejection must not over-reach into legitimate paths.""" + result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://proxy.example.com/litellm"]) + + assert result.exit_code == 0 + assert json.loads(_config_path(isolated_home).read_text()) == {"base_url": "https://proxy.example.com/litellm"} + + def test_set_leaves_no_temp_files_behind(self, cli_runner, isolated_home): + """The atomic write goes through a .tmp-* sibling; it must be renamed away, + never abandoned next to the config.""" + result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"]) + + assert result.exit_code == 0 + config_file = _config_path(isolated_home) + assert stat.S_IMODE(config_file.stat().st_mode) == 0o600 + assert list(config_file.parent.glob(".tmp-*")) == [] + + +class TestConfigGet: + def test_get_prints_only_the_value(self, cli_runner, isolated_home): + """stdout must be exactly the value so scripts can do URL=$(lite config get base_url).""" + set_result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"]) + assert set_result.exit_code == 0 + + result = cli_runner.invoke(cli, ["config", "get", "base_url"]) + + assert result.exit_code == 0 + assert result.stdout.strip() == "https://your-proxy.example.com" + + def test_get_unset_key_exits_one_with_stderr_message(self, cli_runner, isolated_home): + result = cli_runner.invoke(cli, ["config", "get", "base_url"]) + + assert result.exit_code == 1 + assert result.stdout.strip() == "" + assert result.stderr != "" + + def test_get_without_key_lists_entries(self, cli_runner, isolated_home): + set_result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"]) + assert set_result.exit_code == 0 + + result = cli_runner.invoke(cli, ["config", "get"]) + + assert result.exit_code == 0 + assert "base_url = https://your-proxy.example.com" in result.output + + def test_get_without_key_when_nothing_set(self, cli_runner, isolated_home): + result = cli_runner.invoke(cli, ["config", "get"]) + + assert result.exit_code == 0 + assert "no config" in result.output.lower() + + +class TestConfigUnset: + def test_unset_removes_key_from_file(self, cli_runner, isolated_home): + set_result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"]) + assert set_result.exit_code == 0 + + result = cli_runner.invoke(cli, ["config", "unset", "base_url"]) + + assert result.exit_code == 0 + assert "base_url" not in load_config() + assert cli_runner.invoke(cli, ["config", "get", "base_url"]).exit_code == 1 + + def test_unset_missing_key_is_idempotent(self, cli_runner, isolated_home): + result = cli_runner.invoke(cli, ["config", "unset", "base_url"]) + + assert result.exit_code == 0 + assert "not set" in result.output.lower() + + +class TestConfigHelpers: + def test_get_config_file_path_under_home(self, isolated_home): + assert get_config_file_path() == str(isolated_home / ".litellm" / "config.json") + + def test_load_config_missing_file_returns_empty(self, isolated_home): + assert load_config() == {} + + def test_home_unresolvable_does_not_crash_cli(self, cli_runner, isolated_home, monkeypatch): + """Path.home() raises RuntimeError in HOME-less containers; invocations that + never needed the home dir (--api-key supplied) must keep working.""" + monkeypatch.setattr( + "litellm.proxy.client.cli.commands.config.get_config_file_path", + _raise_home_unresolvable, + ) + + assert load_config() == {} + + result = cli_runner.invoke(cli, ["--api-key", "sk-test", "config", "get"]) + assert result.exit_code == 0 + assert "(no config set)" in result.output + + @pytest.mark.parametrize( + "content", + [ + "{not json", + '{"base_url": 123}', + '["https://your-proxy.example.com"]', + '"https://your-proxy.example.com"', + ], + ) + def test_load_config_invalid_content_returns_empty(self, isolated_home, content): + """A corrupt or wrongly-shaped config file must degrade to defaults, never crash the CLI.""" + config_file = _config_path(isolated_home) + config_file.parent.mkdir(parents=True, exist_ok=True) + config_file.write_text(content) + + assert load_config() == {} + + def test_load_config_invalid_utf8_returns_empty(self, isolated_home): + """json.load raises UnicodeDecodeError (a ValueError but not a JSONDecodeError) + on undecodable bytes; before catching ValueError this crashed every CLI + invocation, including the `config set` needed to repair the file.""" + config_file = _config_path(isolated_home) + config_file.parent.mkdir(parents=True, exist_ok=True) + config_file.write_bytes(b"\xff\xfe{}") + + assert load_config() == {} + + def test_save_config_round_trip_creates_dir_and_restricts_permissions(self, isolated_home): + save_config({"base_url": "https://your-proxy.example.com"}) + + assert load_config() == {"base_url": "https://your-proxy.example.com"} + assert stat.S_IMODE(_config_path(isolated_home).stat().st_mode) == 0o600 + + def test_get_config_value_unset_then_set(self, isolated_home): + assert get_config_value("base_url") is None + + save_config({"base_url": "https://your-proxy.example.com"}) + + assert get_config_value("base_url") == "https://your-proxy.example.com" + + def test_corrupt_config_file_warns_on_stderr_but_command_succeeds(self, cli_runner, isolated_home): + """Silently ignoring a broken config file leaves users debugging why their + stored base_url stopped applying; the CLI must keep working but say why.""" + config_file = _config_path(isolated_home) + config_file.parent.mkdir(parents=True, exist_ok=True) + config_file.write_text("{not json") + + result = cli_runner.invoke(cli, ["config", "get"]) + + assert result.exit_code == 0 + assert "Warning: ignoring invalid config file" in result.stderr + + +class TestWritePrivateJson: + def test_failed_write_preserves_previous_file_and_removes_temp(self, tmp_path): + """json.dump can fail partway through serializing; writing to a temp file + and renaming keeps the previous file intact through a crash mid-write.""" + target = tmp_path / "config.json" + original = '{"base_url": "https://original.example.com"}' + target.write_text(original) + + with pytest.raises(TypeError): + write_private_json(str(target), {"bad": object()}) + + assert target.read_text() == original + assert list(tmp_path.glob(".tmp-*")) == [] + + def test_interrupted_write_removes_temp_file(self, tmp_path, monkeypatch): + """Ctrl-C is BaseException, which `except Exception` misses; an interrupt + mid-write must not abandon a .tmp-* file next to the config forever.""" + + def _interrupt(*args: object, **kwargs: object) -> None: + raise KeyboardInterrupt() + + monkeypatch.setattr("litellm.proxy.client.cli.commands.private_json.json.dump", _interrupt) + target = tmp_path / "config.json" + + with pytest.raises(KeyboardInterrupt): + write_private_json(str(target), {"base_url": "https://your-proxy.example.com"}) + + assert not target.exists() + assert list(tmp_path.glob(".tmp-*")) == [] diff --git a/tests/test_litellm/proxy/client/cli/test_global_options.py b/tests/test_litellm/proxy/client/cli/test_global_options.py index 8df763d35c2..9995cb1bca5 100644 --- a/tests/test_litellm/proxy/client/cli/test_global_options.py +++ b/tests/test_litellm/proxy/client/cli/test_global_options.py @@ -1,4 +1,5 @@ # stdlib imports +import json import os import sys from pathlib import Path @@ -7,9 +8,7 @@ from unittest.mock import Mock, patch import pytest from click.testing import CliRunner -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path import litellm.proxy.client.cli @@ -71,13 +70,9 @@ def test_base_url_trailing_slash_normalized(cli_runner): ) as mock_post, patch("requests.get", side_effect=ValueError("stop after start request")), ): - cli_runner.invoke( - cli, ["--base-url", "https://gateway.litellm-sandbox.ai/", "login"] - ) + cli_runner.invoke(cli, ["--base-url", "https://gateway.litellm-sandbox.ai/", "login"]) - mock_post.assert_called_once_with( - "https://gateway.litellm-sandbox.ai/sso/cli/start", timeout=10 - ) + mock_post.assert_called_once_with("https://gateway.litellm-sandbox.ai/sso/cli/start", timeout=10) def test_cli_version_command(cli_runner): @@ -94,3 +89,152 @@ def test_cli_version_command(cli_runner): assert f"LiteLLM Proxy CLI Version: {litellm_version}" in result.output assert "LiteLLM Proxy Server URL: http://localhost:4000" in result.output assert "LiteLLM Proxy Server Version: 1.2.3" in result.output + + +@pytest.fixture +def isolated_home(monkeypatch, tmp_path): + """Point HOME at tmp_path so tests never touch the developer's real ~/.litellm.""" + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + monkeypatch.delenv("LITELLM_PROXY_URL", raising=False) + monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False) + return tmp_path + + +def _write_config_file(home: Path, config: dict[str, str]) -> None: + config_dir = home / ".litellm" + config_dir.mkdir(exist_ok=True) + (config_dir / "config.json").write_text(json.dumps(config)) + + +def _invoke_version(cli_runner: CliRunner, *args: str): + with patch( + "litellm.proxy.client.health.HealthManagementClient.get_server_version", + return_value="1.2.3", + ): + return cli_runner.invoke(cli, [*args, "version"]) + + +def test_base_url_read_from_config_file(cli_runner, isolated_home): + """base_url precedence: flag > env > config file > default.""" + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + + result = _invoke_version(cli_runner) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: https://config-proxy.example.com" in result.output + + +def test_env_var_beats_config_file_base_url(cli_runner, isolated_home, monkeypatch): + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + monkeypatch.setenv("LITELLM_PROXY_URL", "http://env-proxy.example.com:5000") + + result = _invoke_version(cli_runner) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: http://env-proxy.example.com:5000" in result.output + + +def test_base_url_flag_beats_env_var_and_config_file(cli_runner, isolated_home, monkeypatch): + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + monkeypatch.setenv("LITELLM_PROXY_URL", "http://env-proxy.example.com:5000") + + result = _invoke_version(cli_runner, "--base-url", "http://flag-proxy.example.com:9000") + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: http://flag-proxy.example.com:9000" in result.output + + +def test_default_base_url_unchanged_without_config_file(cli_runner, isolated_home): + result = _invoke_version(cli_runner) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: http://localhost:4000" in result.output + + +def test_corrupt_config_file_falls_back_to_default(cli_runner, isolated_home): + """A corrupt config file must never crash the CLI. Exactly one warning proves + the config file is read once per invocation, not once per lookup.""" + config_dir = isolated_home / ".litellm" + config_dir.mkdir(exist_ok=True) + (config_dir / "config.json").write_text("{not json") + + result = _invoke_version(cli_runner) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: http://localhost:4000" in result.output + assert result.stderr.count("Warning: ignoring invalid config file") == 1 + + +def test_empty_base_url_flag_is_not_treated_as_unset(cli_runner, isolated_home): + """`--base-url ""` explicitly provided an (empty) value; falling back to the + config file or localhost would silently redirect auth-sensitive commands.""" + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + + result = _invoke_version(cli_runner, "--base-url", "") + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL:" not in result.output + + +def test_version_flag_reads_config_file_base_url(cli_runner, isolated_home): + """--version resolves through the same precedence chain as every other command.""" + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + + with patch( + "litellm.proxy.client.health.HealthManagementClient.get_server_version", + return_value="1.2.3", + ): + result = cli_runner.invoke(cli, ["--version"]) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: https://config-proxy.example.com" in result.output + + +def test_version_flag_prefers_env_var_over_config_file(cli_runner, isolated_home, monkeypatch): + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + monkeypatch.setenv("LITELLM_PROXY_URL", "http://env-proxy.example.com:5000") + + with patch( + "litellm.proxy.client.health.HealthManagementClient.get_server_version", + return_value="1.2.3", + ): + result = cli_runner.invoke(cli, ["--version"]) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: http://env-proxy.example.com:5000" in result.output + + +def test_version_flag_prefers_explicit_base_url_over_config_file(cli_runner, isolated_home): + """An eager --version could not see the flag and silently queried the config + server instead of the one the user named.""" + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + + with patch( + "litellm.proxy.client.health.HealthManagementClient.get_server_version", + return_value="1.2.3", + ): + result = cli_runner.invoke(cli, ["--base-url", "https://flag-proxy.example.com", "--version"]) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: https://flag-proxy.example.com" in result.output + assert "config-proxy.example.com" not in result.output + + +def test_version_flag_never_sends_api_key_to_unnamed_server(cli_runner, isolated_home, monkeypatch): + """The version request carries a bearer token; it must reach only the server the + user named, never whichever host happens to sit in the config file.""" + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-intended-for-flag-proxy") + + with patch("litellm.proxy.client.http_client.requests.request") as mock_request: + mock_request.return_value.json.return_value = {"litellm_version": "1.2.3"} + mock_request.return_value.raise_for_status.return_value = None + result = cli_runner.invoke(cli, ["--base-url", "https://flag-proxy.example.com", "--version"]) + + assert result.exit_code == 0 + requested_urls = [call.kwargs["url"] for call in mock_request.call_args_list] + assert requested_urls + assert all(url.startswith("https://flag-proxy.example.com") for url in requested_urls) + sent_keys = [call.kwargs["headers"].get("Authorization") for call in mock_request.call_args_list] + assert sent_keys == ["Bearer sk-intended-for-flag-proxy"] * len(requested_urls) diff --git a/tests/test_litellm/proxy/db/test_prisma_client.py b/tests/test_litellm/proxy/db/test_prisma_client.py index eeaf726941f..08b873dfc44 100644 --- a/tests/test_litellm/proxy/db/test_prisma_client.py +++ b/tests/test_litellm/proxy/db/test_prisma_client.py @@ -193,3 +193,25 @@ async def test_recreate_prisma_client_recovers_from_disconnected_client( mock_kill.assert_not_called() assert wrapper._original_prisma is mock_new_prisma mock_new_prisma.connect.assert_awaited_once() + + +def test_db_push_applies_replica_identity_full_when_requested(monkeypatch): + """`prisma db push` bypasses litellm-proxy-extras, so it needs its own call + into the opt-in REPLICA IDENTITY FULL step.""" + from litellm.proxy.db.prisma_client import PrismaManager + from litellm_proxy_extras.replica_identity import REPLICA_IDENTITY_FULL_ENV_VAR + from litellm_proxy_extras.utils import ProxyExtrasDBManager + + monkeypatch.setenv(REPLICA_IDENTITY_FULL_ENV_VAR, "true") + applied = [] + monkeypatch.setattr( + ProxyExtrasDBManager, + "apply_replica_identity_full_if_requested", + staticmethod(lambda: applied.append(True)), + ) + + with patch("litellm.proxy.db.prisma_client.subprocess.run") as mock_run: + assert PrismaManager.setup_database(use_migrate=False) is True + + assert mock_run.call_args[0][0][:3] == ["prisma", "db", "push"] + assert applied == [True] diff --git a/tests/test_litellm/proxy/db/test_replica_identity.py b/tests/test_litellm/proxy/db/test_replica_identity.py new file mode 100644 index 00000000000..ecfc6433ab1 --- /dev/null +++ b/tests/test_litellm/proxy/db/test_replica_identity.py @@ -0,0 +1,85 @@ +"""The opt-in REPLICA IDENTITY FULL step, without a database. + +The behavior against real Postgres is covered by +tests/proxy_migration_tests/test_replica_identity_full.py; these pin the two +things that hold with no database at all: the statement handed to the Prisma +CLI, and the promise that no failure of this optional step escapes into a +migration run that already succeeded. +""" + +import subprocess +from pathlib import Path +from unittest.mock import patch + +import pytest + +from litellm_proxy_extras.replica_identity import ( + REPLICA_IDENTITY_FULL_ENV_VAR, + apply_replica_identity_full, +) +from litellm_proxy_extras.utils import ProxyExtrasDBManager + + +def test_hands_the_alter_statement_to_the_prisma_cli(): + captured = {} + + def capture(cmd, **kwargs): + captured["cmd"] = cmd + captured["sql"] = Path(cmd[cmd.index("--file") + 1]).read_text() + return subprocess.CompletedProcess(cmd, 0) + + with patch( + "litellm_proxy_extras.replica_identity.subprocess.run", side_effect=capture + ): + applied = apply_replica_identity_full( + schema_path="/somewhere/schema.prisma", + prisma_command="prisma", + prisma_env={"DATABASE_URL": "postgresql://x/y"}, + ) + + assert applied is True + assert captured["cmd"][:3] == ["prisma", "db", "execute"] + assert captured["cmd"][-2:] == ["--schema", "/somewhere/schema.prisma"] + + sql = captured["sql"] + assert "ALTER TABLE %s REPLICA IDENTITY FULL" in sql + assert r"c.relname LIKE 'LiteLLM\_%'" in sql + assert "c.relreplident <> 'f'" in sql + assert "lock_timeout" in sql + + +@pytest.mark.parametrize( + "failure", + [ + subprocess.CalledProcessError(1, "prisma", stderr="must be owner of table"), + subprocess.TimeoutExpired("prisma", 60), + OSError(2, "No such file or directory"), + PermissionError(13, "Read-only file system"), + ], + ids=["rejected", "timed-out", "cli-missing", "read-only-fs"], +) +def test_every_failure_is_reported_instead_of_raised(failure): + with patch( + "litellm_proxy_extras.replica_identity.subprocess.run", side_effect=failure + ): + assert ( + apply_replica_identity_full( + schema_path="/somewhere/schema.prisma", + prisma_command="prisma", + prisma_env={}, + ) + is False + ) + + +def test_an_unusable_migrations_dir_skips_the_step_instead_of_killing_the_run( + tmp_path, monkeypatch +): + """LITELLM_MIGRATION_DIR makes the step copy the migrations tree before it + can run, and that copy is filesystem work that can fail on its own.""" + blocker = tmp_path / "blocker" + blocker.write_text("not a directory") + monkeypatch.setenv(REPLICA_IDENTITY_FULL_ENV_VAR, "true") + monkeypatch.setenv("LITELLM_MIGRATION_DIR", str(blocker / "migrations")) + + assert ProxyExtrasDBManager.apply_replica_identity_full_if_requested() is False diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 9337050b61c..a4c42ff601e 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -3165,6 +3165,89 @@ async def test_pre_call_hook_does_not_leak_internal_stash_to_request_body(): assert isinstance(metadata.get(RATE_LIMIT_DESCRIPTORS_KEY), list) +@pytest.mark.asyncio +@pytest.mark.parametrize("caller_metadata", [None, {"user_tag": "abc"}]) +async def test_pre_call_hook_does_not_touch_provider_metadata_on_litellm_metadata_routes( + caller_metadata, +): + """Regression for #35197: routes that own ``litellm_metadata`` (Responses, + /v1/messages, batches, files) send ``metadata`` to the provider, so the + limiter must never create it or write stash keys into it.""" + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _LITELLM_STASH_KEYS, + RATE_LIMIT_DESCRIPTORS_KEY, + RATE_LIMIT_RESPONSE_KEY, + TPM_RESERVED_TOKENS_KEY, + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-responses-metadata"), + tpm_limit=1000, + rpm_limit=5, + ) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache), + ) + + async def mock_should_rate_limit(descriptors, **kwargs): + return { + "overall_code": "OK", + "statuses": [ + { + "code": "OK", + "current_limit": 5, + "limit_remaining": 4, + "descriptor_key": d["key"], + "descriptor_value": d["value"], + "rate_limit_type": "requests", + } + for d in descriptors + ], + } + + async def mock_reserve_tpm_tokens(descriptors, estimated_tokens, **kwargs): + return {"overall_code": "OK", "statuses": []} + + handler.should_rate_limit = mock_should_rate_limit + handler.reserve_tpm_tokens = mock_reserve_tpm_tokens + + data: Dict[str, Any] = { + "model": "responses-model", + "input": "hello", + "litellm_metadata": {}, + } + if caller_metadata is not None: + data["metadata"] = dict(caller_metadata) + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data=data, + call_type="aresponses", + ) + + if caller_metadata is None: + assert "metadata" not in data, f"limiter created provider metadata: {data.get('metadata')!r}" + else: + assert data["metadata"] == caller_metadata + + litellm_metadata = data["litellm_metadata"] + assert litellm_metadata.get(TPM_RESERVED_TOKENS_KEY) + assert isinstance(litellm_metadata.get(RATE_LIMIT_DESCRIPTORS_KEY), list) + assert litellm_metadata.get(RATE_LIMIT_RESPONSE_KEY) + + leaked = [k for k in _LITELLM_STASH_KEYS if k in data] + assert not leaked, f"stash keys leaked to top level: {leaked}" + + for key in _LITELLM_STASH_KEYS: + assert handler._lookup_stashed_value( + kwargs={"litellm_params": {"litellm_metadata": litellm_metadata}}, + standard_logging_metadata=None, + key=key, + ) == litellm_metadata.get(key) + + @pytest.mark.asyncio async def test_pre_call_hook_rejects_caller_supplied_stash_values(): """Caller cannot pre-populate stash keys in body metadata to drive a diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py index 458c7c42eb6..6970e34f759 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py @@ -332,6 +332,21 @@ class TestScimTransformations: assert scim_group.members[1].value == "test2@example.com" assert scim_group.members[1].display == "test2@example.com" + @pytest.mark.asyncio + async def test_transform_team_marks_members_as_users( + self, mock_team, mock_prisma_client + ): + """A LiteLLM team only holds users, and stating the member type keeps the + response from emitting a null ``type`` now that SCIMMember carries one.""" + mock_client, _ = mock_prisma_client + + with patch("litellm.proxy.proxy_server.prisma_client", mock_client): + scim_group = await ScimTransformations.transform_litellm_team_to_scim_group( + mock_team + ) + + assert [member.type for member in scim_group.members] == ["User", "User"] + def test_get_scim_user_name(self, mock_user, mock_user_minimal): # User with email result = ScimTransformations._get_scim_user_name(mock_user) diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py index 94ca0dc11f5..6ced5264267 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py @@ -108,6 +108,19 @@ class TestGetSchemas: assert "displayName" in attr_names assert "members" in attr_names + def test_group_schema_advertises_member_type(self): + """IdPs read the schema to learn we understand ``members.type``, which is how + a nested group announces itself.""" + schemas = _get_schemas() + group_schema = next( + s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:Group" + ) + members = next(a for a in group_schema.attributes if a.name == "members") + member_type = next(a for a in members.subAttributes or [] if a.name == "type") + assert member_type.type == "string" + assert member_type.multiValued is False + assert "Group" in (member_type.description or "") + def test_schema_meta_fields(self): schemas = _get_schemas() user_schema = next( diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index 7bb74285ac6..e333bf1e3fe 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -20,8 +20,10 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import ( _extract_group_member_ids, _extract_ids_from_path_filter, _handle_team_membership_changes, + _parse_member_entries, _process_group_patch_operations, _recompute_scim_member_roles, + _resolve_group_member_ids, create_group, create_user, delete_group, @@ -36,6 +38,8 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import ( ) from litellm.types.proxy.management_endpoints.scim_v2 import ( SCIM_ENTERPRISE_USER_SCHEMA, + SCIM_MANAGED_TEAM_METADATA_KEY, + SCIM_TEAM_DATA_METADATA_KEY, SCIMGroup, SCIMMember, SCIMPatchOp, @@ -1611,7 +1615,10 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch): mock_prisma_client.db.litellm_usertable = mocker.MagicMock() # Mock team operations - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) + def mock_team_lookup(where): + return mock_existing_team if where["team_id"] == group_id else None + + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=mock_team_lookup) # Mock updated team response mock_updated_team = mocker.MagicMock() @@ -1775,6 +1782,7 @@ async def test_extract_group_member_ids_with_flag_true_creates_users(mocker, mon mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) # Mock user lookup - only existing-user exists initially def mock_user_lookup(where): @@ -1842,6 +1850,7 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) # Mock user lookup - only existing-user exists def mock_user_lookup(where): @@ -1902,6 +1911,7 @@ async def test_process_group_patch_operations_with_flag_true_creates_users(mocke # Mock user lookup - new-user-1 doesn't exist mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) # Mock user creation created_user = NewUserResponse(user_id="new-user-1", key="test-key-1") @@ -1956,6 +1966,7 @@ async def test_process_group_patch_operations_with_flag_false_rejects(mocker, mo # Mock user lookup - new-user-1 doesn't exist mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) # Execute the function - should raise HTTPException with pytest.raises(HTTPException) as exc_info: @@ -3519,3 +3530,874 @@ async def test_process_group_patch_replace_empty_value_does_not_use_path_filter( ) assert final_members == set() + + +def _member_resolution_prisma(mocker, *, users: set, teams: set, unmanaged_teams: frozenset = frozenset()): + """Prisma mock where only the given ids resolve to a user row / team row. + + ``teams`` are teams a SCIM group write created, so they carry provenance; + ``unmanaged_teams`` resolve too but look like a team an admin created here. + """ + + def team_row(team_id: str) -> LiteLLM_TeamTable | None: + if team_id in teams: + return LiteLLM_TeamTable(team_id=team_id, metadata={SCIM_MANAGED_TEAM_METADATA_KEY: True}) + if team_id in unmanaged_teams: + return LiteLLM_TeamTable(team_id=team_id, metadata={}) + return None + + prisma_client = mocker.MagicMock() + prisma_client.db = mocker.MagicMock() + prisma_client.db.litellm_usertable = mocker.MagicMock() + prisma_client.db.litellm_usertable.find_unique = AsyncMock( + side_effect=lambda where: ( + LiteLLM_UserTable(user_id=where["user_id"]) if where["user_id"] in users else None + ) + ) + prisma_client.db.litellm_teamtable = mocker.MagicMock() + prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=lambda where: team_row(where["team_id"])) + return prisma_client + + +@pytest.fixture +def scim_upsert_user_enabled(monkeypatch): + from litellm.proxy.proxy_server import proxy_config + + async def mock_get_config(): + return {"litellm_settings": {"scim_upsert_user": True}} + + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) + + +@pytest.fixture +def scim_upsert_user_disabled(monkeypatch): + from litellm.proxy.proxy_server import proxy_config + + async def mock_get_config(): + return {"litellm_settings": {"scim_upsert_user": False}} + + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) + + +@pytest.mark.asyncio +async def test_create_group_ignores_nested_group_members(mocker, scim_upsert_user_enabled): + """Entra sends nested groups as members with ``type: "Group"``. Treating that + GUID as a user id provisioned a phantom internal user per nested group.""" + nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id="parent-group", + displayName="Parent Group", + members=[ + SCIMMember(value="real-user", display="Real User", type="User"), + SCIMMember(value=nested_group_id, display="Nested Group", type="Group"), + ], + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=_member_resolution_prisma(mocker, users={"real-user"}, teams=set())), + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + new_team_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.new_team", + AsyncMock(return_value=mocker.MagicMock()), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=scim_group), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + + await create_group(group=scim_group) + + create_user_mock.assert_not_called() + assert new_team_mock.call_args.kwargs["data"].members_with_roles == [Member(user_id="real-user", role="user")] + + +@pytest.mark.asyncio +async def test_update_group_ignores_nested_group_members(mocker, scim_upsert_user_enabled): + """PUT /Groups must drop nested-group members too, so a full sync from the IdP + neither provisions nor enrolls the nested group's GUID.""" + group_id = "parent-group" + nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested" + existing_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Parent Group", + members=[], + members_with_roles=[], + metadata={}, + ) + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group_id, + displayName="Parent Group", + members=[ + SCIMMember(value="real-user", display="Real User", type="User"), + SCIMMember(value=nested_group_id, display="Nested Group", type="Group"), + ], + ) + + prisma_client = _member_resolution_prisma(mocker, users={"real-user"}, teams={group_id}) + prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=existing_team) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=prisma_client), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._check_team_exists", + AsyncMock(return_value=existing_team), + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + patch_membership_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", + AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=scim_group), + ) + + await update_group(group_id=group_id, group=scim_group) + + create_user_mock.assert_not_called() + enrolled = {call.kwargs["user_id"] for call in patch_membership_mock.call_args_list} + assert enrolled == {"real-user"} + + +@pytest.mark.asyncio +async def test_process_group_patch_operations_ignores_nested_group_members(mocker, scim_upsert_user_enabled): + """PATCH bodies bypass SCIMGroup parsing, so ``type`` must be read off the raw + member dicts; otherwise a nested group is indistinguishable from a user id.""" + nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[ + SCIMPatchOperation( + op="add", + path="members", + value=[ + {"value": "real-user", "display": "Real User", "type": "User"}, + {"value": nested_group_id, "display": "Nested Group", "type": "Group"}, + ], + ) + ], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[Member(user_id="incumbent", role="user")], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"real-user"}, teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == {"incumbent", "real-user"} + + +@pytest.mark.asyncio +async def test_process_group_patch_operations_ignores_lowercase_group_type(mocker, scim_upsert_user_enabled): + """The ``type`` comparison is case-insensitive; IdPs are not consistent about it.""" + nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[ + SCIMPatchOperation(op="add", path="members", value=[{"value": nested_group_id, "type": "group"}]) + ], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == set() + + +@pytest.mark.asyncio +async def test_process_group_patch_operations_skips_member_matching_existing_team(mocker, scim_upsert_user_enabled): + """Okta sends filtered paths and untyped ids, so a nested group arrives with no + ``type`` at all; an id that names an existing team is still not a user.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "child-team"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[Member(user_id="incumbent", role="user")], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users=set(), teams={"child-team", "parent-group"}), + ) + + create_user_mock.assert_not_called() + assert final_members == {"incumbent"} + + +@pytest.mark.asyncio +async def test_process_group_patch_operations_prefers_user_over_team_for_colliding_id( + mocker, scim_upsert_user_enabled +): + """Nothing stops a user id from also being a team id, so the user lookup has to + win; ordering the team check first would silently stop syncing that user.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "dual-id"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"dual-id"}, teams={"dual-id"}), + ) + + assert final_members == {"dual-id"} + + +@pytest.mark.asyncio +async def test_create_group_strict_mode_accepts_group_and_team_members(mocker, scim_upsert_user_disabled): + """Strict mode (scim_upsert_user=False) rejects unknown *users*; a nested group + is not a user, so it must be dropped rather than 400 the whole sync.""" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id="parent-group", + displayName="Parent Group", + members=[ + SCIMMember(value="real-user", type="User"), + SCIMMember(value="nested-group-guid", type="Group"), + SCIMMember(value="child-team"), + ], + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=_member_resolution_prisma(mocker, users={"real-user"}, teams={"child-team"})), + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + new_team_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.new_team", + AsyncMock(return_value=mocker.MagicMock()), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=scim_group), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + + await create_group(group=scim_group) + + create_user_mock.assert_not_called() + assert new_team_mock.call_args.kwargs["data"].members_with_roles == [Member(user_id="real-user", role="user")] + + +@pytest.mark.asyncio +async def test_create_group_strict_mode_still_rejects_unknown_user(mocker, scim_upsert_user_disabled): + """The strict-mode 400 must name the unknown *user* and stay quiet about the + nested group sharing the request.""" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id="parent-group", + displayName="Parent Group", + members=[ + SCIMMember(value="nested-group-guid", type="Group"), + SCIMMember(value="unknown-user"), + ], + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=_member_resolution_prisma(mocker, users=set(), teams=set())), + ) + + with pytest.raises(ProxyException) as exc_info: + await create_group(group=scim_group) + + assert int(exc_info.value.code) == 400 + assert "unknown-user" in str(exc_info.value.message) + assert "nested-group-guid" not in str(exc_info.value.message) + + +@pytest.mark.asyncio +async def test_process_group_patch_remove_unknown_member_does_not_create_user(mocker, scim_upsert_user_enabled): + """A ``remove`` of an id we don't know is an idempotent no-op. Upserting the id + first, only to drop it from the roster, made removals a phantom-user factory.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="remove", path="members", value=[{"value": "long-gone"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[Member(user_id="keep-user", role="user")], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == {"keep-user"} + + +@pytest.mark.parametrize( + "operation", + [ + SCIMPatchOperation(op="remove", path='members[value eq "long-gone"]', value=None), + SCIMPatchOperation(op="remove", path="members", value=[{"value": "long-gone"}]), + SCIMPatchOperation(op="remove", path="members", value=[{"value": "long-gone", "type": "Group"}]), + ], + ids=["path-filter", "unknown-id", "nested-group"], +) +@pytest.mark.asyncio +async def test_process_group_patch_remove_unknown_member_does_not_reject_in_strict_mode( + mocker, scim_upsert_user_disabled, operation +): + """Strict mode must not 400 a removal: refusing to drop an id the IdP already + forgot leaves the roster permanently out of sync.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[operation], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[Member(user_id="keep-user", role="user")], + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()), + ) + + assert final_members == {"keep-user"} + + +@pytest.mark.asyncio +async def test_process_group_patch_remove_drops_member_without_user_row(mocker, scim_upsert_user_enabled): + """Phantom members already on a roster (their user row is gone) must still be + removable, so the removal id is honoured even though it resolves to nothing.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="remove", path="members", value=[{"value": "phantom"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[ + Member(user_id="keep-user", role="user"), + Member(user_id="phantom", role="user"), + ], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == {"keep-user"} + + +_NESTED_GROUP_ID = "8f1e9d70-0000-4a0e-9a1e-nested" + + +@pytest.mark.parametrize( + "member_entry, user_rows, team_rows", + [ + ({"value": _NESTED_GROUP_ID, "type": "Group"}, {"keep-user", _NESTED_GROUP_ID}, set()), + ({"value": _NESTED_GROUP_ID, "type": "Group"}, {"keep-user"}, set()), + ({"value": _NESTED_GROUP_ID, "type": "Group"}, {"keep-user"}, {_NESTED_GROUP_ID}), + ({"value": _NESTED_GROUP_ID}, {"keep-user"}, {_NESTED_GROUP_ID}), + ], + ids=["phantom-user-row-exists", "user-row-already-deleted", "child-group-is-a-team", "untyped-team-id"], +) +@pytest.mark.asyncio +async def test_process_group_patch_remove_discards_non_user_member( + mocker, scim_upsert_user_enabled, member_entry, user_rows, team_rows +): + """Rosters written before nested groups were understood still carry those ids, + and the IdP removes them exactly as it added them; a removal that resolved its + ids first would classify them as non-users and leave them stuck on the team.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="remove", path="members", value=[member_entry])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[ + Member(user_id="keep-user", role="user"), + Member(user_id=_NESTED_GROUP_ID, role="user"), + ], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users=user_rows, teams=team_rows), + ) + + create_user_mock.assert_not_called() + assert final_members == {"keep-user"} + + +@pytest.mark.asyncio +async def test_process_group_patch_add_keeps_member_typed_user_that_collides_with_team_id( + mocker, scim_upsert_user_enabled +): + """The team lookup only exists to catch nested groups that arrive untyped. An id + the IdP calls a User is a user, and IdP ids collide with team ids easily.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "123456", "type": "User"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="123456", key="new-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users=set(), teams={"123456", "parent-group"}), + ) + + assert create_user_mock.call_args.kwargs["user_id"] == "123456" + assert final_members == {"123456"} + + +@pytest.mark.parametrize("member_type", ["Device", " group ", "Machine"]) +@pytest.mark.asyncio +async def test_process_group_patch_operations_skips_non_user_member_types( + mocker, scim_upsert_user_enabled, member_type +): + """A team holds users, so a member that declares itself to be anything else is + dropped; enumerating the types worth skipping would leave the next one to leak.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "not-a-user", "type": member_type}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == set() + + +@pytest.mark.parametrize( + "team_metadata, expect_provisioned", + [ + ({SCIM_MANAGED_TEAM_METADATA_KEY: True}, False), + ({SCIM_TEAM_DATA_METADATA_KEY: {"displayName": "Child.Apps"}}, False), + ({}, True), + (None, True), + ({SCIM_MANAGED_TEAM_METADATA_KEY: False}, True), + ({SCIM_TEAM_DATA_METADATA_KEY: None}, True), + ], + ids=[ + "scim-managed", + "legacy-scim-data", + "admin-created", + "no-metadata", + "marker-unset", + "legacy-key-without-value", + ], +) +@pytest.mark.asyncio +async def test_process_group_patch_team_match_needs_scim_provenance( + mocker, scim_upsert_user_enabled, team_metadata, expect_provisioned +): + """A bare member id that names a team is only evidence of a nested group when the + identity provider is what wrote that team. Teams created here can share an id with + a real user, and skipping those members stops provisioning them entirely.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "child-team"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + prisma_client = _member_resolution_prisma(mocker, users=set(), teams=set()) + prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id="child-team", metadata=team_metadata) + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="child-team", key="new-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=prisma_client, + ) + + assert create_user_mock.called is expect_provisioned + assert final_members == ({"child-team"} if expect_provisioned else set()) + + +@pytest.mark.asyncio +async def test_create_group_strict_mode_rejects_id_matching_admin_created_team(mocker, scim_upsert_user_disabled): + """Strict mode drops nested groups but reports unknown users. A team an admin + created here says nothing about the member, so the member is an unknown user.""" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id="parent-group", + displayName="Parent Group", + members=[SCIMMember(value="admin-team")], + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock( + return_value=_member_resolution_prisma( + mocker, users=set(), teams=set(), unmanaged_teams=frozenset({"admin-team"}) + ) + ), + ) + + with pytest.raises(ProxyException) as exc_info: + await create_group(group=scim_group) + + assert int(exc_info.value.code) == 400 + assert "admin-team" in str(exc_info.value.message) + + +@pytest.mark.asyncio +async def test_create_group_stamps_scim_provenance(mocker, scim_upsert_user_enabled): + """The provenance the classifier reads only exists if the group writes stamp it; + a SCIM-created team that carries no mark looks admin-created forever after.""" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id="child-group", + displayName="Child.Apps", + members=[], + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=_member_resolution_prisma(mocker, users=set(), teams=set())), + ) + new_team_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.new_team", + AsyncMock(return_value=mocker.MagicMock()), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=scim_group), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + + await create_group(group=scim_group) + + assert new_team_mock.call_args.kwargs["data"].metadata == {SCIM_MANAGED_TEAM_METADATA_KEY: True} + + +@pytest.mark.asyncio +async def test_update_group_stamps_scim_provenance(mocker, scim_upsert_user_enabled): + """A PUT full sync adopts a team the identity provider now owns, and the stamp has + to land alongside the existing metadata rather than replacing it.""" + import json + + group_id = "child-group" + existing_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Child.Apps", + members=[], + members_with_roles=[], + metadata={"existing_key": "kept"}, + ) + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group_id, + displayName="Child.Apps", + members=[], + ) + + prisma_client = _member_resolution_prisma(mocker, users=set(), teams=set()) + prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=existing_team) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=prisma_client), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._check_team_exists", + AsyncMock(return_value=existing_team), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=scim_group), + ) + + await update_group(group_id=group_id, group=scim_group) + + written = json.loads(prisma_client.db.litellm_teamtable.update.call_args.kwargs["data"]["metadata"]) + assert written[SCIM_MANAGED_TEAM_METADATA_KEY] is True + assert written["existing_key"] == "kept" + assert SCIM_TEAM_DATA_METADATA_KEY in written + + +@pytest.mark.asyncio +async def test_process_group_patch_stamps_scim_provenance(mocker, scim_upsert_user_enabled): + """PATCH is how Okta adopts a group, so a membership-only patch has to stamp the + team too; otherwise the group it manages never gains provenance.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "real-user"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + metadata={"existing_key": "kept"}, + ) + + update_data, _, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"real-user"}, teams=set()), + ) + + assert update_data["metadata"][SCIM_MANAGED_TEAM_METADATA_KEY] is True + assert update_data["metadata"]["existing_key"] == "kept" + + +@pytest.mark.parametrize("member_type", ["direct", "Device"]) +@pytest.mark.asyncio +async def test_process_group_patch_keeps_existing_user_with_unrecognized_type( + mocker, scim_upsert_user_enabled, member_type +): + """Clients do stamp non-canonical types on real members (RFC 7643 defines + ``direct`` for ``User.groups``). Dropping a member whose id is a live user would + revoke that user's team access on the next full sync, so the type is only + grounds for skipping once the user lookup has missed.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "real-user", "type": member_type}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"real-user"}, teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == {"real-user"} + + +@pytest.mark.parametrize( + "second_creation", + [None, NewUserResponse(user_id="dup-user", key="second-key")], + ids=["second-creation-fails", "both-creations-succeed"], +) +@pytest.mark.asyncio +async def test_resolve_group_member_ids_dedupes_repeated_member(mocker, scim_upsert_user_enabled, second_creation): + """An id the request lists twice is one member. Admitting it twice writes a + duplicate members_with_roles row, and the second creation of the same id fails + against the real unique constraint even when the first one succeeded.""" + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(side_effect=[NewUserResponse(user_id="dup-user", key="first-key"), second_creation]), + ) + + result = await _resolve_group_member_ids( + members=[SCIMMember(value="dup-user"), SCIMMember(value="dup-user")], + created_via="scim_group_membership", + prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()), + ) + + assert result.all_member_ids == ["dup-user"] + + +@pytest.mark.parametrize( + "operation", + [ + SCIMPatchOperation(op="add", path="members", value=[{"value": " "}]), + SCIMPatchOperation(op="remove", path="members", value=[{"value": " "}]), + SCIMPatchOperation(op="remove", path='members[value eq " "]', value=None), + ], + ids=["add", "remove", "remove-path-filter"], +) +@pytest.mark.asyncio +async def test_process_group_patch_rejects_blank_member_id(mocker, scim_upsert_user_enabled, operation): + """A blank id names nobody. The removal path stopped resolving its members, so it + has to keep rejecting one on its own.""" + patch_ops = SCIMPatchOp(schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], Operations=[operation]) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[Member(user_id="keep-user", role="user")], + ) + + with pytest.raises(HTTPException) as exc_info: + await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()), + ) + + assert exc_info.value.status_code == 400 + + +def test_scim_member_round_trips_type(): + """``type`` has to survive parsing; dropping it is what made a nested group + look like a user id.""" + assert SCIMMember.model_validate({"value": "x", "type": "Group"}).type == "Group" + assert SCIMMember(value="x").type is None + + +@pytest.mark.parametrize("junk_type", [123, True, {}, [], 1.5]) +def test_scim_member_treats_non_string_type_as_absent(junk_type): + """Before ``type`` was a field, junk in it was parsed away; typing the field must + not start rejecting those requests, and both parsers have to agree it is typeless.""" + assert SCIMMember.model_validate({"value": "x", "type": junk_type}).type is None + assert _parse_member_entries([{"value": "x", "type": junk_type}])[0].type is None + + +@pytest.mark.asyncio +async def test_get_groups_members_are_typed_as_users(mocker): + """Group members we report back are always users, and saying so keeps the + response from emitting a null ``type``.""" + team = LiteLLM_TeamTable( + team_id="team-1", + team_alias="Team One", + members=[], + members_with_roles=[Member(user_id="member-1", role="user")], + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team]) + mock_prisma_client.db.litellm_teamtable.count = AsyncMock(return_value=1) + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=mocker.MagicMock(user_id="member-1", user_email="member-1@example.com") + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + + response = await get_groups(startIndex=1, count=10, filter=None) + + assert [m.type for m in response.Resources[0].members] == ["User"] diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py index 8eed696e77d..3240ad20edb 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py @@ -160,7 +160,9 @@ async def test_create_access_group_with_model_names_tags_all_deployments(): { "model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o", "api_key": "fake-key"}, + "model_info": {"id": deployment_id, "db_model": True}, } + for deployment_id in ("deploy-A", "deploy-B", "deploy-C") ] ) @@ -318,3 +320,113 @@ async def test_create_access_group_invalid_model_id_returns_400(): await create_model_group(data=request_data, user_api_key_dict=mock_user) assert exc_info.value.status_code == 400 assert "non-existent-id" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_create_access_group_surfaces_dropped_models(): + """An access-group write whose reload does not leave the tagged models live on this + pod must report the drop through this file's HTTPException contract, not a 200.""" + from fastapi import HTTPException + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + create_model_group, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + NewModelGroupRequest, + ) + + deploy_a = MagicMock(model_id="deploy-A", model_name="gpt-4o", model_info={}) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=deploy_a) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + mock_user = UserAPIKeyAuth(user_id="test_admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + wiped_router = MagicMock() + wiped_router.get_model_ids.side_effect = [["deploy-A"], []] + with ( + patch("litellm.proxy.proxy_server.llm_router", wiped_router), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch( + "litellm.proxy.management_endpoints.model_access_group_management_endpoints.clear_cache", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(HTTPException) as exc_info: + await create_model_group( + data=NewModelGroupRequest(access_group="production-models", model_ids=["deploy-A"]), + user_api_key_dict=mock_user, + ) + + assert exc_info.value.status_code == 500 + assert "deploy-A" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_tag_deployment_parses_string_model_info_and_refuses_corrupt(): + """The model_info column can arrive as its JSON string; tagging must parse it rather + than crash, and must refuse to rewrite a present-but-unreadable value.""" + from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + _tag_deployment_with_access_group, + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + pair = await _tag_deployment_with_access_group( + model_id="deploy-str", + model_info='{"access_groups": ["existing"]}', + access_group="new-group", + prisma_client=mock_prisma, + ) + assert pair is not None + assert pair[0] == "deploy-str" + assert pair[1]["access_groups"] == ["existing", "new-group"] + + with pytest.raises(ValueError, match="deploy-corrupt"): + await _tag_deployment_with_access_group( + model_id="deploy-corrupt", + model_info="{not json", + access_group="new-group", + prisma_client=mock_prisma, + ) + + +@pytest.mark.asyncio +async def test_delete_access_group_ignores_models_that_were_already_dead(): + """A metadata-only strip over a model this pod never served must not fail the write; + the model's deadness predates the request, and blaming it here would make a broken + model block every access-group fix that touches it.""" + deploy_broken = MagicMock( + model_id="deploy-broken", model_name="broken-model", model_info={"access_groups": ["doomed-group"]} + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[deploy_broken]) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + delete_access_group, + ) + + never_served_router = MagicMock() + never_served_router.get_model_ids.return_value = [] + with ( + patch("litellm.proxy.proxy_server.llm_router", never_served_router), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch( + "litellm.proxy.management_endpoints.model_access_group_management_endpoints.clear_cache", + new=AsyncMock(return_value=None), + ), + ): + response = await delete_access_group( + access_group="doomed-group", + user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert response.models_updated == 1 + mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once() diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 0769568d6cd..c2a0d34a915 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -650,14 +650,14 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): assert key_data.metrics.spend == 10.0 -def _daily_user_spend_record(*, user_id, api_key, spend): +def _daily_user_spend_record(*, user_id, api_key, spend, model="gpt-4", model_group="gpt-4"): """A LiteLLM_DailyUserSpend row as the per-user breakdown reads it.""" return SimpleNamespace( date="2024-01-01", user_id=user_id, api_key=api_key, - model="gpt-4", - model_group="gpt-4", + model=model, + model_group=model_group, custom_llm_provider="openai", mcp_namespaced_tool_name=None, endpoint="/chat/completions", @@ -731,6 +731,64 @@ async def test_get_daily_activity_applies_resolve_entity_metadata_to_breakdown() assert entities["user-no-email"].metadata == {} +@pytest.mark.asyncio +async def test_model_groups_breakdown_keys_by_public_name_with_model_fallback(): + """The usage UI labels model traffic with the model_groups breakdown. + + Keys must be the requested public model name (model_group), and rows with a + NULL or empty model_group (pre-routing failures, rows written before the + column existed) must fall back to their model name instead of being dropped + from the breakdown. The models breakdown keeps the upstream litellm names. + """ + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + + records = [ + _daily_user_spend_record( + user_id="u1", api_key="key-1", spend=7.0, model="gpt-5.2", model_group="gpt-5.2-eu" + ), + _daily_user_spend_record( + user_id="u1", api_key="key-1", spend=3.0, model="gpt-5.2", model_group=None + ), + _daily_user_spend_record( + user_id="u1", api_key="key-1", spend=2.0, model="claude-x", model_group="" + ), + ] + + mock_table = MagicMock() + mock_table.count = AsyncMock(return_value=len(records)) + mock_table.find_many = AsyncMock(return_value=records) + mock_prisma.db.litellm_dailyuserspend = mock_table + mock_prisma.db.litellm_verificationtoken = MagicMock() + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + + result = await get_daily_activity( + prisma_client=mock_prisma, + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id=None, + entity_metadata_field=None, + start_date="2024-01-01", + end_date="2024-01-01", + model=None, + api_key=None, + page=1, + page_size=1000, + ) + + breakdown = result.results[0].breakdown + + assert set(breakdown.model_groups.keys()) == {"gpt-5.2-eu", "gpt-5.2", "claude-x"} + assert breakdown.model_groups["gpt-5.2-eu"].metrics.spend == 7.0 + assert breakdown.model_groups["gpt-5.2"].metrics.spend == 3.0 + assert breakdown.model_groups["claude-x"].metrics.spend == 2.0 + assert breakdown.model_groups["gpt-5.2"].api_key_breakdown["key-1"].metrics.spend == 3.0 + + assert set(breakdown.models.keys()) == {"gpt-5.2", "claude-x"} + assert breakdown.models["gpt-5.2"].metrics.spend == 10.0 + assert breakdown.models["claude-x"].metrics.spend == 2.0 + + class TestAdjustDatesForTimezone: """ Regression tests for the timezone double-counting bug. @@ -852,6 +910,38 @@ class TestBuildAggregatedSqlQuery: assert "model = $4" in sql assert "api_key = $5" in sql + def test_model_group_rollups_fall_back_to_model_name(self): + """Aggregated model_groups rollups must fall back to model for group-less rows. + + The (date, model_group) grouping level cannot recover the model column + after the fact (it is rolled up), so the fallback has to happen in SQL; + without it, group-less rows silently vanish from the model_groups + breakdown that the usage UI now renders by default. Group-less rows are + stored as empty strings, not NULL (spend_tracking_utils defaults + model_group to ""), so a plain COALESCE is not enough: the fallback must + be NULLIF-wrapped to catch both + """ + sql, _ = _build_aggregated_sql_query( + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id=None, + start_date="2026-07-01", + end_date="2026-07-01", + model=None, + api_key=None, + ) + + normalized = " ".join(sql.split()) + fallback = "COALESCE(NULLIF(model_group, ''), model)" + assert f"{fallback} AS model_group" in normalized + assert ( + f"GROUPING(date, api_key, model, {fallback}, " + "custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level" in normalized + ) + assert f"(date, {fallback}), (date, {fallback}, api_key)," in normalized + assert "(date, model_group)" not in normalized + assert "COALESCE(model_group, model)" not in normalized + @pytest.mark.asyncio async def test_get_daily_activity_aggregated_empty_result_set(): diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 5cbc3e72d83..a37f7ca764d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -2893,6 +2893,7 @@ async def test_user_info_v2_response_shape(mocker): "updated_at", "sso_user_id", "teams", + "object_permission", } assert set(response_dict.keys()) == expected_fields @@ -3667,3 +3668,367 @@ async def test_add_user_to_team_keeps_already_a_member_quiet(mocker, caplog): ) assert [r.getMessage() for r in caplog.records if r.levelno >= logging.ERROR] == [] + + +@pytest.mark.asyncio +async def test_get_user_info_for_proxy_admin_validates_keys_and_teams(): + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy._types import LiteLLM_TeamTable, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _get_user_info_for_proxy_admin, + ) + + raw_rows = [ + { + "teams": [ + {"team_id": "team-b", "team_alias": "beta"}, + {"team_id": "team-a", "team_alias": "alpha"}, + ], + "keys": [ + {"token": "hashed-token-1", "team_id": "team-a", "models": None, "spend": 1.0}, + ], + } + ] + + mock_prisma_client = MagicMock() + mock_prisma_client.db.query_raw = AsyncMock(return_value=raw_rows) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): + result = await _get_user_info_for_proxy_admin(user_api_key_dict=UserAPIKeyAuth(user_id=None)) + + assert all(isinstance(team, LiteLLM_TeamTable) for team in result.teams) + assert [team.team_alias for team in result.teams] == ["alpha", "beta"] + assert len(result.keys) == 1 + returned_key = result.keys[0] + assert returned_key["team_id"] == "team-a" + assert returned_key["models"] == [] + + +def _object_permission_mocks(mocker, existing_object_permission_id=None): + """Prisma double whose user row optionally already links a permission row.""" + mock_prisma_client = mocker.MagicMock() + existing_user = mocker.MagicMock() + existing_user.model_dump.return_value = { + "user_id": "target-user", + "object_permission_id": existing_object_permission_id, + } + existing_user.user_id = "target-user" + existing_user.object_permission_id = existing_object_permission_id + mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock( + return_value=existing_user + ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = mocker.AsyncMock( + return_value=None + ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = mocker.AsyncMock( + return_value=SimpleNamespace(object_permission_id="perm-new") + ) + mock_prisma_client.update_data = mocker.AsyncMock( + return_value={"user_id": "target-user"} + ) + mock_prisma_client.jsonify_object = mocker.MagicMock(side_effect=lambda x: x) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") + mocker.patch( + "litellm.proxy.proxy_server._invalidate_spend_counter", + new=mocker.AsyncMock(), + ) + return mock_prisma_client + + +@pytest.mark.asyncio +async def test_user_update_persists_mcp_entitlement_and_links_it(mocker): + """/user/update documents an object_permission param; it must actually be stored. + + The grants live in their own table, so the endpoint has to upsert them and hand the user row + only the resulting object_permission_id. Passing object_permission through to the user update + would not even be a column. + """ + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + mock_prisma_client = _object_permission_mocks(mocker) + cache = mocker.MagicMock() + cache.async_delete_cache = mocker.AsyncMock() + mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) + + await _update_single_user_helper( + user_request=UpdateUserRequest( + user_id="target-user", + object_permission={ + "mcp_servers": ["github"], + "mcp_tool_permissions": {"github": ["list_issues"]}, + }, + ), + user_api_key_dict=UserAPIKeyAuth( + user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN + ), + ) + + upsert_kwargs = mock_prisma_client.db.litellm_objectpermissiontable.upsert.call_args.kwargs + created = upsert_kwargs["data"]["create"] + assert created["mcp_servers"] == ["github"] + assert json.loads(created["mcp_tool_permissions"]) == {"github": ["list_issues"]} + + written = mock_prisma_client.update_data.call_args.kwargs["data"] + assert written["object_permission_id"] == "perm-new" + assert "object_permission" not in written + + +@pytest.mark.asyncio +async def test_user_update_invalidates_the_cached_entitlement(mocker): + """An admin revoking a tool must take effect now, not at the end of the cache TTL. + + Three entries go stale: the permission row (keyed by its own id), the user -> permission link + (which carries a "no entitlement" sentinel), and the cached user row. + """ + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + _object_permission_mocks(mocker) + cache = mocker.MagicMock() + cache.async_delete_cache = mocker.AsyncMock() + mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) + + await _update_single_user_helper( + user_request=UpdateUserRequest( + user_id="target-user", + object_permission={"mcp_tool_permissions": {"github": []}}, + ), + user_api_key_dict=UserAPIKeyAuth( + user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN + ), + ) + + deleted = {call.kwargs["key"] for call in cache.async_delete_cache.call_args_list} + assert deleted == { + "object_permission_id:perm-new", + "user_object_permission_id:target-user", + "target-user", + } + + +@pytest.mark.asyncio +async def test_admin_can_clear_a_users_mcp_entitlement(mocker): + """An explicit empty object_permission means "no object permission", so it must unlink. + + The merge-based upsert cannot express this: merging an empty grant set over the existing row + leaves every grant in place, and the empty-value filter drops the field before the upsert runs, + so without the explicit clear path the documented operation silently returns success unchanged. + + A clear also leaves no incoming permission id, so invalidation keyed off one would skip it and + the gateway would keep enforcing the cleared grants until the cache expired. + """ + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + mock_prisma_client = _object_permission_mocks(mocker, "perm-existing") + cache = mocker.MagicMock() + cache.async_delete_cache = mocker.AsyncMock() + mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) + + await _update_single_user_helper( + user_request=UpdateUserRequest(user_id="target-user", object_permission={}), + user_api_key_dict=UserAPIKeyAuth( + user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN + ), + ) + + written = mock_prisma_client.update_data.call_args.kwargs["data"] + assert written["object_permission_id"] is None + assert "object_permission" not in written + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_not_called() + + deleted = {call.kwargs["key"] for call in cache.async_delete_cache.call_args_list} + assert deleted == { + "object_permission_id:perm-existing", + "user_object_permission_id:target-user", + "target-user", + } + + +@pytest.mark.asyncio +async def test_user_update_invalidates_both_the_old_and_new_permission_rows(mocker): + """An upsert can mint a new permission row, which leaves the outgoing one cached under its id. + + Only the link cache knows the user moved; the old row's own entry still holds the pre-update + grants, so anything still resolving that id keeps reading them. + """ + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + _object_permission_mocks(mocker, "perm-existing") + cache = mocker.MagicMock() + cache.async_delete_cache = mocker.AsyncMock() + mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) + + await _update_single_user_helper( + user_request=UpdateUserRequest( + user_id="target-user", + object_permission={"mcp_tool_permissions": {"github": ["list_issues"]}}, + ), + user_api_key_dict=UserAPIKeyAuth( + user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN + ), + ) + + deleted = {call.kwargs["key"] for call in cache.async_delete_cache.call_args_list} + assert deleted == { + "object_permission_id:perm-existing", + "object_permission_id:perm-new", + "user_object_permission_id:target-user", + "target-user", + } + + +@pytest.mark.asyncio +async def test_non_admin_cannot_clear_their_own_mcp_entitlement(mocker): + """The empty-value filter drops `object_permission: {}` before the guard saw it, so a non-admin + could clear the very ceiling an admin placed on them. The guard reads the fields the caller SENT. + """ + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + mock_prisma_client = _object_permission_mocks(mocker, "perm-existing") + cache = mocker.MagicMock() + cache.async_delete_cache = mocker.AsyncMock() + mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) + + with pytest.raises(HTTPException) as exc: + await _update_single_user_helper( + user_request=UpdateUserRequest(user_id="target-user", object_permission={}), + user_api_key_dict=UserAPIKeyAuth( + user_id="target-user", user_role=LitellmUserRoles.INTERNAL_USER + ), + ) + + assert exc.value.status_code == 403 + mock_prisma_client.update_data.assert_not_called() + + +@pytest.mark.asyncio +async def test_non_admin_cannot_rewrite_their_own_mcp_entitlement(mocker): + """The entitlement bounds the human, so a self-write is an escalation path: an empty grant list + means "no restriction" and would lift a ceiling the admin placed on them.""" + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + mock_prisma_client = _object_permission_mocks(mocker, "perm-existing") + cache = mocker.MagicMock() + cache.async_delete_cache = mocker.AsyncMock() + mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) + + with pytest.raises(HTTPException) as exc: + await _update_single_user_helper( + user_request=UpdateUserRequest( + user_id="target-user", + object_permission={"mcp_servers": [], "mcp_tool_permissions": {}}, + ), + user_api_key_dict=UserAPIKeyAuth( + user_id="target-user", user_role=LitellmUserRoles.INTERNAL_USER + ), + ) + + assert exc.value.status_code == 403 + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_not_called() + mock_prisma_client.update_data.assert_not_called() + + +@pytest.mark.asyncio +async def test_new_user_persists_the_requested_mcp_entitlement(mocker): + """generate_key_helper_fn only forwards object_permission_id, so /user/new has to create the + grants row itself; otherwise the entitlement the admin sent is silently dropped.""" + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db.litellm_objectpermissiontable.create = mocker.AsyncMock( + return_value=SimpleNamespace(object_permission_id="perm-created") + ) + mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock( + return_value=None + ) + mock_prisma_client.db.litellm_usertable.count = mocker.AsyncMock(return_value=0) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.check_if_default_team_set", + return_value=None, + ) + mock_generate = mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.generate_key_helper_fn", + new=mocker.AsyncMock( + return_value={"user_id": "new-human", "token": "sk-x", "expires": None} + ), + ) + mocker.patch( + "litellm.proxy.hooks.user_management_event_hooks.UserManagementEventHooks.async_user_created_hook", + new=mocker.AsyncMock(), + ) + + await new_user( + data=NewUserRequest( + user_id="new-human", + object_permission={"mcp_tool_permissions": {"github": ["list_issues"]}}, + ), + user_api_key_dict=UserAPIKeyAuth( + user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN + ), + ) + + created = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] + assert json.loads(created["mcp_tool_permissions"]) == {"github": ["list_issues"]} + forwarded = mock_generate.call_args.kwargs + assert forwarded["object_permission_id"] == "perm-created" + assert "object_permission" not in forwarded + + +@pytest.mark.asyncio +async def test_user_info_v2_returns_the_mcp_entitlement(mocker): + """The admin UI reads the current entitlement off this endpoint, so the grants have to come back + with the user row rather than only their id.""" + from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 + + user_row = SimpleNamespace( + object_permission=SimpleNamespace( + object_permission_id="perm-1", + mcp_servers=["github"], + mcp_access_groups=[], + mcp_tool_permissions={"github": ["list_issues"]}, + ), + ) + user_row.model_dump = lambda: { + "user_id": "human-1", + "object_permission": { + "object_permission_id": "perm-1", + "mcp_servers": ["github"], + "mcp_access_groups": [], + "mcp_tool_permissions": {"github": ["list_issues"]}, + }, + } + mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.MagicMock()) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_user_info_v2_access", + new=mocker.AsyncMock(return_value=user_row), + ) + + response = await user_info_v2( + request=SimpleNamespace(query_params={}), + user_id="human-1", + user_api_key_dict=UserAPIKeyAuth( + user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN + ), + ) + + assert response.object_permission is not None + assert response.object_permission.mcp_servers == ["github"] + assert response.object_permission.mcp_tool_permissions == { + "github": ["list_issues"] + } diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 51f72f91dc3..867ef759fb3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -2507,6 +2507,195 @@ async def test_update_key_nonexistent_key_returns_404(monkeypatch): assert "Authentication Error" not in str(exc_info.value.message) +def _setup_update_key_mocks(monkeypatch, mock_prisma_client): + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.store_audit_logs", False) + + +@pytest.mark.asyncio +async def test_update_key_by_alias_only(monkeypatch): + """ + /key/update identified by key_alias alone resolves the key row via + find_many on the alias and updates using the resolved token. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b" + key_in_db = LiteLLM_VerificationToken( + token=hashed_token, + key_alias="prod-alias", + user_id="test-user", + max_budget=200.0, + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[key_in_db] + ) + mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock( + return_value=None + ) + mock_prisma_client.update_data = AsyncMock( + return_value={"data": {"max_budget": 50.0}} + ) + _setup_update_key_mocks(monkeypatch, mock_prisma_client) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ) + + request_data = UpdateKeyRequest(key_alias="prod-alias", max_budget=50.0) + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + ) as mock_delete_cache: + mock_delete_cache.return_value = None + result = await update_key_fn( + request=MagicMock(), + data=request_data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + mock_prisma_client.db.litellm_verificationtoken.find_many.assert_called_once_with( + where={"key_alias": "prod-alias"}, take=2 + ) + assert request_data.key == hashed_token + mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_not_called() + mock_prisma_client.update_data.assert_awaited_once() + assert mock_prisma_client.update_data.call_args.kwargs["token"] == hashed_token + assert ( + mock_prisma_client.update_data.call_args.kwargs["data"]["token"] == hashed_token + ) + assert result["key"] == hashed_token + + +@pytest.mark.asyncio +async def test_update_key_by_alias_not_found_returns_404(monkeypatch): + """ + /key/update with a key_alias matching no key returns 404. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[] + ) + _setup_update_key_mocks(monkeypatch, mock_prisma_client) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ) + + with pytest.raises(ProxyException) as exc_info: + await update_key_fn( + request=MagicMock(), + data=UpdateKeyRequest(key_alias="no-such-alias", max_budget=50.0), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert str(exc_info.value.code) == "404" + assert "not found" in str(exc_info.value.message).lower() + mock_prisma_client.update_data.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_key_by_duplicate_alias_returns_400(monkeypatch): + """ + /key/update with a key_alias shared by multiple keys returns 400 + instead of silently updating one of them. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + rows = [ + LiteLLM_VerificationToken(token="hashed-token-1", key_alias="dup-alias"), + LiteLLM_VerificationToken(token="hashed-token-2", key_alias="dup-alias"), + ] + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=rows + ) + _setup_update_key_mocks(monkeypatch, mock_prisma_client) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ) + + with pytest.raises(ProxyException) as exc_info: + await update_key_fn( + request=MagicMock(), + data=UpdateKeyRequest(key_alias="dup-alias", max_budget=50.0), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert str(exc_info.value.code) == "400" + assert "multiple keys" in str(exc_info.value.message).lower() + mock_prisma_client.update_data.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_key_with_key_and_alias_selects_by_key(monkeypatch): + """ + Regression: passing both key and key_alias keeps today's behavior. The key + identifies the row (find_unique, never find_many) and key_alias is the new + alias to set; the response echoes the caller-passed key. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b" + key_in_db = LiteLLM_VerificationToken( + token=hashed_token, + key_alias="old-name", + user_id="test-user", + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=key_in_db + ) + mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock( + return_value=None + ) + mock_prisma_client.update_data = AsyncMock( + return_value={"data": {"key_alias": "new-name"}} + ) + _setup_update_key_mocks(monkeypatch, mock_prisma_client) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ) + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + ) as mock_delete_cache: + mock_delete_cache.return_value = None + result = await update_key_fn( + request=MagicMock(), + data=UpdateKeyRequest(key="sk-test-key", key_alias="new-name"), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + mock_prisma_client.db.litellm_verificationtoken.find_many.assert_not_called() + mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once() + assert mock_prisma_client.update_data.call_args.kwargs["token"] == "sk-test-key" + assert result["key"] == "sk-test-key" + + @pytest.mark.asyncio async def test_block_key_existing_key_succeeds(monkeypatch): """ 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 bd5eda0197b..1e8add52f74 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 @@ -819,6 +819,7 @@ class TestUpdateModel: ) mock_router = MagicMock() + mock_router.get_model_ids.return_value = [model_id] admin_user = UserAPIKeyAuth( user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN ) @@ -838,7 +839,7 @@ class TestUpdateModel: ), patch( "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", - new=AsyncMock(return_value=None), + new=AsyncMock(return_value=True), ) as mock_clear_cache, ): await update_model( @@ -1888,6 +1889,7 @@ class TestAddAndDeleteModelLifecycle: mock_router = MagicMock() mock_router.delete_deployment = MagicMock() + mock_router.get_model_ids.return_value = [model_id] _PS = "litellm.proxy.proxy_server" _ENCRYPT = "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper" @@ -2031,8 +2033,7 @@ class TestDeleteTeamBYOKModelGhost: mock_refresh.assert_awaited_once() assert mock_refresh.await_args.kwargs["team_row"] is updated_team_row - # BYOK internal name can't be an alias value -> the alias-table scan is skipped. - mock_prisma.db.litellm_modeltable.find_many.assert_not_awaited() + mock_prisma.db.litellm_modeltable.find_many.assert_awaited() @pytest.mark.asyncio async def test_delete_non_internal_team_model_still_scans_aliases(self): @@ -2186,6 +2187,171 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma.db.litellm_teamtable.update.assert_not_awaited() mock_refresh.assert_not_awaited() + @pytest.mark.asyncio + async def test_delete_legacy_team_model_scrubs_stale_alias_and_keeps_gateway_access( + self, + ): + """Regression: legacy team models store {public_name: internal model_name} in + the team's model_aliases. delete_model skipped the alias scan for + internal-shaped names, so the stale alias kept rewriting requests for the + public name to a deployment that no longer existed ("no healthy deployments + for model_name_{team_id}_..."). Deleting the deployment must scrub the + alias, and the public name must stay in team.models while a gateway-level + deployment still serves it.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelInfoDelete, + delete_model as delete_model_endpoint, + ) + + team_id = "team-legacy-alias" + model_id = "legacy-alias-model-1" + public_name = "gpt-4" + internal_name = f"model_name_{team_id}_abc-uuid" + + db_row = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name=internal_name, + litellm_params={"model": "openai/gpt-4.1-nano"}, + model_info={"id": model_id, "team_id": team_id}, + created_by="admin", + updated_by="admin", + ) + team_row = LiteLLM_TeamTable( + team_id=team_id, + team_alias="legacy-alias-team", + members_with_roles=[Member(user_id="admin", role="admin")], + models=[public_name], + ) + alias_row = MagicMock( + id="alias-row-1", model_aliases={public_name: internal_name} + ) + alias_row.team = MagicMock() + alias_row.team.team_id = team_id + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_proxymodeltable = AsyncMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( + return_value=db_row + ) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_teamtable = AsyncMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_modeltable = AsyncMock() + mock_prisma.db.litellm_modeltable.find_many = AsyncMock( + return_value=[alias_row] + ) + mock_prisma.db.litellm_modeltable.update = AsyncMock() + + mock_router = MagicMock() + mock_router.model_name_to_deployment_indices = {public_name: [0]} + + admin_user = UserAPIKeyAuth( + user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + _PS = "litellm.proxy.proxy_server" + _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), + patch(f"{_PS}.store_model_in_db", True), + patch(f"{_PS}.premium_user", True), + patch(f"{_PS}.llm_router", mock_router), + patch(f"{_PS}.proxy_logging_obj", MagicMock()), + patch(f"{_PS}.user_api_key_cache", MagicMock()), + patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()) as mock_refresh, + ): + result = await delete_model_endpoint( + model_info=ModelInfoDelete(id=model_id), + user_api_key_dict=admin_user, + ) + + assert "deleted successfully" in result["message"] + + mock_prisma.db.litellm_modeltable.update.assert_awaited_once() + alias_update_kwargs = mock_prisma.db.litellm_modeltable.update.await_args.kwargs + assert alias_update_kwargs["where"] == {"id": "alias-row-1"} + assert json.loads(alias_update_kwargs["data"]["model_aliases"]) == {} + + # A gateway-level deployment still serves the public name -> team access stays. + mock_prisma.db.litellm_teamtable.update.assert_not_awaited() + mock_refresh.assert_not_awaited() + + @pytest.mark.asyncio + async def test_delete_replica_keeps_alias_while_surviving_replica_serves_it(self): + """Deleting one replica of a load-balanced legacy team model (several + deployment rows sharing one internal model_name) must not scrub the team + alias: the surviving replicas still serve the aliased name, so removing + the alias would break routing that works.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelInfoDelete, + delete_model as delete_model_endpoint, + ) + + team_id = "team-lb-legacy" + model_id = "lb-replica-1" + internal_name = f"model_name_{team_id}_shared-uuid" + + db_row = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name=internal_name, + litellm_params={"model": "openai/gpt-4.1-nano"}, + model_info={"id": model_id, "team_id": team_id}, + created_by="admin", + updated_by="admin", + ) + team_row = LiteLLM_TeamTable( + team_id=team_id, + team_alias="lb-legacy-team", + members_with_roles=[Member(user_id="admin", role="admin")], + models=["gpt-4"], + ) + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_proxymodeltable = AsyncMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( + return_value=db_row + ) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_teamtable = AsyncMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_modeltable = AsyncMock() + mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_modeltable.update = AsyncMock() + + mock_router = MagicMock() + mock_router.model_name_to_deployment_indices = {internal_name: [0]} + + admin_user = UserAPIKeyAuth( + user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + _PS = "litellm.proxy.proxy_server" + _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), + patch(f"{_PS}.store_model_in_db", True), + patch(f"{_PS}.premium_user", True), + patch(f"{_PS}.llm_router", mock_router), + patch(f"{_PS}.proxy_logging_obj", MagicMock()), + patch(f"{_PS}.user_api_key_cache", MagicMock()), + patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + ): + result = await delete_model_endpoint( + model_info=ModelInfoDelete(id=model_id), + user_api_key_dict=admin_user, + ) + + assert "deleted successfully" in result["message"] + mock_prisma.db.litellm_modeltable.find_many.assert_not_awaited() + mock_prisma.db.litellm_modeltable.update.assert_not_awaited() + mock_prisma.db.litellm_teamtable.update.assert_not_awaited() + class TestDeleteModelTeamAuth: """Team auth on the /model/delete path. @@ -3002,7 +3168,7 @@ class TestPatchModelBlockedAuthGate: with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), - patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]})), patch("litellm.proxy.proxy_server.store_model_in_db", True), patch("litellm.proxy.proxy_server.premium_user", True), patch( @@ -3048,7 +3214,7 @@ class TestPatchModelBlockedAuthGate: with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), - patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]})), patch("litellm.proxy.proxy_server.store_model_in_db", True), patch("litellm.proxy.proxy_server.premium_user", True), patch( @@ -3057,7 +3223,7 @@ class TestPatchModelBlockedAuthGate: ), patch( "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", - new=AsyncMock(return_value=None), + new=AsyncMock(return_value=True), ), ): result = await patch_model( @@ -3067,3 +3233,313 @@ class TestPatchModelBlockedAuthGate: ) assert result is updated_row mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once() + + +class TestWriteSurfacesReloadDrop: + """A model-write endpoint may report success only if every row it wrote is, after the + reload it triggered, live in this pod's router or deliberately environment-inactive.""" + + def test_reload_serving_verdict_matrix(self, monkeypatch): + import litellm + from litellm.proxy.management_endpoints.model_management_endpoints import ( + reload_serving_verdict, + ) + + live_router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "gpt-4o"}, + "model_info": {"id": "m-live", "db_model": True}, + } + ] + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", live_router) + monkeypatch.setenv("LITELLM_ENVIRONMENT", "development") + + written = [ + ("m-live", {"id": "m-live"}), + ("m-gone", {"id": "m-gone"}), + ("m-env", {"id": "m-env", "supported_environments": ["production"]}), + ("m-env-str", '{"id": "m-env-str", "supported_environments": ["production"]}'), + ("m-env-misconfigured", {"id": "m-env-misconfigured", "supported_environments": ["bogus"]}), + ("m-corrupt", "{not json"), + ] + missing, collateral = reload_serving_verdict( + before=frozenset({"m-live", "m-collateral"}), written_models=written, written_must_serve=True + ) + assert missing == ("m-gone", "m-env-misconfigured", "m-corrupt") + assert collateral == ("m-collateral",) + + missing, collateral = reload_serving_verdict( + before=frozenset({"m-live", "m-was-live"}), + written_models=[("m-live", None), ("m-was-live", None), ("m-never-lived", None)], + written_must_serve=False, + ) + assert missing == ("m-was-live",) + assert collateral == () + + def test_raise_if_reload_degraded_serving_contract(self, monkeypatch): + import litellm + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + raise_if_reload_degraded_serving, + ) + + live_router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "gpt-4o"}, + "model_info": {"id": "m-live", "db_model": True}, + } + ] + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", live_router) + + assert ( + raise_if_reload_degraded_serving( + before=frozenset({"m-live"}), written_models=[("m-live", None)], action="update" + ) + is None + ) + + with pytest.raises(ProxyException, match="m-gone"): + raise_if_reload_degraded_serving( + before=frozenset(), written_models=[("m-gone", None)], action="update" + ) + + with pytest.raises(ProxyException, match="m-collateral"): + raise_if_reload_degraded_serving( + before=frozenset({"m-live", "m-collateral"}), written_models=[("m-live", None)], action="update" + ) + + +class TestModelInfoAsMapping: + """The model_info column reaches consumers as a dict or as its JSON string; this is + the single owner of that parse, and None means no usable mapping.""" + + def test_contract(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + model_info_as_mapping, + ) + + assert model_info_as_mapping({"id": "m1"}) == {"id": "m1"} + assert model_info_as_mapping('{"id": "m1"}') == {"id": "m1"} + assert model_info_as_mapping(None) is None + assert model_info_as_mapping("{not json") is None + assert model_info_as_mapping('["a", "b"]') is None + assert model_info_as_mapping(42) is None + + +class TestStrategyRouterWriteValidation: + """Management write paths must reject litellm_params.model values that would + corrupt a strategy router's pseudo-model (LIT-4663). The router loads these + deployments by the auto_router/ discriminator, so a mangled string makes it + drop the deployment silently under ignore_invalid_deployments; the mistake + has to fail loudly at the API boundary instead.""" + + def _stored_complexity_params(self) -> LiteLLM_Params: + return LiteLLM_Params( + model="auto_router/complexity_router", + complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini"}}, + ) + + def _db_complexity_router(self, model_id: str) -> Deployment: + return Deployment( + model_name="my-auto-router", + litellm_params=self._stored_complexity_params(), + model_info={"id": model_id}, + ) + + def test_double_prefix_rejected_against_stored_params(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + violation = _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(model="auto_router/auto_router/complexity_router"), + existing_params=self._stored_complexity_params(), + ) + assert violation is not None + assert "repeats" in violation + + def test_prefix_strip_rejected_against_stored_params(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + violation = _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(model="complexity_router"), + existing_params=self._stored_complexity_params(), + ) + assert violation is not None + assert "does not start with" in violation + + def test_patch_without_model_is_not_judged(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + assert ( + _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(rpm=10), + existing_params=self._stored_complexity_params(), + ) + is None + ) + assert _strategy_router_write_violation(incoming_params=None, existing_params=None) is None + + def test_restore_of_corrupted_row_is_allowed(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + corrupted = LiteLLM_Params( + model="auto_router/auto_router/complexity_router", + complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini"}}, + ) + assert ( + _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(model="auto_router/complexity_router"), + existing_params=corrupted, + ) + is None + ) + + def test_create_semantic_router_missing_embedding_rejected(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + + violation = _strategy_router_write_violation( + incoming_params=LiteLLM_Params( + model="auto_router/my-router", + auto_router_config="{}", + auto_router_default_model="gpt-4o-mini", + ), + existing_params=None, + ) + assert violation is not None + assert "auto_router_embedding_model" in violation + + @pytest.mark.asyncio + async def test_patch_model_rejects_double_prefix(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + patch_model, + ) + from litellm.types.router import updateLiteLLMParams + + model_id = "strategy-router-patch-test" + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.get_db_model", + new=AsyncMock(return_value=self._db_complexity_router(model_id)), + ), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints._update_team_model_in_db", + new=AsyncMock(), + ) as mock_update, + ): + with pytest.raises(ProxyException) as exc_info: + await patch_model( + model_id=model_id, + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams(model="auto_router/auto_router/complexity_router") + ), + user_api_key_dict=admin, + ) + assert "repeats" in str(exc_info.value.message) + mock_update.assert_not_awaited() + + @pytest.mark.asyncio + async def test_add_new_model_rejects_prefixed_model_without_config(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + mock_prisma = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model( + model_params=Deployment( + model_name="my-auto-router", + litellm_params=LiteLLM_Params(model="auto_router/complexity_router"), + model_info={"id": "strategy-router-create-test"}, + ), + user_api_key_dict=admin, + ) + assert "requires" in str(exc_info.value.message) + mock_prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_update_model_rejects_prefix_strip(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + update_model, + ) + from litellm.types.router import ModelInfo, updateLiteLLMParams + + model_id = "strategy-router-update-test" + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + existing_row = MagicMock() + existing_row.model_dump.return_value = { + "model_name": "my-auto-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": {"tiers": {"SIMPLE": "gpt-4o-mini"}}, + }, + "model_info": {"id": model_id}, + } + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(ProxyException) as exc_info: + await update_model( + model_params=updateDeployment( + litellm_params=updateLiteLLMParams(model="complexity_router"), + model_info=ModelInfo(id=model_id), + ), + user_api_key_dict=admin, + ) + assert "does not start with" in str(exc_info.value.message) + mock_prisma.db.litellm_proxymodeltable.update.assert_not_awaited() diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 5202c8cbfc0..a658a4b7353 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -6578,6 +6578,7 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit(): return_value=mock_existing_team ) mock_cache.async_set_cache = AsyncMock() + mock_logging.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() # Mock team update mock_updated_team = MagicMock(spec=LiteLLM_TeamTable) @@ -10223,3 +10224,68 @@ def test_patch_team_route_publishes_its_request_body_schema(): assert schema == {"$ref": "#/components/schemas/PatchTeamRequest"} properties = app.openapi()["components"]["schemas"]["PatchTeamRequest"]["properties"] assert "tpm_limit" in properties and "metadata" in properties + + +@pytest.mark.asyncio +async def test_get_all_team_memberships_validates_rows(): + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.management_endpoints.team_endpoints import ( + get_all_team_memberships, + ) + + membership_row = MagicMock() + membership_row.model_dump = lambda: { + "user_id": "member-1", + "team_id": "team-1", + "spend": 2.5, + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership_row]) + + result = await get_all_team_memberships(mock_prisma_client, ["team-1"], user_id="member-1") + + assert len(result) == 1 + assert isinstance(result[0], LiteLLM_TeamMembership) + assert result[0].user_id == "member-1" + assert result[0].team_id == "team-1" + assert result[0].spend == 2.5 + find_many_kwargs = mock_prisma_client.db.litellm_teammembership.find_many.call_args.kwargs + assert find_many_kwargs["where"] == {"team_id": {"in": ["team-1"]}, "user_id": {"in": ["member-1"]}} + + +@pytest.mark.asyncio +async def test_list_available_teams_filters_joined_and_validates_rows(monkeypatch): + from fastapi import Request + + import litellm + from litellm.proxy.management_endpoints.team_endpoints import list_available_teams + + monkeypatch.setattr( + litellm, + "default_internal_user_params", + {"available_teams": ["team-open", "team-joined"]}, + ) + + user_row = MagicMock() + user_row.model_dump = lambda: {"user_id": "u-1", "teams": ["team-joined"]} + + open_team_row = MagicMock() + open_team_row.model_dump = lambda: {"team_id": "team-open", "team_alias": "open team"} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[open_team_row]) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): + result = await list_available_teams( + http_request=MagicMock(spec=Request), + user_api_key_dict=UserAPIKeyAuth(user_id="u-1"), + ) + + assert len(result) == 1 + assert isinstance(result[0], LiteLLM_TeamTable) + assert result[0].team_id == "team-open" + assert result[0].team_alias == "open team" + find_many_kwargs = mock_prisma_client.db.litellm_teamtable.find_many.call_args.kwargs + assert find_many_kwargs["where"] == {"team_id": {"in": ["team-open"]}} diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 46ecb31e1c8..ac01c6ae1d1 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -2610,3 +2610,444 @@ def test_list_files_with_all_proxy_models_team_uses_openai_deployment( assert captured_kwargs.get("api_key") == "team-openai-key" assert captured_kwargs.get("custom_llm_provider") == "openai" proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def _setup_vertex_named_credential_router(monkeypatch) -> Router: + from litellm.types.utils import CredentialItem + + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="vertex-named-cred", + credential_info={}, + credential_values={ + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + }, + ) + ], + ) + return Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "litellm_credential_name": "vertex-named-cred", + }, + } + ] + ) + + +def _assert_vertex_named_credentials_attached(captured_kwargs: dict) -> None: + assert captured_kwargs.get("custom_llm_provider") == "vertex_ai" + assert captured_kwargs.get("vertex_project") == "customer-project" + assert captured_kwargs.get("vertex_location") == "us-central1" + assert captured_kwargs.get("vertex_credentials") == "/creds/customer-sa.json" + assert captured_kwargs.get("model") is None + + +def test_create_file_provider_only_resolves_named_vertex_credentials( + mocker: MockerFixture, monkeypatch +): + """ + POST /v1/files with only a custom-llm-provider header (no model, no + target_model_names) must attach the configured named vertex credential to + the upstream call instead of falling through to google.auth.default(), + which uploads into the hosting environment's GCP project. + """ + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + router = _setup_vertex_named_credential_router(monkeypatch) + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_acreate_file(**kwargs): + captured_kwargs.update(kwargs) + return OpenAIFileObject( + id="file-vertex-123", + object="file", + bytes=2, + created_at=1234567890, + filename="batch.jsonl", + purpose="batch", + status="uploaded", + ) + + monkeypatch.setattr(litellm, "acreate_file", _mock_acreate_file) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + + try: + response = client.post( + "/v1/files", + files={"file": ("batch.jsonl", b"{}", "application/jsonl")}, + data={"purpose": "batch"}, + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "vertex_ai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + _assert_vertex_named_credentials_attached(captured_kwargs) + proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def test_get_file_provider_only_resolves_named_vertex_credentials( + mocker: MockerFixture, monkeypatch +): + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + router = _setup_vertex_named_credential_router(monkeypatch) + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_afile_retrieve(**kwargs): + captured_kwargs.update(kwargs) + return OpenAIFileObject( + id="file-abc123", + object="file", + bytes=2, + created_at=1234567890, + filename="batch.jsonl", + purpose="batch", + status="uploaded", + ) + + monkeypatch.setattr(litellm, "afile_retrieve", _mock_afile_retrieve) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + + try: + response = client.get( + "/v1/files/file-abc123", + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "vertex_ai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert captured_kwargs.get("file_id") == "file-abc123" + _assert_vertex_named_credentials_attached(captured_kwargs) + proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def test_get_file_content_provider_only_resolves_named_vertex_credentials( + mocker: MockerFixture, monkeypatch +): + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + router = _setup_vertex_named_credential_router(monkeypatch) + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_afile_content(**kwargs): + captured_kwargs.update(kwargs) + return HttpxBinaryResponseContent( + response=httpx.Response( + status_code=200, + content=b"vertex-bytes", + headers={"content-type": "application/octet-stream"}, + ) + ) + + monkeypatch.setattr(litellm, "afile_content", _mock_afile_content) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + + try: + response = client.get( + "/v1/files/file-abc123/content", + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "vertex_ai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert response.content == b"vertex-bytes" + assert captured_kwargs.get("file_id") == "file-abc123" + _assert_vertex_named_credentials_attached(captured_kwargs) + proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def test_delete_file_provider_only_resolves_named_vertex_credentials( + mocker: MockerFixture, monkeypatch +): + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + router = _setup_vertex_named_credential_router(monkeypatch) + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_afile_delete(**kwargs): + captured_kwargs.update(kwargs) + return OpenAIFileObject( + id="file-abc123", + object="file", + bytes=2, + created_at=1234567890, + filename="batch.jsonl", + purpose="batch", + status="uploaded", + ) + + monkeypatch.setattr(litellm, "afile_delete", _mock_afile_delete) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + + try: + response = client.delete( + "/v1/files/file-abc123", + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "vertex_ai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert captured_kwargs.get("file_id") == "file-abc123" + _assert_vertex_named_credentials_attached(captured_kwargs) + proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def test_create_file_provider_only_skips_other_team_vertex_deployment( + mocker: MockerFixture, monkeypatch +): + """ + Regression: with a team-scoped vertex deployment indexed before a global + one under the same model name, a provider-only upload from a different + team must use the global deployment's credentials, never the other + team's. + """ + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + router = Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "team-b-project", + }, + "model_info": { + "id": "team-b-vertex", + "team_id": "team-b", + "team_public_model_name": "gemini-2.5-pro", + }, + }, + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "shared-project", + }, + }, + ] + ) + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_acreate_file(**kwargs): + captured_kwargs.update(kwargs) + return OpenAIFileObject( + id="file-vertex-456", + object="file", + bytes=2, + created_at=1234567890, + filename="batch.jsonl", + purpose="batch", + status="uploaded", + ) + + monkeypatch.setattr(litellm, "acreate_file", _mock_acreate_file) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="test-user", + team_id="team-a", + team_models=["gemini-2.5-pro"], + ) + + try: + response = client.post( + "/v1/files", + files={"file": ("batch.jsonl", b"{}", "application/jsonl")}, + data={"purpose": "batch"}, + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "vertex_ai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert captured_kwargs.get("vertex_project") == "shared-project" + proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def _team_openai_plus_global_anthropic_router() -> Router: + return Router( + model_list=[ + { + "model_name": "team-gpt", + "litellm_params": { + "model": "openai/gpt-4o", + "api_key": "team-openai-key", + }, + "model_info": { + "id": "team-a-openai", + "team_id": "team-a", + "team_public_model_name": "team-gpt", + }, + }, + { + "model_name": "claude-opus-4-6", + "litellm_params": { + "model": "anthropic/claude-opus-4-6", + "api_key": "anthropic-key", + }, + }, + ] + ) + + +def _list_files_captured_kwargs( + mocker: MockerFixture, monkeypatch, router: Router, key_models: list +) -> dict: + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_success_hook = mocker.AsyncMock(return_value=[]) + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_afile_list(**kwargs): + captured_kwargs.update(kwargs) + return [] + + monkeypatch.setattr(litellm, "afile_list", _mock_afile_list) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="test-user", + team_id="team-a", + team_models=["team-gpt", "claude-opus-4-6"], + models=key_models, + ) + + try: + response = client.get( + "/v1/files", + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "openai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + return captured_kwargs + + +def test_list_files_key_restricted_to_other_provider_does_not_leak_team_openai_credentials( + mocker: MockerFixture, monkeypatch +): + """ + Regression: a key restricted to an anthropic model on a team that also has + an openai deployment must not attach the team's openai credentials to a + provider-only openai files call; key-level model restrictions apply to + credential resolution, not just completions. + """ + captured_kwargs = _list_files_captured_kwargs( + mocker, monkeypatch, _team_openai_plus_global_anthropic_router(), ["claude-opus-4-6"] + ) + assert captured_kwargs.get("api_key") != "team-openai-key" + + +def test_list_files_key_allowed_openai_model_still_resolves_team_credentials( + mocker: MockerFixture, monkeypatch +): + """ + A key whose allowlist includes the team's openai model keeps resolving that + deployment's credentials for provider-only openai files calls. + """ + captured_kwargs = _list_files_captured_kwargs( + mocker, monkeypatch, _team_openai_plus_global_anthropic_router(), ["team-gpt"] + ) + assert captured_kwargs.get("api_key") == "team-openai-key" diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index cf3351c4ff8..181846fe289 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -3022,3 +3022,172 @@ class TestCursorProxyRoute: assert call_args["target"] == "https://api.cursor.com/v0/agents" assert result["id"] == "bc_abc123" assert result["status"] == "CREATING" + + +class TestVertexRawPredictStreamingClassification: + """ + Regression coverage for LIT-4761. + + `_base_vertex_proxy_route` classified any target URL containing "stream" as a + streaming request. `:streamRawPredict` carries that substring, so a unary + Anthropic-on-Vertex call (no `stream` field in the body) was sent with + `?alt=sse` and logged through the streaming chunk collector, which parses + Anthropic SSE deltas and finds no usage in a complete `"type": "message"` + body; the spend log recorded 0 tokens and $0 cost. + + Streaming for the rawPredict family is decided by the request body, per the + Anthropic Messages contract. The Gemini generateContent family stays + URL-signalled because the Gemini REST body has no `stream` field. + """ + + RAW_PREDICT_ENDPOINT = ( + "v1/projects/test-project/locations/us-east5/publishers/anthropic/models/" + "claude-sonnet-4-6:streamRawPredict" + ) + GENERATE_CONTENT_ENDPOINT = ( + "v1/projects/test-project/locations/us-east5/publishers/google/models/" + "gemini-2.5-flash:streamGenerateContent" + ) + + async def _capture_passthrough_kwargs(self, endpoint: str, body: object) -> dict: + raw_body = json.dumps(body).encode("utf-8") + + async def receive(): + return {"type": "http.request", "body": raw_body, "more_body": False} + + request = Request( + { + "type": "http", + "method": "POST", + "path": f"/vertex_ai/{endpoint}", + "headers": [(b"content-type", b"application/json")], + "query_string": b"", + }, + receive=receive, + ) + + captured: dict = {} + + def fake_create_pass_through_route(**kwargs): + captured.update(kwargs) + return AsyncMock(return_value={"status": "success"}) + + mock_credentials = Mock() + mock_credentials.token = "test-token" + + base_url = "https://us-east5-aiplatform.googleapis.com/" + mock_handler = Mock() + mock_handler.get_default_base_target_url.return_value = base_url + mock_handler.update_base_target_url_with_credential_location = Mock(return_value=base_url) + + module = "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints" + with ( + mock.patch( + "litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth", + return_value=(mock_credentials, "test-project"), + ), + mock.patch(f"{module}.create_pass_through_route", side_effect=fake_create_pass_through_route), + mock.patch(f"{module}.get_litellm_virtual_key", return_value="Bearer test-key"), + mock.patch(f"{module}.user_api_key_auth", new=AsyncMock(return_value={"api_key": "test-key"})), + mock.patch(f"{module}.get_vertex_pass_through_handler", return_value=mock_handler), + ): + await vertex_proxy_route( + endpoint=endpoint, + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(token="test-key"), + ) + + assert captured, "create_pass_through_route was never called" + return captured + + @pytest.mark.asyncio + async def test_raw_predict_without_stream_field_is_not_streaming(self): + captured = await self._capture_passthrough_kwargs( + endpoint=self.RAW_PREDICT_ENDPOINT, + body={ + "anthropic_version": "vertex-2023-10-16", + "messages": [{"role": "user", "content": "Explain MLOps"}], + "max_tokens": 5000, + }, + ) + + assert captured["is_streaming_request"] is False + assert "alt=sse" not in captured["target"] + + @pytest.mark.asyncio + async def test_raw_predict_with_stream_false_is_not_streaming(self): + captured = await self._capture_passthrough_kwargs( + endpoint=self.RAW_PREDICT_ENDPOINT, + body={ + "anthropic_version": "vertex-2023-10-16", + "stream": False, + "messages": [{"role": "user", "content": "Explain MLOps"}], + "max_tokens": 5000, + }, + ) + + assert captured["is_streaming_request"] is False + assert "alt=sse" not in captured["target"] + + @pytest.mark.asyncio + async def test_raw_predict_with_stream_true_still_streams(self): + captured = await self._capture_passthrough_kwargs( + endpoint=self.RAW_PREDICT_ENDPOINT, + body={ + "anthropic_version": "vertex-2023-10-16", + "stream": True, + "messages": [{"role": "user", "content": "Explain MLOps"}], + "max_tokens": 5000, + }, + ) + + assert captured["is_streaming_request"] is True + assert captured["target"].endswith("?alt=sse") + + @pytest.mark.asyncio + async def test_gemini_stream_generate_content_stays_url_signalled(self): + captured = await self._capture_passthrough_kwargs( + endpoint=self.GENERATE_CONTENT_ENDPOINT, + body={"contents": [{"role": "user", "parts": [{"text": "Explain MLOps"}]}]}, + ) + + assert captured["is_streaming_request"] is True + assert captured["target"].endswith("?alt=sse") + + @pytest.mark.asyncio + async def test_raw_predict_with_non_object_body_is_not_streaming(self): + captured = await self._capture_passthrough_kwargs( + endpoint=self.RAW_PREDICT_ENDPOINT, + body=[{"role": "user", "content": "Explain MLOps"}], + ) + + assert captured["is_streaming_request"] is False + assert "alt=sse" not in captured["target"] + + +@pytest.mark.parametrize( + "request_body, expected", + [ + ({"stream": True}, True), + ({"stream": "true"}, True), + ({"stream": False}, False), + ({}, False), + ([{"role": "user"}], False), + ([], False), + ("stream", False), + (7, False), + (None, False), + ], +) +def test_is_passthrough_request_streaming_tolerates_non_object_bodies(request_body, expected): + """ + A JSON request body is not required to be an object. Every passthrough + streaming decision funnels through this predicate, so a list or scalar body + must answer False instead of raising AttributeError. + """ + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + is_passthrough_request_streaming, + ) + + assert is_passthrough_request_streaming(request_body) is expected diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py index 5de682ec8a0..52da7a4a81d 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py @@ -590,22 +590,20 @@ class TestVertexAIBatchCostCalculation: assert usage.completion_tokens == 0 assert usage.total_tokens == 0 - def test_openai_shaped_output_records_nonzero_cost_and_usage(self): + @pytest.mark.asyncio + async def test_openai_shaped_output_records_nonzero_cost_and_usage(self): """ Regression test for the bug where Vertex batch cost/usage was always 0. After PR #25627 (transform_file_content_response), the GCS predictions.jsonl is rewritten into OpenAI batch shape before the cost-tracking path sees it. - With disable_vertex_batch_output_transformation=False (default), the content - is OpenAI-shaped, so _batch_cost_calculator must fall through to the generic - path rather than calling calculate_vertex_ai_batch_cost_and_usage (which only - reads raw usageMetadata fields). + With disable_vertex_batch_output_transformation=False (default), the cost + dispatch must fall through to the generic aggregation path rather than + calling calculate_vertex_ai_batch_cost_and_usage (which only reads raw + usageMetadata fields). """ import litellm - from litellm.batches.batch_utils import ( - _batch_cost_calculator, - _get_batch_job_total_usage_from_file_content, - ) + from litellm.batches.batch_utils import calculate_batch_cost_and_usage openai_shaped_responses = [ { @@ -668,12 +666,7 @@ class TestVertexAIBatchCostCalculation: try: litellm.disable_vertex_batch_output_transformation = False - cost = _batch_cost_calculator( - file_content_dictionary=openai_shaped_responses, - custom_llm_provider="vertex_ai", - model_name="gemini-2.0-flash-001", - ) - usage = _get_batch_job_total_usage_from_file_content( + cost, usage, _ = await calculate_batch_cost_and_usage( file_content_dictionary=openai_shaped_responses, custom_llm_provider="vertex_ai", model_name="gemini-2.0-flash-001", @@ -694,16 +687,14 @@ class TestVertexAIBatchCostCalculation: cost > 0 ), f"expected non-zero cost for completed Vertex batch, got {cost}" - def test_raw_vertex_output_still_works_when_transformation_disabled(self): + @pytest.mark.asyncio + async def test_raw_vertex_output_still_works_when_transformation_disabled(self): """ When disable_vertex_batch_output_transformation=True the GCS file is returned as raw Vertex predictions.jsonl; the specialized reader must be used. """ import litellm - from litellm.batches.batch_utils import ( - _batch_cost_calculator, - _get_batch_job_total_usage_from_file_content, - ) + from litellm.batches.batch_utils import calculate_batch_cost_and_usage raw_vertex_responses = [ { @@ -727,12 +718,7 @@ class TestVertexAIBatchCostCalculation: try: litellm.disable_vertex_batch_output_transformation = True - cost = _batch_cost_calculator( - file_content_dictionary=raw_vertex_responses, - custom_llm_provider="vertex_ai", - model_name="gemini-2.0-flash-001", - ) - usage = _get_batch_job_total_usage_from_file_content( + cost, usage, _ = await calculate_batch_cost_and_usage( file_content_dictionary=raw_vertex_responses, custom_llm_provider="vertex_ai", model_name="gemini-2.0-flash-001", diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py index 00c3c5b1e74..b0e8a85d3fa 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py @@ -292,3 +292,149 @@ def test_model_group_info_invalid_method(client, auth_as, null_router): response = client.post("/model_group/info", json={}) assert response.status_code == 405 assert len(response.content) > 0 + + +# --------------------------------------------------------------------------- +# GET /v2/model/info?exclude_auto_routers +# --------------------------------------------------------------------------- + + +@pytest.fixture +def mixed_auto_router_router(monkeypatch): + """Router carrying one ordinary deployment per auto-router strategy plus two plain ones.""" + model_list = [ + { + "model_name": "gpt-4o-mini", + "litellm_params": {"model": "openai/gpt-4o-mini"}, + "model_info": {"id": "plain-1", "db_model": False}, + }, + { + "model_name": "tri-tier-router", + "litellm_params": {"model": "auto_router/complexity_router"}, + "model_info": {"id": "auto-complexity", "db_model": True}, + }, + { + "model_name": "support-router", + "litellm_params": {"model": "auto_router/support-router"}, + "model_info": {"id": "auto-semantic", "db_model": True}, + }, + { + "model_name": "adaptive-router", + "litellm_params": {"model": "auto_router/adaptive_router"}, + "model_info": {"id": "auto-adaptive", "db_model": True}, + }, + { + "model_name": "claude-opus", + "litellm_params": {"model": "anthropic/claude-opus-4-6"}, + "model_info": {"id": "plain-2", "db_model": False}, + }, + ] + from unittest.mock import AsyncMock + + router = MagicMock() + router.model_list = model_list + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "llm_model_list", model_list) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + monkeypatch.setattr(proxy_server, "user_model", None) + monkeypatch.setattr(proxy_server.proxy_config, "get_config", AsyncMock(return_value={})) + monkeypatch.setattr( + proxy_server, + "_apply_search_filter_to_models", + AsyncMock(side_effect=lambda all_models, **kw: (all_models, len(all_models))), + ) + monkeypatch.setattr(proxy_server, "_enrich_model_info_with_litellm_data", lambda model, **kw: model) + + import litellm.proxy.agent_endpoints.model_list_helpers as mlh + + monkeypatch.setattr(mlh, "append_agents_to_model_info", AsyncMock(side_effect=lambda models, **kw: models)) + yield router + + +def _model_names(payload) -> list: + return [m["model_name"] for m in payload["data"]] + + +def test_v2_model_info_includes_auto_routers_by_default(client, auth_as, mixed_auto_router_router): + """The new param is opt-in; omitting it must not change what any existing caller sees.""" + with auth_as(): + response = client.get("/v2/model/info") + assert response.status_code == 200 + payload = response.json() + assert "tri-tier-router" in _model_names(payload) + assert payload["total_count"] == 5 + + +def test_v2_model_info_excludes_every_auto_router_strategy(client, auth_as, mixed_auto_router_router): + """All four `auto_router/*` strategies go, not just the semantic one that + Router._is_auto_router_deployment recognises.""" + with auth_as(): + response = client.get("/v2/model/info", params={"exclude_auto_routers": "true"}) + assert response.status_code == 200 + payload = response.json() + assert _model_names(payload) == ["gpt-4o-mini", "claude-opus"] + + +def test_v2_model_info_exclude_auto_routers_shrinks_total_count(client, auth_as, mixed_auto_router_router): + """The filter must run before the count, or the table pages off a total that + includes rows it never renders (49 shown, 50 claimed).""" + with auth_as(): + response = client.get("/v2/model/info", params={"exclude_auto_routers": "true"}) + payload = response.json() + assert payload["total_count"] == 2 + assert len(payload["data"]) == payload["total_count"] + + +def test_v2_model_info_exclude_auto_routers_paginates_over_the_filtered_set( + client, auth_as, mixed_auto_router_router +): + """Page size applies to the filtered list, so no page silently comes back short.""" + with auth_as(): + response = client.get( + "/v2/model/info", params={"exclude_auto_routers": "true", "page": 1, "size": 1} + ) + payload = response.json() + assert payload["total_count"] == 2 + assert payload["total_pages"] == 2 + assert len(payload["data"]) == 1 + + +@pytest.mark.asyncio +async def test_model_info_v2_query_sentinel_does_not_filter(monkeypatch, mixed_auto_router_router): + """Called directly (not through FastAPI) the default arrives as a truthy Query object. + Guarding on `is True` is what stops every direct-call test from silently filtering.""" + from unittest.mock import AsyncMock + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + monkeypatch.setattr(proxy_server.proxy_config, "get_config", AsyncMock(return_value={})) + monkeypatch.setattr( + proxy_server, + "_apply_search_filter_to_models", + AsyncMock(side_effect=lambda all_models, **kw: (all_models, len(all_models))), + ) + monkeypatch.setattr(proxy_server, "_enrich_model_info_with_litellm_data", lambda model, **kw: model) + + import litellm.proxy.agent_endpoints.model_list_helpers as mlh + + monkeypatch.setattr(mlh, "append_agents_to_model_info", AsyncMock(side_effect=lambda models, **kw: models)) + + admin = UserAPIKeyAuth(user_id="u", user_role=LitellmUserRoles.PROXY_ADMIN) + # Deliberately omit exclude_auto_routers, exactly as the pre-existing direct-call tests do. + resp = await proxy_server.model_info_v2( + user_api_key_dict=admin, + model=None, + user_models_only=False, + include_team_models=False, + debug=False, + page=1, + size=50, + search=None, + modelId=None, + teamId=None, + sortBy=None, + sortOrder="asc", + ) + + assert "tri-tier-router" in [m["model_name"] for m in resp["data"]] diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 67945436987..795a99ec266 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -1548,6 +1548,7 @@ async def test_ui_view_session_spend_logs_pagination(client, monkeypatch): assert page_size == 1 assert skip == 1 # page=2, page_size=1 assert 'ORDER BY "startTime" DESC' in sql_query + assert '"user" = $4' not in sql_query return [mock_spend_logs[0]] class MockPrismaClient: @@ -1558,20 +1559,144 @@ async def test_ui_view_session_spend_logs_pagination(client, monkeypatch): mock_prisma_client = MockPrismaClient() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - response = client.get( - "/spend/logs/session/ui", - params={"session_id": "session-123", "page": 2, "page_size": 1}, - headers={"Authorization": "Bearer sk-test"}, + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" ) - assert response.status_code == 200 - data = response.json() - assert data["total"] == 2 - assert data["page"] == 2 - assert data["page_size"] == 1 - assert data["total_pages"] == 2 - assert len(data["data"]) == 1 - assert data["data"][0]["request_id"] == "req1" + try: + response = client.get( + "/spend/logs/session/ui", + params={"session_id": "session-123", "page": 2, "page_size": 1}, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["total"] == 2 + assert data["page"] == 2 + assert data["page_size"] == 1 + assert data["total_pages"] == 2 + assert len(data["data"]) == 1 + assert data["data"][0]["request_id"] == "req1" + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_session_spend_logs_scopes_non_admin_to_own_logs(client, monkeypatch): + own_log = { + "id": "log1", + "request_id": "req1", + "session_id": "session-123", + "user": "user-1", + "startTime": "2024-01-01T00:00:00Z", + } + + class MockDB: + async def count(self, *args, **kwargs): + assert kwargs.get("where") == {"session_id": "session-123", "user": "user-1"} + return 1 + + async def query_raw(self, sql_query, session_id, page_size, skip, scoped_user): + assert session_id == "session-123" + assert scoped_user == "user-1" + assert '"user" = $4' in sql_query + return [own_log] + + class MockPrismaClient: + def __init__(self): + self.db = MockDB() + self.db.litellm_spendlogs = self.db + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + + async def no_permitted_teams(*args, **kwargs): + return [] + + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", + no_permitted_teams, + ) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user-1" + ) + + try: + response = client.get( + "/spend/logs/session/ui", + params={"session_id": "session-123", "page": 1, "page_size": 50}, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["total"] == 1 + assert [row["request_id"] for row in data["data"]] == ["req1"] + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_session_spend_logs_includes_permitted_team_logs(client, monkeypatch): + class MockDB: + async def count(self, *args, **kwargs): + assert kwargs.get("where") == { + "session_id": "session-123", + "OR": [ + {"user": "user-1"}, + {"team_id": {"in": ["team-9"]}}, + ], + } + return 1 + + async def query_raw(self, sql_query, session_id, page_size, skip, scoped_user, team_ids): + assert session_id == "session-123" + assert scoped_user == "user-1" + assert team_ids == ["team-9"] + assert '("user" = $4 OR team_id = ANY($5::text[]))' in sql_query + return [ + { + "id": "log2", + "request_id": "req2", + "session_id": "session-123", + "team_id": "team-9", + "startTime": "2024-01-02T00:00:00Z", + } + ] + + class MockPrismaClient: + def __init__(self): + self.db = MockDB() + self.db.litellm_spendlogs = self.db + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + + async def permitted_teams(*args, **kwargs): + return ["team-9"] + + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", + permitted_teams, + ) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user-1" + ) + + try: + response = client.get( + "/spend/logs/session/ui", + params={"session_id": "session-123", "page": 1, "page_size": 50}, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["total"] == 1 + assert [row["request_id"] for row in data["data"]] == ["req2"] + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) @pytest.mark.asyncio @@ -2271,7 +2396,7 @@ class TestSpendLogsPayload: "model": "gpt-4o", "user": "", "team_id": "", - "metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "compression_savings": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}', + "metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "guardrail_information": null, "compression_savings": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}', "cache_key": "Cache OFF", "spend": 0.00022500000000000002, "total_tokens": 30, @@ -2367,7 +2492,7 @@ class TestSpendLogsPayload: "model": "claude-4-sonnet-20250514", "user": "", "team_id": "", - "metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "compression_savings": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-4-sonnet-20250514", "model_map_value": {"key": "claude-4-sonnet-20250514", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": {"accepted_prediction_tokens": null, "audio_tokens": null, "reasoning_tokens": null, "rejected_prediction_tokens": null, "text_tokens": 503, "image_tokens": null}, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}', + "metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "guardrail_information": null, "compression_savings": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-4-sonnet-20250514", "model_map_value": {"key": "claude-4-sonnet-20250514", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": {"accepted_prediction_tokens": null, "audio_tokens": null, "reasoning_tokens": null, "rejected_prediction_tokens": null, "text_tokens": 503, "image_tokens": null}, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}', "cache_key": "Cache OFF", "spend": 0.01383, "total_tokens": 2598, @@ -2461,7 +2586,7 @@ class TestSpendLogsPayload: "model": "claude-4-sonnet-20250514", "user": "", "team_id": "", - "metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "compression_savings": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-4-sonnet-20250514", "model_map_value": {"key": "claude-4-sonnet-20250514", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": {"accepted_prediction_tokens": null, "audio_tokens": null, "reasoning_tokens": null, "rejected_prediction_tokens": null, "text_tokens": 503, "image_tokens": null}, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}', + "metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "guardrail_information": null, "compression_savings": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-4-sonnet-20250514", "model_map_value": {"key": "claude-4-sonnet-20250514", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": {"accepted_prediction_tokens": null, "audio_tokens": null, "reasoning_tokens": null, "rejected_prediction_tokens": null, "text_tokens": 503, "image_tokens": null}, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}', "cache_key": "Cache OFF", "spend": 0.01383, "total_tokens": 2598, diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index cc1e2943c8f..c6f2a6f1792 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -2902,3 +2902,17 @@ async def test_compression_savings_survive_to_spend_log_payload_metadata(monkeyp "tokens_saved": 7000, "source": "compression_interception", } + + +def test_no_routing_decision_key_defaults_to_none_in_spend_log_metadata(): + payload = get_logging_payload( + kwargs={ + "model": "gpt-4o-mini", + "litellm_params": {"metadata": {"user_api_key": "test-key"}}, + }, + response_obj=litellm.ModelResponse(id="chatcmpl-no-routing-decision", choices=[], usage=litellm.Usage()), + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), + ) + metadata = json.loads(payload["metadata"]) + assert metadata["routing_decision"] is None diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index e6ccbc579b5..1a584423fac 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -611,12 +611,12 @@ async def test_should_reserve_team_member_and_org_budget_counters(spend_counter_ @pytest.mark.asyncio -async def test_should_reserve_user_budget_counter_for_team_key(spend_counter_state): - """A user's personal budget must be reserved even when the key belongs to a team. +async def test_should_not_reserve_user_budget_counter_for_team_key(spend_counter_state): + """The reservation path mirrors the read path: no personal user counter for a team key. - Regression for GitHub issue #12905: previously the reservation path skipped the - user spend counter whenever the key had a team, so a team key could overshoot the - user's personal max_budget under concurrency. + A team-scoped key reserves against the key and team counters only, so the key + owner's personal max_budget never gates a team request. Fails if the user + counter is reserved for team keys again. """ counter_cache, key_cache = spend_counter_state proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) @@ -645,44 +645,7 @@ async def test_should_reserve_user_budget_counter_for_team_key(spend_counter_sta proxy_logging_obj=proxy_logging_obj, ) - assert counter_cache.in_memory_cache.get_cache(key="spend:user:user-on-team") == pytest.approx(0.3) - - await release_budget_reservation(reservation) - - -@pytest.mark.asyncio -async def test_should_skip_user_budget_counter_for_team_key_when_flag_set(spend_counter_state): - """skip_user_budget_on_team_key=True restores the legacy behavior where a user's - personal budget is not reserved for a team key.""" - counter_cache, key_cache = spend_counter_state - proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) - valid_token = UserAPIKeyAuth( - token="key-user-on-team-skip", - spend=0.0, - user_id="user-on-team-skip", - team_id="team-no-budget-skip", - ) - team_object = LiteLLM_TeamTable(team_id="team-no-budget-skip", spend=0.0, max_budget=None) - user_object = LiteLLM_UserTable(user_id="user-on-team-skip", spend=0.0, max_budget=5.0) - - with patch( - "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", - return_value=0.3, - ): - reservation = await reserve_budget_for_request( - request_body=_request_body(), - route="/chat/completions", - llm_router=None, - valid_token=valid_token, - team_object=team_object, - user_object=user_object, - prisma_client=None, - user_api_key_cache=key_cache, - proxy_logging_obj=proxy_logging_obj, - skip_user_budget_on_team_key=True, - ) - - assert counter_cache.in_memory_cache.get_cache(key="spend:user:user-on-team-skip") is None + assert counter_cache.in_memory_cache.get_cache(key="spend:user:user-on-team") is None await release_budget_reservation(reservation) diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 143884c8a0d..e8acd7e6b75 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -671,6 +671,7 @@ async def test_add_litellm_data_to_request_strips_user_control_fields(): "applied_guardrails": ["spoofed"], "applied_policies": ["spoofed-policy"], "policy_sources": {"spoofed-policy": "request"}, + "routing_decision": {"cause": "forged", "routed_model": "spoofed"}, "_guardrail_pipelines": [{"name": "spoofed"}], "_pipeline_managed_guardrails": ["evaded"], "safe_user_metadata": "kept", @@ -681,6 +682,7 @@ async def test_add_litellm_data_to_request_strips_user_control_fields(): "mock_response": "free response", "mock_tool_calls": [{"id": "call_1"}], "disable_global_guardrails": True, + "routing_decision": {"cause": "forged", "routed_model": "spoofed"}, "metadata": copy.deepcopy(malicious_metadata), "litellm_metadata": copy.deepcopy(malicious_metadata), } @@ -697,6 +699,7 @@ async def test_add_litellm_data_to_request_strips_user_control_fields(): assert "mock_response" not in updated assert "mock_tool_calls" not in updated assert "disable_global_guardrails" not in updated + assert "routing_decision" not in updated stripped_keys = { "disable_global_guardrails", @@ -710,6 +713,7 @@ async def test_add_litellm_data_to_request_strips_user_control_fields(): "applied_guardrails", "applied_policies", "policy_sources", + "routing_decision", "_guardrail_pipelines", "_pipeline_managed_guardrails", } @@ -5457,3 +5461,94 @@ def test_get_sanitized_user_information_from_key_drops_callback_config(): # UserAPIKeyAuth is the live auth object; the per-key callbacks are resolved # from it during pre-call, so it must not be mutated by building the log view assert "logging" in (user_api_key_dict.metadata or {}) + + +def test_team_alias_targeting_deleted_team_deployment_keeps_requested_model(monkeypatch): + """ + Regression: a team's model_aliases can point at the internal routing key + (model_name_{team_id}_{uuid}) of a team deployment that was since deleted, + e.g. after an admin replaces per-team duplicates with one gateway-level + model. Rewriting to the dead internal name made every request fail with + "no healthy deployments for model_name_..." even though the requested + public name resolves at the gateway level. The rewrite must be skipped + when the alias target has no live deployment. + """ + import litellm.proxy.litellm_pre_call_utils as pre_call_utils + from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists + + monkeypatch.delenv("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", raising=False) + pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None + + class _MockRouter: + model_name_to_deployment_indices = {"gpt-4": [0]} + team_model_to_deployment_indices = {} + + test_data = {"model": "gpt-4"} + user_api_key_dict = UserAPIKeyAuth( + api_key="test_key", + team_id="team-1", + team_model_aliases={"gpt-4": "model_name_team-1_dead-uuid"}, + ) + + with patch("litellm.proxy.proxy_server.llm_router", _MockRouter()): + _update_model_if_team_alias_exists( + data=test_data, user_api_key_dict=user_api_key_dict + ) + + assert test_data.get("model") == "gpt-4" + + +def test_team_alias_targeting_live_team_deployment_still_rewrites(monkeypatch): + import litellm.proxy.litellm_pre_call_utils as pre_call_utils + from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists + + monkeypatch.delenv("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", raising=False) + pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None + + class _MockRouter: + model_name_to_deployment_indices = {"model_name_team-1_live-uuid": [0]} + team_model_to_deployment_indices = {} + + test_data = {"model": "gpt-4"} + user_api_key_dict = UserAPIKeyAuth( + api_key="test_key", + team_id="team-1", + team_model_aliases={"gpt-4": "model_name_team-1_live-uuid"}, + ) + + with patch("litellm.proxy.proxy_server.llm_router", _MockRouter()): + _update_model_if_team_alias_exists( + data=test_data, user_api_key_dict=user_api_key_dict + ) + + assert test_data.get("model") == "model_name_team-1_live-uuid" + + +def test_warn_stale_team_alias_once_logs_once_per_key(monkeypatch): + from collections import OrderedDict + + import litellm.proxy.litellm_pre_call_utils as pre_call_utils + + monkeypatch.setattr(pre_call_utils, "_STALE_TEAM_ALIAS_WARNING_KEYS", OrderedDict()) + + with patch.object(pre_call_utils.verbose_proxy_logger, "warning") as mock_warning: + pre_call_utils._warn_stale_team_alias_once("team-1:gpt-4", "stale alias %s", "gpt-4") + pre_call_utils._warn_stale_team_alias_once("team-1:gpt-4", "stale alias %s", "gpt-4") + + assert mock_warning.call_count == 1 + + +def test_warn_stale_team_alias_once_evicts_oldest_key_beyond_cap(monkeypatch): + from collections import OrderedDict + + import litellm.proxy.litellm_pre_call_utils as pre_call_utils + + monkeypatch.setattr(pre_call_utils, "_STALE_TEAM_ALIAS_WARNING_KEYS", OrderedDict()) + monkeypatch.setattr(pre_call_utils, "_MAX_STALE_ALIAS_WARNING_KEYS", 2) + + with patch.object(pre_call_utils.verbose_proxy_logger, "warning"): + pre_call_utils._warn_stale_team_alias_once("key-1", "stale alias") + pre_call_utils._warn_stale_team_alias_once("key-2", "stale alias") + pre_call_utils._warn_stale_team_alias_once("key-3", "stale alias") + + assert list(pre_call_utils._STALE_TEAM_ALIAS_WARNING_KEYS) == ["key-2", "key-3"] diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 536d24d4b4e..5646d202e31 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -1571,14 +1571,14 @@ async def test_get_all_team_models(): with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class: # Configure the mock class to return proper instances - def mock_team_table_constructor(**kwargs): + def mock_team_table_constructor(data): mock_instance = MagicMock() - mock_instance.team_id = kwargs["team_id"] - mock_instance.models = kwargs["models"] - mock_instance.access_group_ids = kwargs.get("access_group_ids") + mock_instance.team_id = data["team_id"] + mock_instance.models = data["models"] + mock_instance.access_group_ids = data.get("access_group_ids") return mock_instance - mock_team_table_class.side_effect = mock_team_table_constructor + mock_team_table_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams="*", @@ -1607,7 +1607,7 @@ async def test_get_all_team_models(): mock_litellm_teamtable.find_many.return_value = [mock_team1] with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class: - mock_team_table_class.side_effect = mock_team_table_constructor + mock_team_table_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams=["team1"], @@ -1658,7 +1658,7 @@ async def test_get_all_team_models(): mock_router.get_model_list.side_effect = mock_get_model_list_with_none with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class: - mock_team_table_class.side_effect = mock_team_table_constructor + mock_team_table_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams=["team1"], @@ -2373,14 +2373,14 @@ async def test_get_all_team_models_with_access_groups(): with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_tt_class: - def mock_team_table_constructor(**kwargs): + def mock_team_table_constructor(data): mock_instance = MagicMock() - mock_instance.team_id = kwargs["team_id"] - mock_instance.models = kwargs["models"] - mock_instance.access_group_ids = kwargs.get("access_group_ids") + mock_instance.team_id = data["team_id"] + mock_instance.models = data["models"] + mock_instance.access_group_ids = data.get("access_group_ids") return mock_instance - mock_tt_class.side_effect = mock_team_table_constructor + mock_tt_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams=["team1"], @@ -9201,38 +9201,6 @@ def test_get_config_list_includes_cancel_on_disconnect(monkeypatch): app.dependency_overrides.clear() -def test_get_config_list_includes_skip_user_budget_on_team_key(monkeypatch): - """Related to #12905: the opt-out flag must be discoverable via /config/list so - it renders as a Boolean toggle on the Admin UI General Settings table. This - requires both the ConfigGeneralSettings field and the allowed_args entry.""" - import types - from unittest.mock import AsyncMock, MagicMock - - from fastapi.testclient import TestClient - - import litellm.proxy.proxy_server as ps - from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.proxy_server import app - - mock_prisma = MagicMock() - mock_config_table = MagicMock() - mock_config_table.find_first = AsyncMock(return_value=None) - mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table) - monkeypatch.setattr(ps, "prisma_client", mock_prisma) - app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) - try: - client = TestClient(app) - resp = client.get("/config/list", params={"config_type": "general_settings"}) - assert resp.status_code == 200, resp.text - fields = {item["field_name"]: item for item in resp.json()} - assert "skip_user_budget_on_team_key" in fields - assert fields["skip_user_budget_on_team_key"]["field_type"] == "Boolean" - finally: - app.dependency_overrides.clear() - - def test_get_config_list_includes_budget_exceeded_throttle_percentage(monkeypatch): """The throttle fraction is a litellm_settings scalar surfaced on the General Settings table as a Float field so it sits with the other global limits; it diff --git a/tests/test_litellm/proxy/test_proxy_types.py b/tests/test_litellm/proxy/test_proxy_types.py index c8e0b3a730a..5354de182a0 100644 --- a/tests/test_litellm/proxy/test_proxy_types.py +++ b/tests/test_litellm/proxy/test_proxy_types.py @@ -139,3 +139,23 @@ def test_key_request_router_settings_keeps_enable_tag_filtering(): dumped = req.router_settings.model_dump(exclude_none=True) assert dumped["enable_tag_filtering"] is True assert dumped["num_retries"] == 2 + + +def test_update_key_request_requires_key_or_key_alias(): + """``/key/update`` can be addressed by ``key`` or by ``key_alias``; + a request with neither has no way to identify the target key and must + fail validation before hitting the endpoint.""" + import pydantic + + from litellm.proxy._types import UpdateKeyRequest + + with pytest.raises(pydantic.ValidationError, match="either key or key_alias must be provided"): + UpdateKeyRequest(max_budget=10.0) + + by_key = UpdateKeyRequest(key="sk-1234") + assert by_key.key == "sk-1234" + assert by_key.key_alias is None + + by_alias = UpdateKeyRequest(key_alias="my-alias") + assert by_alias.key is None + assert by_alias.key_alias == "my-alias" diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 4673807a135..3421751d962 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -7,8 +7,10 @@ import pytest from fastapi import HTTPException from litellm.caching.caching import DualCache +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._types import ProxyErrorTypes from litellm.proxy.utils import ProxyLogging +from litellm.types.guardrails import GuardrailEventHooks sys.path.insert( 0, os.path.abspath("../../..") @@ -947,3 +949,139 @@ class TestSendEmailStartTls: assert isinstance(context, ssl.SSLContext) assert context.verify_mode == ssl.CERT_REQUIRED assert context.check_hostname is True + + +class _RecordingMCPGuardrail(CustomGuardrail): + """Unified guardrail that masks every text it is handed.""" + + def __init__(self, event_hook, masked_text="", raises=None): + super().__init__(guardrail_name="mcp-output-guardrail", event_hook=event_hook, default_on=True) + self.masked_text = masked_text + self.raises = raises + self.call_count = 0 + self.last_input_type = None + + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + self.call_count += 1 + self.last_input_type = input_type + if self.raises is not None: + raise self.raises + return {"texts": [self.masked_text for _ in inputs.get("texts", [])]} + + +class _NativeMCPGuardrail(CustomGuardrail): + """Guardrail that only implements the MCP logging hook (cisco-style).""" + + def __init__(self): + super().__init__( + guardrail_name="native-mcp-guardrail", + event_hook=GuardrailEventHooks.post_mcp_call, + default_on=True, + ) + self.considered_count = 0 + + def should_run_guardrail(self, data, event_type): + self.considered_count += 1 + return super().should_run_guardrail(data=data, event_type=event_type) + + async def async_post_mcp_tool_call_hook(self, kwargs, response_obj, start_time, end_time): + return None + + +@pytest.fixture +def restore_callbacks(): + """Restore the process-wide callback state post_mcp_call_hook reads. + + ProxyLogging caches callback capabilities keyed on id()s of litellm.callbacks, + so a restored-but-different list can collide with a stale entry after GC and + leak a has_guardrail verdict into unrelated tests in the same worker. + """ + original = list(litellm.callbacks) + yield + litellm.callbacks = original + ProxyLogging._callback_capabilities_cache.clear() + + +@pytest.mark.asyncio +async def test_post_mcp_call_hook_masks_tool_result(restore_callbacks): + """A post_mcp_call guardrail must see the tool result text and mask it in the returned result.""" + from mcp.types import CallToolResult, TextContent + + guardrail = _RecordingMCPGuardrail(event_hook=GuardrailEventHooks.post_mcp_call) + litellm.callbacks = [guardrail] + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False) + + returned = await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data={"mcp_tool_name": "echo"}, + user_api_key_dict=None, + ) + + assert guardrail.call_count == 1 + assert guardrail.last_input_type == "response" + assert [item.text for item in returned.content] == [""] + + +@pytest.mark.asyncio +async def test_post_mcp_call_hook_skips_guardrail_configured_for_other_hooks(restore_callbacks): + """A guardrail not configured for post_mcp_call must not scan MCP tool results.""" + from mcp.types import CallToolResult, TextContent + + guardrail = _RecordingMCPGuardrail(event_hook=GuardrailEventHooks.post_call) + litellm.callbacks = [guardrail] + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False) + + returned = await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data={"mcp_tool_name": "echo"}, + user_api_key_dict=None, + ) + + assert guardrail.call_count == 0 + assert [item.text for item in returned.content] == ["jane@example.com"] + + +@pytest.mark.asyncio +async def test_post_mcp_call_hook_skips_guardrail_without_apply_guardrail(restore_callbacks): + """Guardrails that implement async_post_mcp_tool_call_hook are dispatched by the + logging object, so this hook must not run them a second time.""" + from mcp.types import CallToolResult, TextContent + + guardrail = _NativeMCPGuardrail() + litellm.callbacks = [guardrail] + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False) + + returned = await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data={"mcp_tool_name": "echo"}, + user_api_key_dict=None, + ) + + assert guardrail.considered_count == 0 + assert [item.text for item in returned.content] == ["jane@example.com"] + + +@pytest.mark.asyncio +async def test_post_mcp_call_hook_propagates_guardrail_block(restore_callbacks): + """A guardrail rejecting the tool result must raise out of the hook.""" + from mcp.types import CallToolResult, TextContent + + from litellm.exceptions import BlockedPiiEntityError + + guardrail = _RecordingMCPGuardrail( + event_hook=GuardrailEventHooks.post_mcp_call, + raises=BlockedPiiEntityError(entity_type="EMAIL_ADDRESS", guardrail_name="mcp-output-guardrail"), + ) + litellm.callbacks = [guardrail] + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False) + + with pytest.raises(BlockedPiiEntityError): + await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data={"mcp_tool_name": "echo"}, + user_api_key_dict=None, + ) diff --git a/tests/test_litellm/repositories/test_repositories.py b/tests/test_litellm/repositories/test_repositories.py index 6308faf8fc7..c923b722991 100644 --- a/tests/test_litellm/repositories/test_repositories.py +++ b/tests/test_litellm/repositories/test_repositories.py @@ -203,23 +203,32 @@ class TestBaseRepository: assert len(budgets) == 1 def test_record_to_dict_branches(self): - from litellm.repositories.base_repository import _record_to_dict + from litellm.repositories.base_repository import record_to_dict - assert _record_to_dict({"a": 1}) == {"a": 1} + assert record_to_dict({"a": 1}) == {"a": 1} class WithModelDump: def model_dump(self): return {"src": "model_dump"} - assert _record_to_dict(WithModelDump()) == {"src": "model_dump"} + assert record_to_dict(WithModelDump()) == {"src": "model_dump"} class WithDict: def dict(self): return {"src": "dict"} - assert _record_to_dict(WithDict()) == {"src": "dict"} + assert record_to_dict(WithDict()) == {"src": "dict"} - assert _record_to_dict([("k", "v")]) == {"k": "v"} + assert record_to_dict([("k", "v")]) == {"k": "v"} + + class WithBoth: + def model_dump(self): + return {"src": "model_dump"} + + def dict(self): + return {"src": "dict"} + + assert record_to_dict(WithBoth()) == {"src": "model_dump"} class TestBudgetRepository: diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_async_pre_routing.py b/tests/test_litellm/router_strategy/adaptive_router/test_async_pre_routing.py index fb43cf403d6..0a8c230e773 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_async_pre_routing.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_async_pre_routing.py @@ -240,3 +240,25 @@ async def test_invalid_min_quality_tier_header_treated_as_none(): assert ( r.pick_model.await_args.kwargs["min_quality_tier"] is None # type: ignore[union-attr] ) + + +@pytest.mark.asyncio +async def test_routing_decision_reports_bandit_choice(): + r = _make_router() + r.pick_model = AsyncMock(return_value="smart") # type: ignore[method-assign] + + response = await r.async_pre_routing_hook( + model="smart-cheap-router", + request_kwargs={}, + messages=[{"role": "user", "content": "Write a Python function"}], + ) + + assert response is not None + decision = response.routing_decision + assert decision is not None + assert decision["router_model_name"] == "smart-cheap-router" + assert decision["router_type"] == "adaptive" + assert decision["cause"] == "bandit" + assert decision["routed_model"] == "smart" + assert decision["request_type"] == RequestType.CODE_GENERATION.value + assert "tier" not in decision diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index ef70687bd97..e734d8ec876 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -24,6 +24,7 @@ from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY from litellm.router_strategy.complexity_router.complexity_router import ( ComplexityRouter, DimensionScore, + KeywordOverride, ) from litellm.router_strategy.complexity_router.config import ( DEFAULT_COMPLEXITY_CONFIG, @@ -1381,21 +1382,27 @@ class TestLLMClassifier: async def test_aclassify_heuristic_skips_llm_call(self, complexity_router, mock_router_instance): """When classifier_type is 'heuristic' (default), aclassify must not call the LLM.""" mock_router_instance.acompletion = AsyncMock() - tier, score, signals = await complexity_router.aclassify("Hello!") + outcome = await complexity_router.aclassify("Hello!") mock_router_instance.acompletion.assert_not_called() - assert tier == ComplexityTier.SIMPLE + assert outcome.tier == ComplexityTier.SIMPLE + assert outcome.cause == "heuristic_scorer" + assert outcome.score is not None @pytest.mark.asyncio async def test_aclassify_llm_success_routes_by_llm_verdict(self, llm_complexity_router, mock_router_instance): """A well-formed structured LLM response should decide the tier directly. Uses a prompt that heuristic scoring alone would classify as SIMPLE, to prove - the LLM verdict -- not the heuristic scorer -- is what decided the tier. + the LLM verdict -- not the heuristic scorer -- is what decided the tier. The + outcome must say so (cause) and must not fabricate a score: the LLM path + produces a tier label only. """ mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "COMPLEX"}')) - tier, score, signals = await llm_complexity_router.aclassify("hi") - assert tier == ComplexityTier.COMPLEX - assert "llm-classifier:COMPLEX" in signals + outcome = await llm_complexity_router.aclassify("hi") + assert outcome.tier == ComplexityTier.COMPLEX + assert outcome.cause == "llm_classifier" + assert outcome.score is None + assert "llm-classifier:COMPLEX" in outcome.signals mock_router_instance.acompletion.assert_awaited_once() call_kwargs = mock_router_instance.acompletion.call_args.kwargs assert call_kwargs["model"] == "haiku-classifier" @@ -1417,6 +1424,102 @@ class TestLLMClassifier: call_kwargs = mock_router_instance.acompletion.call_args.kwargs assert call_kwargs["metadata"] == request_metadata + @pytest.mark.asyncio + async def test_aclassify_forwards_metadata_key_used_by_chat_completions( + self, llm_complexity_router, mock_router_instance + ): + """/v1/chat/completions puts the request metadata under "metadata", not "litellm_metadata". + + Only the routes in LITELLM_METADATA_ROUTES (/v1/messages, /v1/responses, ...) get a + "litellm_metadata" bucket; chat completions gets "metadata". Reading only + "litellm_metadata" leaves the classifier call unattributed on the most common route, + so _should_track_cost_callback drops it and no spend-log row is written at all, + which also makes the captured request body unreachable in the Logs UI. + """ + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}')) + request_metadata = {"user_api_key": "sk-abc", "user_api_key_team_id": "team-1"} + await llm_complexity_router.aclassify("hi", request_kwargs={"metadata": request_metadata}) + call_kwargs = mock_router_instance.acompletion.call_args.kwargs + assert call_kwargs["metadata"] == request_metadata + + @pytest.mark.asyncio + async def test_aclassify_captures_request_body_in_proxy_server_request( + self, llm_complexity_router, mock_router_instance + ): + """The classifier call must supply proxy_server_request so its request body is logged. + + proxy_server_request["body"] is populated only by the proxy's HTTP ingress + middleware, which never runs for this internally-initiated router.acompletion + call. Without it _get_proxy_server_request_for_spend_logs_payload reads nothing + and stores "{}" for the request, so the classifier's spend-log row shows a + populated response but an empty request and the log cannot show which prompt + drove the tier decision. The captured body must carry the classification prompt + actually sent, so the classifier model, the classification prompt, and the user + text are all asserted here. + """ + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "COMPLEX"}')) + await llm_complexity_router.aclassify("explain quantum tunneling in depth") + call_kwargs = mock_router_instance.acompletion.call_args.kwargs + body = call_kwargs["proxy_server_request"]["body"] + assert body["model"] == "haiku-classifier" + assert body["messages"] == call_kwargs["messages"] + assert "explain quantum tunneling in depth" in body["messages"][0]["content"] + assert body["response_format"]["type"] == "json_schema" + assert body["response_format"]["json_schema"]["schema"]["properties"]["tier"]["enum"] == [ + "SIMPLE", + "MEDIUM", + "COMPLEX", + "REASONING", + ] + + @pytest.mark.asyncio + async def test_aclassify_propagates_top_level_turn_off_message_logging( + self, llm_complexity_router, mock_router_instance + ): + """A caller's top-level turn_off_message_logging must reach the classifier call. + + Without this, a caller who opts a request out of message logging still has their + prompt captured in full by the classifier's proxy_server_request: the spend-log + redaction gate (should_redact_message_logging) reads turn_off_message_logging off + the classifier call's own kwargs, and this internal call is not the caller's + request, so it never inherits the opt-out unless it's forwarded explicitly. + """ + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}')) + await llm_complexity_router.aclassify("secret prompt", request_kwargs={"turn_off_message_logging": True}) + call_kwargs = mock_router_instance.acompletion.call_args.kwargs + assert call_kwargs["turn_off_message_logging"] is True + + @pytest.mark.asyncio + async def test_aclassify_propagates_metadata_slot_turn_off_message_logging( + self, llm_complexity_router, mock_router_instance + ): + """turn_off_message_logging set inside metadata/litellm_metadata must also propagate. + + initialize_standard_callback_dynamic_params reads this flag from either the + top-level request kwargs or the metadata/litellm_metadata dicts (the same slots a + real HTTP request populates), so the classifier call must resolve it from there too. + """ + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}')) + await llm_complexity_router.aclassify( + "secret prompt", request_kwargs={"litellm_metadata": {"turn_off_message_logging": True}} + ) + call_kwargs = mock_router_instance.acompletion.call_args.kwargs + assert call_kwargs["turn_off_message_logging"] is True + + @pytest.mark.asyncio + async def test_aclassify_defaults_turn_off_message_logging_to_none( + self, llm_complexity_router, mock_router_instance + ): + """With no caller opt-out, the classifier call must not force redaction on or off. + + Passing None (rather than omitting the kwarg or defaulting to False) preserves the + existing header- and global-setting fallbacks in should_redact_message_logging. + """ + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}')) + await llm_complexity_router.aclassify("hi") + call_kwargs = mock_router_instance.acompletion.call_args.kwargs + assert call_kwargs["turn_off_message_logging"] is None + @pytest.mark.asyncio async def test_aclassify_strips_budget_reservation_from_classifier_metadata( self, llm_complexity_router, mock_router_instance @@ -1460,9 +1563,13 @@ class TestLLMClassifier: ): """A timeout/error from the classifier model must fall back to heuristic scoring.""" mock_router_instance.acompletion = AsyncMock(side_effect=TimeoutError("classifier timed out")) - tier, score, signals = await llm_complexity_router.aclassify("Hello!") - assert tier == llm_complexity_router.classify("Hello!")[0] - assert tier == ComplexityTier.SIMPLE + outcome = await llm_complexity_router.aclassify("Hello!") + assert outcome.tier == llm_complexity_router.classify("Hello!")[0] + assert outcome.tier == ComplexityTier.SIMPLE + # The fallback ran the heuristic, and the outcome must say so even though + # the configured classifier_type is "llm". + assert outcome.cause == "heuristic_scorer" + assert outcome.score is not None @pytest.mark.asyncio async def test_aclassify_falls_back_to_heuristic_on_unparseable_response( @@ -1470,8 +1577,9 @@ class TestLLMClassifier: ): """Non-JSON or schema-violating output must fall back to heuristic scoring, not raise.""" mock_router_instance.acompletion = AsyncMock(return_value=_llm_response("not json")) - tier, score, signals = await llm_complexity_router.aclassify("Hello!") - assert tier == ComplexityTier.SIMPLE + outcome = await llm_complexity_router.aclassify("Hello!") + assert outcome.tier == ComplexityTier.SIMPLE + assert outcome.cause == "heuristic_scorer" @pytest.mark.asyncio async def test_aclassify_falls_back_to_heuristic_on_empty_content( @@ -1479,8 +1587,9 @@ class TestLLMClassifier: ): """Empty/None message content (e.g. provider quirk) must fall back, not raise.""" mock_router_instance.acompletion = AsyncMock(return_value=_llm_response(None)) - tier, score, signals = await llm_complexity_router.aclassify("Hello!") - assert tier == ComplexityTier.SIMPLE + outcome = await llm_complexity_router.aclassify("Hello!") + assert outcome.tier == ComplexityTier.SIMPLE + assert outcome.cause == "heuristic_scorer" @pytest.mark.asyncio async def test_pre_routing_hook_uses_llm_classifier_end_to_end(self, llm_complexity_router, mock_router_instance): @@ -1975,8 +2084,12 @@ class TestLexicalKeywordTierRules: litellm_router_instance=mock_router_instance, complexity_router_config=config, ) - assert router._lexical_tier_override("hi there, please advise") == ComplexityTier.COMPLEX - assert router._lexical_tier_override("just saying hi") == ComplexityTier.SIMPLE + assert router._lexical_tier_override("hi there, please advise") == KeywordOverride( + tier=ComplexityTier.COMPLEX, matched_keyword="advise" + ) + assert router._lexical_tier_override("just saying hi") == KeywordOverride( + tier=ComplexityTier.SIMPLE, matched_keyword="hi" + ) assert router._lexical_tier_override("nothing relevant here") is None @pytest.mark.asyncio @@ -2012,7 +2125,9 @@ class TestLexicalKeywordTierRules: litellm_router_instance=mock_router_instance, complexity_router_config=config, ) - assert router._lexical_tier_override("running my k8s cluster") == ComplexityTier.REASONING + assert router._lexical_tier_override("running my k8s cluster") == KeywordOverride( + tier=ComplexityTier.REASONING, matched_keyword="k8s" + ) assert router._lexical_tier_override("what is a k8scluster thing") is None @@ -2169,6 +2284,69 @@ class TestSemanticKeywordTierRules: assert fake_router.async_embedding_kwargs[0]["metadata"] == caller_metadata assert fake_router.async_embedding_kwargs[0]["litellm_metadata"] == caller_litellm_metadata + @pytest.mark.asyncio + async def test_semantic_embedding_call_captures_request_body_in_proxy_server_request(self, basic_config): + """The query embedding call must supply proxy_server_request so its request is logged. + + Like the LLM classifier, this embedding is fired internally and never passes + through the proxy's HTTP ingress middleware, so proxy_server_request is unset and + the embedding's spend-log row stores "{}" for the request while its response is + captured. The captured body must carry the embedded input so the log shows what + was classified. + """ + fake_router = FakeEmbeddingRouter() + config = { + **basic_config, + "keyword_tier_rules": [{"keywords": ["kubernetes deployment"], "tier": "REASONING"}], + "semantic_keyword_matching": True, + "embedding_model": "fake-embed", + "match_threshold": 0.5, + } + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=fake_router, + complexity_router_config=config, + ) + await router.async_pre_routing_hook( + model="test-model", + request_kwargs={}, + messages=[{"role": "user", "content": "roll out my k8s cluster"}], + ) + assert fake_router.async_embedding_kwargs, "expected an embedding call for the prompt" + body = fake_router.async_embedding_kwargs[0]["proxy_server_request"]["body"] + assert body["model"] == "fake-embed" + assert body["input"] == ["roll out my k8s cluster"] + + @pytest.mark.asyncio + async def test_semantic_embedding_call_propagates_turn_off_message_logging(self, basic_config): + """A caller's turn_off_message_logging must reach the query embedding call. + + The embedding now captures the user's prompt in proxy_server_request, so a caller + who opts out of message logging must have that opt-out forwarded; otherwise the + embedding's spend-log row stores the prompt in the clear despite the parent request + being redacted, exposing it to anyone authorized to read the team's spend logs. + """ + fake_router = FakeEmbeddingRouter() + config = { + **basic_config, + "keyword_tier_rules": [{"keywords": ["kubernetes deployment"], "tier": "REASONING"}], + "semantic_keyword_matching": True, + "embedding_model": "fake-embed", + "match_threshold": 0.5, + } + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=fake_router, + complexity_router_config=config, + ) + await router.async_pre_routing_hook( + model="test-model", + request_kwargs={"turn_off_message_logging": True}, + messages=[{"role": "user", "content": "roll out my k8s cluster"}], + ) + assert fake_router.async_embedding_kwargs, "expected an embedding call for the prompt" + assert fake_router.async_embedding_kwargs[0]["turn_off_message_logging"] is True + @pytest.mark.asyncio async def test_semantic_embedding_call_strips_budget_reservation(self, basic_config): """The embedding call must not carry the parent request's budget reservation. @@ -2653,7 +2831,7 @@ class TestRoutingDecisionCauseLogging: request_kwargs={}, messages=[{"role": "user", "content": "What is the boiling point of water at sea level?"}], ) - assert "routing decision cause=complexity_scorer" in router_log_capture.text + assert "routing decision cause=heuristic_scorer" in router_log_capture.text assert "score=" in router_log_capture.text assert "cause=literal_keyword_match" not in router_log_capture.text assert "cause=semantic_keyword_match" not in router_log_capture.text @@ -3157,9 +3335,9 @@ class TestEscalationKeywords: assert complexity_router.escalation_keywords == ["LITELLM ESCALATE"] def test_escalation_triggered_is_case_sensitive(self, complexity_router): - assert complexity_router._escalation_triggered("please LITELLM ESCALATE now") is True - assert complexity_router._escalation_triggered("please litellm escalate now") is False - assert complexity_router._escalation_triggered("how do I escalate this ticket") is False + assert complexity_router._matched_escalation_keyword("please LITELLM ESCALATE now") == "LITELLM ESCALATE" + assert complexity_router._matched_escalation_keyword("please litellm escalate now") is None + assert complexity_router._matched_escalation_keyword("how do I escalate this ticket") is None def test_escalate_tier_bumps_one_step(self, complexity_router): assert complexity_router._escalate_tier(ComplexityTier.SIMPLE) == ComplexityTier.MEDIUM @@ -3400,3 +3578,584 @@ class TestEscalationKeywords: messages=[{"role": "user", "content": "LITELLM ESCALATE do better"}], ) assert result.model == "o1-b" # unchanged: no random hop to o1-a / o1-c + + +class TestRoutingDecisionContents: + """Every routing path must return a PreRoutingHookResponse carrying a routing_decision + that names the mechanism that actually decided, with the facts of that path only.""" + + @pytest.mark.asyncio + async def test_heuristic_decision_carries_score_signals_and_boundary_snapshot(self, complexity_router): + response = await complexity_router.async_pre_routing_hook( + model="test-complexity-router", + request_kwargs={}, + messages=[{"role": "user", "content": "Hello!"}], + ) + assert response is not None + decision = response.routing_decision + assert decision is not None + assert decision["router_model_name"] == "test-complexity-router" + assert decision["router_type"] == "complexity" + assert decision["cause"] == "heuristic_scorer" + assert decision["tier"] == "SIMPLE" + assert decision["routed_model"] == response.model == "gpt-4o-mini" + assert isinstance(decision["score"], float) + assert any("short" in signal for signal in decision["signals"]) + # The snapshot must reflect the CONFIGURED boundaries (the fixture overrides the + # 0.15/0.35/0.60 defaults), so a logged row stays truthful after config edits. + assert decision["tier_boundaries"] == { + "simple_medium": 0.25, + "medium_complex": 0.50, + "complex_reasoning": 0.75, + } + assert "escalated" not in decision + assert "classifier_model" not in decision + + @pytest.mark.asyncio + async def test_llm_classifier_decision_names_judge_and_omits_score( + self, llm_complexity_router, mock_router_instance + ): + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "REASONING"}')) + response = await llm_complexity_router.async_pre_routing_hook( + model="test-complexity-router", + request_kwargs={}, + messages=[{"role": "user", "content": "hi"}], + ) + assert response is not None + decision = response.routing_decision + assert decision is not None + assert decision["cause"] == "llm_classifier" + assert decision["classifier_model"] == "haiku-classifier" + assert decision["tier"] == "REASONING" + # The LLM path produces a tier label, not a score: no synthetic score and no + # boundary snapshot may appear on these rows. + assert "score" not in decision + assert "tier_boundaries" not in decision + + @pytest.mark.asyncio + async def test_llm_classifier_fallback_decision_reports_heuristic( + self, llm_complexity_router, mock_router_instance + ): + """A failed LLM classifier falls back to the heuristic, and the persisted cause + must say heuristic_scorer even though classifier_type is 'llm'.""" + mock_router_instance.acompletion = AsyncMock(side_effect=TimeoutError("classifier timed out")) + response = await llm_complexity_router.async_pre_routing_hook( + model="test-complexity-router", + request_kwargs={}, + messages=[{"role": "user", "content": "Hello!"}], + ) + assert response is not None + decision = response.routing_decision + assert decision is not None + assert decision["cause"] == "heuristic_scorer" + assert "classifier_model" not in decision + assert isinstance(decision["score"], float) + + @pytest.mark.asyncio + async def test_keyword_override_decision_carries_matched_keyword(self, mock_router_instance, basic_config): + config = { + **basic_config, + "keyword_tier_rules": [{"keywords": ["deploy to k8s"], "tier": "REASONING"}], + } + router = ComplexityRouter( + model_name="test-complexity-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + response = await router.async_pre_routing_hook( + model="test-complexity-router", + request_kwargs={}, + messages=[{"role": "user", "content": "please deploy to k8s now"}], + ) + assert response is not None + decision = response.routing_decision + assert decision is not None + assert decision["cause"] == "literal_keyword_match" + assert decision["matched_keyword"] == "deploy to k8s" + assert decision["tier"] == "REASONING" + assert "score" not in decision + + @pytest.mark.asyncio + async def test_no_user_message_decision_is_default_fallback(self, complexity_router): + response = await complexity_router.async_pre_routing_hook( + model="test-complexity-router", + request_kwargs={}, + messages=[{"role": "system", "content": "be nice"}], + ) + assert response is not None + decision = response.routing_decision + assert decision is not None + assert decision["cause"] == "default_fallback" + assert decision["routed_model"] == response.model + assert "tier" not in decision + + @pytest.mark.asyncio + async def test_session_pin_decision(self, mock_router_instance, basic_config): + mock_router_instance.cache = DualCache() + router = ComplexityRouter( + model_name="test-complexity-router", + litellm_router_instance=mock_router_instance, + complexity_router_config={**basic_config, "session_affinity": True}, + ) + request_kwargs = {"metadata": {"session_id": "session-decision"}} + cache_key = router._get_session_affinity_cache_key("session-decision", request_kwargs) + await mock_router_instance.cache.async_set_cache(key=cache_key, value="gpt-4o") + response = await router.async_pre_routing_hook( + model="test-complexity-router", + request_kwargs=request_kwargs, + messages=[{"role": "user", "content": "hi again"}], + ) + assert response is not None + decision = response.routing_decision + assert decision is not None + assert decision["cause"] == "session_affinity_pin" + assert decision["routed_model"] == "gpt-4o" + assert "escalated" not in decision + + @pytest.mark.asyncio + async def test_reasoning_override_is_its_own_cause(self, complexity_router): + """The override is the fact that the score did NOT choose the tier, so it is a + cause rather than a marker inside `signals`; anything that filters signals would + otherwise change what the row claims.""" + response = await complexity_router.async_pre_routing_hook( + model="test-complexity-router", + request_kwargs={}, + messages=[{"role": "user", "content": "Let's think step by step and prove the theorem."}], + ) + decision = response.routing_decision + assert decision["tier"] == "REASONING" + assert decision["cause"] == "reasoning_override" + # The score is still recorded, but the cause is what says it did not decide. + assert decision["score"] < decision["tier_boundaries"]["complex_reasoning"] + + +class TestSignalsNeverQuoteTheSystemPrompt: + """Signals are persisted to the caller-readable spend log, so they may name a matched + term only when the caller supplied it. A term matched solely in the system prompt is + reported as a count, which still explains the score without letting a caller recover + configured terms from a prompt it cannot see.""" + + @pytest.mark.asyncio + async def test_system_prompt_only_terms_are_reported_as_a_count(self, complexity_router): + response = await complexity_router.async_pre_routing_hook( + model="test-complexity-router", + request_kwargs={}, + messages=[ + {"role": "system", "content": "You operate the kubernetes database api for the deployment pipeline."}, + {"role": "user", "content": "say hi"}, + ], + ) + assert response is not None + signals = response.routing_decision["signals"] + joined = " ".join(signals) + # The system prompt drove these matches, so no signal may name them. + for term in ("kubernetes", "database", "api", "deployment"): + assert term not in joined + # The match is still reported, as a count, so the score stays explainable. + assert any("matches" in signal for signal in signals) + + @pytest.mark.asyncio + async def test_terms_the_caller_supplied_are_still_named(self, complexity_router): + response = await complexity_router.async_pre_routing_hook( + model="test-complexity-router", + request_kwargs={}, + messages=[ + {"role": "system", "content": "You operate the kubernetes cluster."}, + {"role": "user", "content": "help me debug the database api timeout in production"}, + ], + ) + assert response is not None + signals = " ".join(response.routing_decision["signals"]) + # The caller typed these, so quoting them discloses nothing. + assert "database" in signals or "api" in signals + # It did not type this one. + assert "kubernetes" not in signals + + def test_scoring_still_reads_the_system_prompt(self, complexity_router): + """Redaction is a disclosure rule, not a scoring change: the system prompt must + still count toward the tier exactly as before.""" + with_system = complexity_router.classify( + "say hi", "You operate the kubernetes database api for the deployment pipeline." + ) + without_system = complexity_router.classify("say hi") + assert with_system[1] > without_system[1] + + +class TestRoutingDecisionSurvivesToSpendLogOnEveryMetadataShape: + """The decision must reach the spend-log row on every request surface. + + `/v1/chat/completions` carries proxy state in `metadata`; `/v1/messages` and the + batch-style routes carry it in `litellm_metadata` (so the provider's own `metadata` + field stays untouched), and a caller may supply either, both, or neither. Logging + snapshots `litellm_metadata` by value (`function_setup`, litellm/utils.py), so a + stash written to the wrong bucket, or read after a copy, is dropped silently and + only on the surfaces nobody exercised. This drives the real hook and then the real + spend-log payload builder for every shape. + """ + + MODEL_LIST = [ + { + "model_name": "smart-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": ["gpt-4o-mini"], "MEDIUM": ["gpt-4o"]}, + "session_affinity": False, + }, + }, + }, + {"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini"}}, + {"model_name": "gpt-4o", "litellm_params": {"model": "openai/gpt-4o"}}, + ] + + @pytest.mark.parametrize( + "request_kwargs, expected_bucket", + [ + pytest.param({}, "metadata", id="no-caller-metadata"), + pytest.param({"metadata": {"caller_tag": "x"}}, "metadata", id="caller-metadata"), + pytest.param({"litellm_metadata": {}}, "litellm_metadata", id="litellm-metadata-seeded"), + pytest.param( + {"litellm_metadata": {"caller_tag": "x"}}, "litellm_metadata", id="litellm-metadata-with-caller-value" + ), + pytest.param( + {"litellm_metadata": {}, "metadata": {"user_id": "end-user-1"}}, + "litellm_metadata", + id="both-buckets", + ), + ], + ) + @pytest.mark.asyncio + async def test_decision_reaches_the_spend_log_payload(self, request_kwargs, expected_bucket): + import datetime + import json + + from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload + + router = Router(model_list=self.MODEL_LIST) + response = await router.async_pre_routing_hook( + model="smart-router", + request_kwargs=request_kwargs, + messages=[{"role": "user", "content": "Hello!"}], + ) + assert response is not None + assert "routing_decision" in request_kwargs[expected_bucket] + if expected_bucket == "litellm_metadata" and isinstance(request_kwargs.get("metadata"), dict): + # On these routes `metadata` is the provider's own field, forwarded upstream. + assert "routing_decision" not in request_kwargs["metadata"] + + # Mirror function_setup: it copies `litellm_metadata` by value into + # litellm_params AFTER the router hook has run, so the copy must carry + # the decision. Reading the stash any earlier would lose it. + litellm_params: Dict = {} + if "metadata" in request_kwargs: + litellm_params["metadata"] = request_kwargs["metadata"] + if isinstance(request_kwargs.get("litellm_metadata"), dict): + litellm_params["litellm_metadata"] = request_kwargs["litellm_metadata"].copy() + + payload = get_logging_payload( + kwargs={"model": "gpt-4o-mini", "litellm_params": litellm_params}, + response_obj=litellm.ModelResponse(id="chatcmpl-shape", choices=[], usage=litellm.Usage()), + start_time=datetime.datetime.now(datetime.timezone.utc), + end_time=datetime.datetime.now(datetime.timezone.utc), + ) + persisted = json.loads(payload["metadata"])["routing_decision"] + assert persisted is not None, f"routing_decision dropped for {expected_bucket}" + assert persisted["router_model_name"] == "smart-router" + + +class TestRoutingDecisionIsPerAttempt: + """The stash must describe the attempt that actually served the request. + + Fallbacks re-enter `async_pre_routing_hook` with the SAME request_kwargs, so a + decision left behind by a failed auto-router attempt would be attributed to the + plain model group that served the retry, making the spend row claim a tier the + request never used. The bucket is also resolved through the shared owner, so a + non-dict value in the bucket slot is replaced rather than silently skipped. + """ + + MODEL_LIST = [ + { + "model_name": "smart-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": ["gpt-4o-mini"], "MEDIUM": ["gpt-4o"]}, + "session_affinity": False, + }, + }, + }, + {"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini"}}, + {"model_name": "gpt-4o", "litellm_params": {"model": "openai/gpt-4o"}}, + ] + + @pytest.mark.parametrize( + "seed, bucket", [({}, "metadata"), ({"litellm_metadata": {}}, "litellm_metadata")] + ) + @pytest.mark.asyncio + async def test_fallback_to_plain_model_group_clears_the_earlier_decision(self, seed, bucket): + router = Router(model_list=self.MODEL_LIST) + request_kwargs: Dict = dict(seed) + messages = [{"role": "user", "content": "Hello!"}] + + await router.async_pre_routing_hook( + model="smart-router", request_kwargs=request_kwargs, messages=messages + ) + assert "routing_decision" in request_kwargs[bucket] + + # The fallback attempt reuses the same kwargs and selects no strategy. + response = await router.async_pre_routing_hook( + model="gpt-4o-mini", request_kwargs=request_kwargs, messages=messages + ) + assert response is None + assert "routing_decision" not in request_kwargs[bucket] + + @pytest.mark.parametrize("unusable_bucket", [None, "not-a-dict"]) + @pytest.mark.asyncio + async def test_non_dict_bucket_is_replaced_not_skipped(self, unusable_bucket): + """A caller can send `litellm_metadata` as a non-dict (unparsed string, null). + Skipping the write there would drop provenance on a successfully routed + request with no error, so the shared bucket owner replaces the value.""" + router = Router(model_list=self.MODEL_LIST) + request_kwargs: Dict = {"litellm_metadata": unusable_bucket} + + response = await router.async_pre_routing_hook( + model="smart-router", + request_kwargs=request_kwargs, + messages=[{"role": "user", "content": "Hello!"}], + ) + + assert response is not None + bucket = request_kwargs["litellm_metadata"] + assert isinstance(bucket, dict) + assert bucket["routing_decision"]["router_model_name"] == "smart-router" + + +class TestRecordRoutingDecision: + """Direct coverage of the single recording point, whose contract is write-or-clear: + the request's metadata must describe the current attempt and nothing else.""" + + DECISION = {"router_model_name": "smart-router", "router_type": "complexity", "routed_model": "gpt-4o-mini"} + + def test_none_clears_a_previous_decision_from_both_buckets(self): + request_kwargs: Dict = { + "metadata": {"routing_decision": self.DECISION, "keep": 1}, + "litellm_metadata": {"routing_decision": self.DECISION}, + } + Router._record_routing_decision(request_kwargs=request_kwargs, routing_decision=None) + assert "routing_decision" not in request_kwargs["metadata"] + assert "routing_decision" not in request_kwargs["litellm_metadata"] + assert request_kwargs["metadata"]["keep"] == 1 + + def test_none_creates_no_bucket_on_a_request_that_had_none(self): + request_kwargs: Dict = {} + Router._record_routing_decision(request_kwargs=request_kwargs, routing_decision=None) + assert request_kwargs == {} + + +class TestEscalationIsRecordedConsistently: + """An escalation keyword records two separate facts on every path: that the caller + asked, and whether the tier actually moved. Dropping the ask when there is nowhere + higher to go makes a request look like an ordinary route, and reporting a bump that + never happened is the opposite error; both must be avoided identically everywhere.""" + + CEILING_CONFIG = { + "tiers": {"SIMPLE": ["gpt-4o-mini"], "REASONING": ["o1-preview"]}, + "session_affinity": False, + } + + @pytest.mark.asyncio + async def test_scorer_path_at_ceiling_keeps_the_keyword_and_reports_no_bump(self, mock_router_instance): + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config={ + **self.CEILING_CONFIG, + "tier_boundaries": {"simple_medium": -99, "medium_complex": -99, "complex_reasoning": -99}, + }, + ) + response = await router.async_pre_routing_hook( + model="test-router", + request_kwargs={}, + messages=[{"role": "user", "content": "LITELLM ESCALATE already at the top"}], + ) + decision = response.routing_decision + assert decision["tier"] == "REASONING" + assert decision["escalation_keyword"] == "LITELLM ESCALATE" + assert decision["escalated"] is False + + @pytest.mark.asyncio + async def test_scorer_path_below_ceiling_reports_the_bump(self, complexity_router): + response = await complexity_router.async_pre_routing_hook( + model="test-router", + request_kwargs={}, + messages=[{"role": "user", "content": "LITELLM ESCALATE what is 2+2"}], + ) + decision = response.routing_decision + assert decision["escalation_keyword"] == "LITELLM ESCALATE" + assert decision["escalated"] is True + + @pytest.mark.asyncio + async def test_session_pin_at_ceiling_still_records_the_ask(self, mock_router_instance): + mock_router_instance.cache = DualCache() + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config={**self.CEILING_CONFIG, "session_affinity": True}, + ) + request_kwargs = {"metadata": {"session_id": "session-ceiling"}} + cache_key = router._get_session_affinity_cache_key("session-ceiling", request_kwargs) + await mock_router_instance.cache.async_set_cache(key=cache_key, value="o1-preview") + + response = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=request_kwargs, + messages=[{"role": "user", "content": "LITELLM ESCALATE go higher"}], + ) + decision = response.routing_decision + assert decision["routed_model"] == "o1-preview" + assert decision["cause"] == "session_affinity_pin" + # Previously the keyword was dropped here, so the row was indistinguishable + # from a turn that never asked to escalate. + assert decision["escalation_keyword"] == "LITELLM ESCALATE" + assert decision["escalated"] is False + + @pytest.mark.asyncio + async def test_session_pin_below_ceiling_reports_the_bump(self, mock_router_instance): + mock_router_instance.cache = DualCache() + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config={**self.CEILING_CONFIG, "session_affinity": True}, + ) + request_kwargs = {"metadata": {"session_id": "session-below"}} + cache_key = router._get_session_affinity_cache_key("session-below", request_kwargs) + await mock_router_instance.cache.async_set_cache(key=cache_key, value="gpt-4o-mini") + + response = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=request_kwargs, + messages=[{"role": "user", "content": "LITELLM ESCALATE go higher"}], + ) + decision = response.routing_decision + assert decision["cause"] == "session_affinity_escalation" + assert decision["escalation_keyword"] == "LITELLM ESCALATE" + assert decision["escalated"] is True + + @pytest.mark.asyncio + async def test_signals_are_a_json_array_not_a_stringified_tuple(self, complexity_router): + """The dashboard maps over `signals`, so the persisted shape has to be an array + regardless of how any given serializer treats sequence types.""" + import json + + from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + + response = await complexity_router.async_pre_routing_hook( + model="test-router", + request_kwargs={}, + messages=[{"role": "user", "content": "Hello!"}], + ) + signals = response.routing_decision["signals"] + assert isinstance(signals, list) + assert isinstance(json.loads(safe_dumps({"d": response.routing_decision}))["d"]["signals"], list) + + +class TestRedactedLoggingDropsPromptText: + """An operator who turns message logging off has said prompt content must not reach + the logs. The routing decision quotes the prompt in its matched keywords and in the + signals that name them, so those are dropped while the derived values that make the + row explainable are kept.""" + + MODEL_LIST = [ + { + "model_name": "smart-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": ["gpt-4o-mini"], "REASONING": ["gpt-4o"]}, + "session_affinity": False, + "keyword_tier_rules": [{"keywords": ["deploy to k8s"], "tier": "REASONING"}], + }, + }, + }, + {"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini"}}, + {"model_name": "gpt-4o", "litellm_params": {"model": "openai/gpt-4o"}}, + ] + + MESSAGES = [{"role": "user", "content": "LITELLM ESCALATE please deploy to k8s now"}] + + async def _decision(self, request_kwargs: Dict) -> Dict: + router = Router(model_list=self.MODEL_LIST) + response = await router.async_pre_routing_hook( + model="smart-router", request_kwargs=request_kwargs, messages=self.MESSAGES + ) + assert response is not None + return request_kwargs["metadata"]["routing_decision"] + + @pytest.mark.asyncio + async def test_prompt_text_is_persisted_when_logging_is_not_redacted(self): + decision = await self._decision({}) + # Control: without redaction the terms are the point of the feature. + assert decision["matched_keyword"] == "deploy to k8s" + assert decision["escalation_keyword"] == "LITELLM ESCALATE" + + @pytest.mark.asyncio + async def test_redaction_drops_quoted_prompt_text_but_keeps_the_explanation(self, monkeypatch): + # The usual deployment shape: `litellm_settings: turn_off_message_logging: true` + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + decision = await self._decision({}) + + for field in ("signals", "matched_keyword", "escalation_keyword"): + assert field not in decision, f"{field} quotes the prompt and must be dropped" + # Nothing here reproduces the prompt, so the row stays explainable. + assert decision["cause"] == "literal_keyword_match" + assert decision["tier"] == "REASONING" + assert decision["routed_model"] == "gpt-4o" + assert decision["escalated"] is False + + def test_only_verbatim_prompt_fields_are_classified_as_prompt_text(self, monkeypatch): + """The field classification is the whole contract, so pin it directly: anything + that quotes the prompt goes, anything derived from it stays.""" + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + full = { + "router_model_name": "smart-router", + "router_type": "complexity", + "routed_model": "gpt-4o", + "cause": "literal_keyword_match", + "tier": "REASONING", + "score": 0.8, + "tier_boundaries": {"simple_medium": 0.15, "medium_complex": 0.35, "complex_reasoning": 0.6}, + "classifier_model": "claude-haiku", + "escalated": True, + "signals": ["code (python)"], + "matched_keyword": "deploy to k8s", + "escalation_keyword": "LITELLM ESCALATE", + } + kept = Router._redact_prompt_text_if_needed(request_kwargs={}, routing_decision=full) + assert set(full) - set(kept) == {"signals", "matched_keyword", "escalation_keyword"} + + @pytest.mark.asyncio + async def test_redaction_via_request_header_is_honored(self): + request_kwargs: Dict = {"metadata": {"headers": {"x-litellm-enable-message-redaction": True}}} + decision = await self._decision(request_kwargs) + assert "matched_keyword" not in decision + assert decision["cause"] == "literal_keyword_match" + + +def test_every_routing_decision_field_is_classified(): + """Redaction is derived from a declaration, not a list at the call site, so every + field has to be classified as quoting the prompt or aggregating it. A field added + without a decision fails here rather than silently shipping unredacted or, worse, + being over-redacted and taking a load-bearing fact with it.""" + from litellm.types.utils import ( + DERIVED_ROUTING_DECISION_FIELDS, + PROMPT_QUOTING_ROUTING_DECISION_FIELDS, + StandardLoggingRoutingDecision, + ) + + declared = set(StandardLoggingRoutingDecision.__annotations__) + classified = PROMPT_QUOTING_ROUTING_DECISION_FIELDS | DERIVED_ROUTING_DECISION_FIELDS + assert declared == classified, ( + "classify new routing-decision fields in litellm/types/utils.py: " + f"unclassified={declared - classified}, stale={classified - declared}" + ) + assert not (PROMPT_QUOTING_ROUTING_DECISION_FIELDS & DERIVED_ROUTING_DECISION_FIELDS) diff --git a/tests/test_litellm/router_strategy/test_lowest_latency.py b/tests/test_litellm/router_strategy/test_lowest_latency.py new file mode 100644 index 00000000000..4edc1e21d6b --- /dev/null +++ b/tests/test_litellm/router_strategy/test_lowest_latency.py @@ -0,0 +1,170 @@ +#### What this tests #### +# Latency values recorded by lowest-latency routing must be JSON +# serializable for non-chat responses too (embeddings/speech/image skip +# the ModelResponse branch, so the raw timedelta used to leak into the +# latency list and break the Redis cache sync). Issue #33169. + +import json +import os +import sys +from datetime import datetime, timedelta + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + +import litellm +from litellm.caching.caching import DualCache +from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler + +DEPLOYMENT_ID = "9876" +KWARGS = { + "litellm_params": { + "metadata": { + "model_group": "gemini-embedding-001", + "deployment": "vertex_ai/gemini-embedding-001", + }, + "model_info": {"id": DEPLOYMENT_ID}, + } +} + + +def _embedding_response(): + return litellm.EmbeddingResponse( + model="gemini-embedding-001", + data=[{"embedding": [0.1, 0.2], "index": 0, "object": "embedding"}], + object="list", + usage=litellm.Usage(prompt_tokens=5, completion_tokens=0, total_tokens=5), + ) + + +def _recorded_latencies(cache: DualCache): + cached = cache.get_cache(key="gemini-embedding-001_map") or {} + return cached.get(DEPLOYMENT_ID, {}).get("latency", []) + + +def test_sync_embedding_latency_is_json_serializable(): + """log_success_event with datetime start/end (as the proxy passes) must not + record a raw timedelta for non-ModelResponse results.""" + cache = DualCache() + handler = LowestLatencyLoggingHandler(router_cache=cache) + + start_time = datetime(2026, 1, 1, 12, 0, 0) + end_time = datetime(2026, 1, 1, 12, 0, 2) + + handler.log_success_event( + response_obj=_embedding_response(), + kwargs=KWARGS, + start_time=start_time, + end_time=end_time, + ) + + latencies = _recorded_latencies(cache) + assert latencies, "expected a latency entry to be recorded" + assert all( + not isinstance(value, timedelta) for value in latencies + ), f"raw timedelta leaked into latency list: {latencies}" + assert latencies[-1] == pytest.approx(2.0) + # the exact failure mode from production: redis cache sync json.dumps + json.dumps({"latency": latencies}) + + +@pytest.mark.asyncio +async def test_async_embedding_latency_is_json_serializable(): + """async_log_success_event is the path the proxy actually hits.""" + cache = DualCache() + handler = LowestLatencyLoggingHandler(router_cache=cache) + + start_time = datetime(2026, 1, 1, 12, 0, 0) + end_time = datetime(2026, 1, 1, 12, 0, 3) + + await handler.async_log_success_event( + response_obj=_embedding_response(), + kwargs=KWARGS, + start_time=start_time, + end_time=end_time, + ) + + latencies = _recorded_latencies(cache) + assert latencies, "expected a latency entry to be recorded" + assert all( + not isinstance(value, timedelta) for value in latencies + ), f"raw timedelta leaked into latency list: {latencies}" + assert latencies[-1] == pytest.approx(3.0) + json.dumps({"latency": latencies}) + + +def _chat_response(completion_tokens: int): + return litellm.ModelResponse( + model="gpt-4o-mini", + choices=[ + litellm.Choices( + finish_reason="stop", + index=0, + message=litellm.Message(content="hi", role="assistant"), + ) + ], + usage=litellm.Usage( + prompt_tokens=10, + completion_tokens=completion_tokens, + total_tokens=10 + completion_tokens, + ), + ) + + +@pytest.mark.asyncio +async def test_async_chat_latency_normalized_per_token(): + """Chat responses go through the per-token normalization branch — with the + up-front timedelta conversion the stored value must be seconds/token.""" + cache = DualCache() + handler = LowestLatencyLoggingHandler(router_cache=cache) + + await handler.async_log_success_event( + response_obj=_chat_response(completion_tokens=4), + kwargs=KWARGS, + start_time=datetime(2026, 1, 1, 12, 0, 0), + end_time=datetime(2026, 1, 1, 12, 0, 2), + ) + + latencies = _recorded_latencies(cache) + assert latencies and latencies[-1] == pytest.approx(0.5) # 2s / 4 tokens + json.dumps({"latency": latencies}) + + +@pytest.mark.asyncio +async def test_async_chat_zero_completion_tokens_falls_back_to_seconds(): + """safe_divide_seconds returns None for zero tokens — the fallback branch + must store plain float seconds, not a timedelta.""" + cache = DualCache() + handler = LowestLatencyLoggingHandler(router_cache=cache) + + await handler.async_log_success_event( + response_obj=_chat_response(completion_tokens=0), + kwargs=KWARGS, + start_time=datetime(2026, 1, 1, 12, 0, 0), + end_time=datetime(2026, 1, 1, 12, 0, 3), + ) + + latencies = _recorded_latencies(cache) + assert latencies and latencies[-1] == pytest.approx(3.0) + assert not isinstance(latencies[-1], timedelta) + json.dumps({"latency": latencies}) + + +def test_sync_chat_zero_completion_tokens_falls_back_to_seconds(): + cache = DualCache() + handler = LowestLatencyLoggingHandler(router_cache=cache) + + handler.log_success_event( + response_obj=_chat_response(completion_tokens=0), + kwargs=KWARGS, + start_time=datetime(2026, 1, 1, 12, 0, 0), + end_time=datetime(2026, 1, 1, 12, 0, 2), + ) + + latencies = _recorded_latencies(cache) + assert latencies and latencies[-1] == pytest.approx(2.0) + assert not isinstance(latencies[-1], timedelta) + json.dumps({"latency": latencies}) diff --git a/tests/test_litellm/router_strategy/test_quality_router.py b/tests/test_litellm/router_strategy/test_quality_router.py index 01574cb980d..b2e901739da 100644 --- a/tests/test_litellm/router_strategy/test_quality_router.py +++ b/tests/test_litellm/router_strategy/test_quality_router.py @@ -1031,3 +1031,53 @@ class TestRouterQualityDeploymentMethods: ) router.init_quality_router_deployment(deployment) assert "auto_router/quality_router/test-router" in router.quality_routers + + +class TestRoutingDecisionProvenance: + """Every quality-router path must attach a routing_decision to its hook response, + including the no-user-message default path that previously recorded nothing.""" + + @pytest.mark.asyncio + async def test_quality_tier_decision(self, quality_router): + resp = await quality_router.async_pre_routing_hook( + model="quality-router-test", + request_kwargs={}, + messages=[{"role": "user", "content": "hi"}], + ) + assert resp is not None + decision = resp.routing_decision + assert decision is not None + assert decision["router_model_name"] == "quality-router-test" + assert decision["router_type"] == "quality" + assert decision["cause"] == "quality_tier" + assert decision["routed_model"] == "haiku" + assert decision["tier"] == "1" + assert isinstance(decision["score"], float) + + @pytest.mark.asyncio + async def test_keyword_decision_carries_matched_keyword(self, keyword_router): + resp = await keyword_router.async_pre_routing_hook( + model="quality-router-test", + request_kwargs={}, + messages=[{"role": "user", "content": "please write python for me"}], + ) + assert resp is not None + decision = resp.routing_decision + assert decision is not None + assert decision["cause"] == "keyword" + assert decision["matched_keyword"] == "python" + assert decision["routed_model"] == resp.model + + @pytest.mark.asyncio + async def test_no_user_message_decision_is_default_fallback(self, quality_router): + resp = await quality_router.async_pre_routing_hook( + model="quality-router-test", + request_kwargs={}, + messages=[{"role": "system", "content": "You are a helpful assistant."}], + ) + assert resp is not None + decision = resp.routing_decision + assert decision is not None + assert decision["cause"] == "default_fallback" + assert decision["routed_model"] == "haiku" + assert "tier" not in decision diff --git a/tests/test_litellm/router_utils/test_auto_router_model_naming.py b/tests/test_litellm/router_utils/test_auto_router_model_naming.py new file mode 100644 index 00000000000..ca290caac0b --- /dev/null +++ b/tests/test_litellm/router_utils/test_auto_router_model_naming.py @@ -0,0 +1,77 @@ +import pytest + +from litellm.router_utils.auto_router_model_naming import ( + classify_strategy_router_model, + validate_strategy_router_model_write, +) + +COMPLEXITY_FIELDS = frozenset({"complexity_router_config"}) +SEMANTIC_FIELDS = frozenset( + {"auto_router_config", "auto_router_default_model", "auto_router_embedding_model"} +) + + +@pytest.mark.parametrize( + "model,expected", + [ + ("anthropic/claude-sonnet-5", None), + ("complexity_router", None), + ("autorouter/complexity_router", None), + ("auto_router/my-router", "semantic"), + ("auto_router/complexity_router", "complexity"), + ("auto_router/complexity_router-eu", "complexity"), + ("auto_router/adaptive_router", "adaptive"), + ("auto_router/quality_router", "quality"), + ("auto_router/auto_router/complexity_router", "semantic"), + ("auto_router/", "semantic"), + ], +) +def test_classify_strategy_router_model(model, expected): + assert classify_strategy_router_model(model) == expected + + +@pytest.mark.parametrize( + "model,present_fields,expected_fragment", + [ + ("auto_router/auto_router/complexity_router", COMPLEXITY_FIELDS, "repeats"), + ("complexity_router", COMPLEXITY_FIELDS, "does not start with"), + ("anthropic/claude-sonnet-5", COMPLEXITY_FIELDS, "does not start with"), + ("auto_router/", frozenset(), "missing the router name"), + ("auto_router/complexity_router", frozenset(), "requires"), + ("auto_router/my-router", frozenset({"auto_router_config"}), "requires"), + ("auto_router/adaptive_router", frozenset(), "requires"), + ("auto_router/quality_router", frozenset(), "requires"), + ], +) +def test_validate_rejects_incoherent_writes(model, present_fields, expected_fragment): + violation = validate_strategy_router_model_write(model=model, present_fields=present_fields) + assert violation is not None + assert expected_fragment in violation + + +@pytest.mark.parametrize( + "model,present_fields", + [ + ("anthropic/claude-sonnet-5", frozenset()), + ("openai/gpt-4o-mini", frozenset({"api_key"})), + ("auto_router/complexity_router", COMPLEXITY_FIELDS), + ("auto_router/complexity_router", frozenset({"complexity_router_default_model"})), + ("auto_router/complexity_router-eu", COMPLEXITY_FIELDS), + ("auto_router/my-router", SEMANTIC_FIELDS), + ( + "auto_router/my-router", + frozenset( + { + "auto_router_config_path", + "auto_router_default_model", + "auto_router_embedding_model", + } + ), + ), + ("auto_router/adaptive_router", frozenset({"adaptive_router_config"})), + ("auto_router/quality_router", frozenset({"quality_router_default_model"})), + ("auto_router/quality_router", frozenset({"quality_router_config"})), + ], +) +def test_validate_accepts_coherent_writes(model, present_fields): + assert validate_strategy_router_model_write(model=model, present_fields=present_fields) is None diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 276ee96ed65..e6ae1f85cfd 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -1014,9 +1014,9 @@ def test_tiered_pricing_only_deployment_selects_router_model_id(): router = Router( model_list=[ { - "model_name": "qwen-3.7-plus", + "model_name": "qwen-tier-only", "litellm_params": { - "model": "dashscope/qwen3.7-plus", + "model": "dashscope/qwen-tier-only-test", "api_key": "sk-fake", }, "model_info": { @@ -1037,10 +1037,12 @@ def test_tiered_pricing_only_deployment_selects_router_model_id(): assert entry.get("input_cost_per_token") is None assert entry.get("tiered_pricing") is not None # The stripped shared alias must not carry tiered pricing. - assert litellm.model_cost["dashscope/qwen3.7-plus"].get("tiered_pricing") is None + assert ( + litellm.model_cost["dashscope/qwen-tier-only-test"].get("tiered_pricing") is None + ) selected = _select_model_name_for_cost_calc( - model="dashscope/qwen3.7-plus", + model="dashscope/qwen-tier-only-test", completion_response=None, custom_pricing=True, custom_llm_provider="dashscope", diff --git a/tests/test_litellm/test_gpt_5_4_model_metadata.py b/tests/test_litellm/test_gpt_5_4_model_metadata.py new file mode 100644 index 00000000000..f93e6187dcb --- /dev/null +++ b/tests/test_litellm/test_gpt_5_4_model_metadata.py @@ -0,0 +1,81 @@ +import json +from functools import lru_cache +from pathlib import Path + +import pytest + +REPO_ROOT = Path(__file__).parents[2] +MAIN_PATH = REPO_ROOT / "model_prices_and_context_window.json" +BACKUP_PATH = REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json" + +DOCUMENTED_MAX_INPUT_TOKENS = 272000 +DOCUMENTED_MAX_OUTPUT_TOKENS = 128000 + +SMALL_MODEL_NAMES = ( + "gpt-5.4-mini", + "gpt-5.4-mini-2026-03-17", + "gpt-5.4-nano", + "gpt-5.4-nano-2026-03-17", +) +SMALL_MODELS = tuple(f"{prefix}{name}" for prefix in ("", "azure/", "azure_ai/") for name in SMALL_MODEL_NAMES) + +STANDARD_PRICING = { + "gpt-5.4-mini": (7.5e-07, 4.5e-06, 7.5e-08), + "gpt-5.4-nano": (2e-07, 1.25e-06, 2e-08), +} + +LONG_CONTEXT_MODELS = ("gpt-5.4", "gpt-5.4-pro") + + +@lru_cache(maxsize=2) +def _load(path: Path) -> dict[str, dict[str, object]]: + with open(path) as f: + return json.load(f) + + +def _pricing_key(model: str) -> str: + return "gpt-5.4-nano" if "nano" in model else "gpt-5.4-mini" + + +@pytest.mark.parametrize("model", SMALL_MODELS) +def test_gpt_5_4_small_models_use_documented_token_limits(model: str) -> None: + """gpt-5.4-mini/nano are 400K-window models: 272K in, 128K out, not gpt-5.4's 1.05M window.""" + info = _load(MAIN_PATH).get(model) + assert info is not None, f"{model} not found in model_prices_and_context_window.json" + + assert info["max_input_tokens"] == DOCUMENTED_MAX_INPUT_TOKENS + assert info["max_output_tokens"] == DOCUMENTED_MAX_OUTPUT_TOKENS + assert info["max_tokens"] == DOCUMENTED_MAX_OUTPUT_TOKENS + + +@pytest.mark.parametrize("model", SMALL_MODELS) +def test_gpt_5_4_small_models_have_no_long_context_surcharge(model: str) -> None: + """OpenAI prices prompts above 272K at 2x input / 1.5x output for the 1.05M-window models only.""" + info = _load(MAIN_PATH)[model] + assert [key for key in info if "above_272k" in key] == [] + + +@pytest.mark.parametrize("model", SMALL_MODELS) +def test_gpt_5_4_small_models_standard_pricing(model: str) -> None: + info = _load(MAIN_PATH)[model] + input_cost, output_cost, cache_read_cost = STANDARD_PRICING[_pricing_key(model)] + + assert info["input_cost_per_token"] == input_cost + assert info["output_cost_per_token"] == output_cost + assert info["cache_read_input_token_cost"] == cache_read_cost + + +@pytest.mark.parametrize("model", LONG_CONTEXT_MODELS) +def test_gpt_5_4_long_context_models_keep_surcharge(model: str) -> None: + """The mini/nano correction must leave gpt-5.4 and gpt-5.4-pro tiered pricing intact.""" + info = _load(MAIN_PATH)[model] + + assert info["input_cost_per_token_above_272k_tokens"] == pytest.approx(info["input_cost_per_token"] * 2) + assert info["output_cost_per_token_above_272k_tokens"] == pytest.approx(info["output_cost_per_token"] * 1.5) + + +@pytest.mark.parametrize("model", SMALL_MODELS) +def test_gpt_5_4_small_models_backup_matches_main(model: str) -> None: + assert _load(BACKUP_PATH).get(model) == _load(MAIN_PATH).get(model), ( + f"{model} differs between main and backup model cost maps" + ) diff --git a/tests/test_litellm/test_model_block_unblock.py b/tests/test_litellm/test_model_block_unblock.py index 318b2f519c4..dc0098e405e 100644 --- a/tests/test_litellm/test_model_block_unblock.py +++ b/tests/test_litellm/test_model_block_unblock.py @@ -34,9 +34,9 @@ def _setup_model_block_mocks(monkeypatch, *, updated_blocked: bool): mock_prisma_client.db.litellm_proxymodeltable = model_table mock_router = MagicMock() - mock_router.get_deployment.return_value = None + mock_router.get_model_ids.return_value = [model_id] - mock_clear_cache = AsyncMock(return_value=None) + mock_clear_cache = AsyncMock(return_value=True) mock_audit_log = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -197,3 +197,51 @@ async def test_route_request_returns_403_when_model_is_fully_blocked(monkeypatch assert exc_info.value.status_code == 403 assert "Model is blocked" in exc_info.value.message + + +@pytest.mark.asyncio +async def test_model_block_surfaces_wholesale_reload_failure(monkeypatch): + """The write endpoints owe the caller an error when the pod failed to reload at all; + the DB row is saved but this pod is not serving the change.""" + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import block_model + + model_id, model_table, updated_row, mock_clear_cache, mock_audit_log = _setup_model_block_mocks( + monkeypatch, updated_blocked=True + ) + wiped_router = MagicMock() + wiped_router.get_model_ids.side_effect = [[model_id], []] + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", wiped_router) + + with pytest.raises(ProxyException, match=model_id): + await block_model( + data=BlockModelRequest(model_id=model_id), + http_request=MagicMock(), + user_api_key_dict=_proxy_admin(), + litellm_changed_by="operator@example.com", + ) + + assert mock_audit_log.call_args.kwargs["object_id"] == model_id + + +@pytest.mark.asyncio +async def test_model_block_surfaces_model_dropped_by_reload(monkeypatch): + """A reload that completes but drops the written model (ignore_invalid_deployments + swallowed its re-add) must not produce an unqualified success.""" + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import block_model + + model_id, model_table, updated_row, mock_clear_cache, _ = _setup_model_block_mocks( + monkeypatch, updated_blocked=True + ) + dropped_router = MagicMock() + dropped_router.get_model_ids.return_value = [] + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", dropped_router) + + with pytest.raises(ProxyException, match=model_id): + await block_model( + data=BlockModelRequest(model_id=model_id), + http_request=MagicMock(), + user_api_key_dict=_proxy_admin(), + litellm_changed_by="operator@example.com", + ) diff --git a/tests/test_litellm/test_model_prices_schema.py b/tests/test_litellm/test_model_prices_schema.py new file mode 100644 index 00000000000..f6f1bf16742 --- /dev/null +++ b/tests/test_litellm/test_model_prices_schema.py @@ -0,0 +1,99 @@ +from __future__ import annotations + +import importlib.util +import json +from pathlib import Path + +import jsonschema +import pytest + +REPO_ROOT = Path(__file__).parents[2] +GENERATOR_PATH = REPO_ROOT / "ci_cd" / "generate_model_prices_schema.py" +PRICES_PATH = REPO_ROOT / "model_prices_and_context_window.json" +SCHEMA_PATH = REPO_ROOT / "model_prices_and_context_window.schema.json" + + +def build_validator(schema: dict) -> jsonschema.Draft202012Validator: + return jsonschema.Draft202012Validator(schema, format_checker=jsonschema.Draft202012Validator.FORMAT_CHECKER) + + +def load_generator(): + spec = importlib.util.spec_from_file_location("generate_model_prices_schema", GENERATOR_PATH) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +@pytest.fixture(scope="module") +def committed_schema() -> dict: + return json.loads(SCHEMA_PATH.read_text()) + + +@pytest.fixture(scope="module") +def prices() -> dict: + return json.loads(PRICES_PATH.read_text()) + + +def test_committed_schema_matches_generator_output(prices: dict, committed_schema: dict): + generator = load_generator() + regenerated = json.loads(generator.render(generator.build_schema(prices))) + assert regenerated == committed_schema, ( + "model_prices_and_context_window.schema.json is out of sync; " + "run `python ci_cd/generate_model_prices_schema.py` and commit the result" + ) + + +def test_prices_file_validates_against_committed_schema(prices: dict, committed_schema: dict): + validator = build_validator(committed_schema) + errors = [ + f"{'.'.join(str(part) for part in error.absolute_path)}: {error.message}" + for error in validator.iter_errors(prices) + ] + assert errors == [] + + +@pytest.mark.parametrize( + "entry", + [ + {"litellm_provider": "openai", "mode": "chat", "input_cost_per_token": "0.01"}, + {"litellm_provider": "openai", "mode": "chat", "input_cost_per_token": -1}, + {"litellm_provider": "openai", "mode": "not_a_real_mode"}, + {"mode": "chat"}, + {"litellm_provider": "openai", "deprecation_date": "June 2026"}, + {"litellm_provider": "openai", "deprecation_date": "2026-99-99"}, + {"litellm_provider": "openai", "deprecation_date": "2026-13-01"}, + {"litellm_provider": "openai", "deprecation_date": "2026-01-32"}, + {"litellm_provider": "openai", "deprecation_date": "2026-01-00"}, + {"litellm_provider": "openai", "deprecation_date": "2026-02-31"}, + {"litellm_provider": "openai", "supported_modalities": ["smell"]}, + {"litellm_provider": "openai", "supports_vision": "yes"}, + {"litellm_provider": "openai", "max_tokens": 8191.5}, + {"litellm_provider": "openai", "tiered_pricing": [{"unknown_tier_field": 1}]}, + ], + ids=[ + "cost_as_string", + "negative_cost", + "unknown_mode", + "missing_provider", + "non_iso_deprecation_date", + "impossible_month_and_day", + "month_out_of_range", + "day_out_of_range", + "day_zero", + "calendar_impossible_day", + "unknown_modality", + "boolean_flag_as_string", + "fractional_max_tokens", + "unknown_tiered_pricing_field", + ], +) +def test_schema_rejects_malformed_entries(committed_schema: dict, entry: dict): + validator = build_validator(committed_schema) + assert not validator.is_valid({"some-model": entry}) + + +def test_schema_accepts_minimal_and_unknown_optional_fields(committed_schema: dict): + validator = build_validator(committed_schema) + assert validator.is_valid({"some-model": {"litellm_provider": "openai"}}) + assert validator.is_valid({"some-model": {"litellm_provider": "openai", "brand_new_field": {"nested": True}}}) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 0818237655d..e0fa800723d 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -737,3 +737,111 @@ def test_connection_pool_env_redis_ssl_false_uses_plain_connection(monkeypatch): assert pool is not None assert pool.connection_class is async_redis.Connection assert "ssl" not in pool.connection_kwargs + + +@pytest.mark.parametrize( + "redis_config", + [ + pytest.param({"host": "redis-host", "port": 6379}, id="host_port"), + pytest.param({"url": "redis://redis-host:6379"}, id="url"), + ], +) +def test_connection_pool_keeps_socket_timeout(redis_config, monkeypatch): + """The async pool must carry socket_timeout however Redis was configured. + + The url branch used to rebuild pool kwargs from scratch as {timeout, url, + max_connections}, dropping socket_timeout. redis-py then leaves both + socket_timeout and socket_connect_timeout (which falls back to it) unset, so a + Redis host that drops packets rather than refusing them blocks every caller + indefinitely instead of failing fast. + """ + monkeypatch.delenv("REDIS_URL", raising=False) + monkeypatch.delenv("REDIS_HOST", raising=False) + monkeypatch.delenv("REDIS_CLUSTER_NODES", raising=False) + + pool = get_redis_connection_pool(socket_timeout=5.0, **redis_config) + + assert pool is not None + assert pool.connection_kwargs.get("socket_timeout") == 5.0 + + +@pytest.mark.parametrize( + "redis_config", + [ + pytest.param({"host": "redis-host", "port": 6379}, id="host_port"), + pytest.param({"url": "redis://redis-host:6379"}, id="url"), + ], +) +def test_sync_client_keeps_socket_timeout(redis_config, monkeypatch): + """The sync client is built during RedisCache.__init__ and blocks the caller. + + Without socket_timeout it stalls for the OS TCP timeout against an unreachable + host, so merely constructing the cache stops the process. + """ + monkeypatch.delenv("REDIS_URL", raising=False) + monkeypatch.delenv("REDIS_HOST", raising=False) + monkeypatch.delenv("REDIS_CLUSTER_NODES", raising=False) + + client = get_redis_client(socket_timeout=5.0, **redis_config) + + assert client.connection_pool.connection_kwargs.get("socket_timeout") == 5.0 + + +@pytest.mark.parametrize( + "redis_config", + [ + pytest.param({"host": "redis-host", "port": 6379}, id="host_port"), + pytest.param({"url": "redis://redis-host:6379"}, id="url"), + ], +) +def test_async_client_keeps_socket_timeout(redis_config, monkeypatch): + """Same invariant for the async client built without an injected pool.""" + monkeypatch.delenv("REDIS_URL", raising=False) + monkeypatch.delenv("REDIS_HOST", raising=False) + monkeypatch.delenv("REDIS_CLUSTER_NODES", raising=False) + + client = get_redis_async_client(socket_timeout=5.0, **redis_config) + + assert client.connection_pool.connection_kwargs.get("socket_timeout") == 5.0 + + +def test_url_config_does_not_forward_ssl_kwarg(monkeypatch): + """ssl stays consumed rather than forwarded on the url path. + + TLS is selected by the rediss:// scheme there; handing ssl=True to a redis:// + url yields a plain Connection that rejects the kwarg when it first connects. + """ + monkeypatch.delenv("REDIS_URL", raising=False) + monkeypatch.delenv("REDIS_HOST", raising=False) + monkeypatch.delenv("REDIS_CLUSTER_NODES", raising=False) + + client = get_redis_client(url="redis://redis-host:6379", ssl=True) + + assert "ssl" not in client.connection_pool.connection_kwargs + + +@pytest.mark.parametrize( + "client_only_kwarg", + [ + pytest.param({"single_connection_client": True}, id="single_connection_client"), + pytest.param({"auto_close_connection_pool": True}, id="auto_close_connection_pool"), + pytest.param({"ssl_ca_certs": "/tmp/ca.pem"}, id="ssl_ca_certs"), + pytest.param({"ssl": True}, id="ssl"), + ], +) +def test_url_config_drops_kwargs_the_connection_cannot_accept(client_only_kwarg, monkeypatch): + """Only kwargs the connection accepts may be forwarded on the url path. + + from_url hands its kwargs down to the connection class, so client-level settings and + the SSLConnection-only ssl_* family raise TypeError the first time a connection is + created. TLS on a url config comes from the rediss:// scheme instead. + """ + monkeypatch.delenv("REDIS_URL", raising=False) + monkeypatch.delenv("REDIS_HOST", raising=False) + monkeypatch.delenv("REDIS_CLUSTER_NODES", raising=False) + + pool = get_redis_connection_pool(url="redis://redis-host:6379", socket_timeout=5.0, **client_only_kwarg) + + assert pool is not None + pool.make_connection() + assert pool.connection_kwargs.get("socket_timeout") == 5.0 diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index ad4e430c603..46b5ce65c3f 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -3755,6 +3755,182 @@ def test_get_deployment_credentials_with_provider_team_wildcard_priority(): assert global_credentials["api_key"] == "global-key" +def test_get_deployment_credentials_with_provider_skips_other_team_deployment(): + """ + Regression: a team-scoped deployment sharing a model_name with a global + deployment must never resolve for another team's (or an unscoped) caller, + even when it is indexed first; the shared global deployment wins instead. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "team-b-project", + }, + "model_info": { + "id": "team-b-vertex", + "team_id": "team-b", + "team_public_model_name": "gemini-2.5-pro", + }, + }, + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "shared-project", + }, + }, + ], + ) + + other_team_credentials = router.get_deployment_credentials_with_provider( + model_id="gemini-2.5-pro", team_id="team-a" + ) + assert other_team_credentials is not None + assert other_team_credentials["vertex_project"] == "shared-project" + + unscoped_credentials = router.get_deployment_credentials_with_provider( + model_id="gemini-2.5-pro" + ) + assert unscoped_credentials is not None + assert unscoped_credentials["vertex_project"] == "shared-project" + + owner_credentials = router.get_deployment_credentials_with_provider( + model_id="gemini-2.5-pro", team_id="team-b" + ) + assert owner_credentials is not None + assert owner_credentials["vertex_project"] == "team-b-project" + + +def test_get_deployment_credentials_with_provider_no_fallback_to_other_team_only_name(): + """ + When the only deployments under a model name belong to another team, other + callers must get None (env fallback) instead of that team's credentials. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "team-b-project", + }, + "model_info": { + "id": "team-b-vertex", + "team_id": "team-b", + "team_public_model_name": "gemini-2.5-pro", + }, + }, + ], + ) + + assert ( + router.get_deployment_credentials_with_provider( + model_id="gemini-2.5-pro", team_id="team-a" + ) + is None + ) + assert ( + router.get_deployment_credentials_with_provider(model_id="gemini-2.5-pro") + is None + ) + + +def test_deployment_usable_by_team_helpers(): + """ + Direct coverage of the team-ownership filter: a team-scoped deployment is + usable only by its owning team, shared deployments by anyone, and the + model-group picker returns the first usable deployment or None. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "team-b-project", + }, + "model_info": { + "id": "team-b-vertex", + "team_id": "team-b", + "team_public_model_name": "gemini-2.5-pro", + }, + }, + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "shared-project", + }, + }, + ], + ) + + team_owned, shared = router.model_list + assert router._deployment_usable_by_team(team_owned, "team-b") is True + assert router._deployment_usable_by_team(team_owned, "team-a") is False + assert router._deployment_usable_by_team(team_owned, None) is False + assert router._deployment_usable_by_team(shared, "team-a") is True + assert router._deployment_usable_by_team(shared, None) is True + + picked = router._get_model_group_deployment_usable_by_team( + model_group_name="gemini-2.5-pro", team_id="team-a" + ) + assert picked is not None + assert picked.litellm_params.vertex_project == "shared-project" + + owner_picked = router._get_model_group_deployment_usable_by_team( + model_group_name="gemini-2.5-pro", team_id="team-b" + ) + assert owner_picked is not None + assert owner_picked.litellm_params.vertex_project == "team-b-project" + + assert ( + router._get_model_group_deployment_usable_by_team( + model_group_name="unknown-model", team_id="team-a" + ) + is None + ) + + +def test_get_deployment_credentials_with_provider_skips_other_team_wildcard(): + """ + Global wildcard resolution must skip a team-scoped wildcard deployment for + callers outside that team, falling through to the shared wildcard entry. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "openai/*", + "litellm_params": {"model": "openai/*", "api_key": "team-b-key"}, + "model_info": { + "id": "team-b-wildcard", + "team_id": "team-b", + "team_public_model_name": "openai/*", + }, + }, + { + "model_name": "openai/*", + "litellm_params": {"model": "openai/*", "api_key": "global-key"}, + }, + ], + ) + + other_team_credentials = router.get_deployment_credentials_with_provider( + model_id="openai/gpt-5.2", team_id="team-a" + ) + assert other_team_credentials is not None + assert other_team_credentials["api_key"] == "global-key" + + owner_credentials = router.get_deployment_credentials_with_provider( + model_id="openai/gpt-5.2", team_id="team-b" + ) + assert owner_credentials is not None + assert owner_credentials["api_key"] == "team-b-key" + + def test_team_wildcard_credentials_not_usable_after_delete_deployment(): """ Regression: team_pattern_routers retained deleted deployments, so a team @@ -6379,3 +6555,22 @@ class TestPreRoutingStrategyRegistryLifecycle: litellm_params=LiteLLM_Params(**params) ) assert actual is expected, params["model"] + + +def test_model_info_is_active_for_environment_matrix(monkeypatch): + """The model-write endpoints consult this predicate to tell a deliberately + environment-inactive model from one dropped by a failed reload; the Router's own + deployment gate delegates to it, so the two can never diverge.""" + from litellm.router import model_info_is_active_for_environment + + assert model_info_is_active_for_environment(model_info=None) is True + assert model_info_is_active_for_environment(model_info={"id": "m1"}) is True + assert model_info_is_active_for_environment(model_info={"supported_environments": None}) is True + + monkeypatch.setenv("LITELLM_ENVIRONMENT", "development") + assert model_info_is_active_for_environment(model_info={"supported_environments": ["development"]}) is True + assert model_info_is_active_for_environment(model_info={"supported_environments": ["production"]}) is False + + monkeypatch.delenv("LITELLM_ENVIRONMENT") + with pytest.raises(ValueError, match="LITELLM_ENVIRONMENT"): + model_info_is_active_for_environment(model_info={"supported_environments": ["production"]}) diff --git a/tests/test_litellm/test_ssl_verify_unit.py b/tests/test_litellm/test_ssl_verify_unit.py index 7cc15703a3b..c39362c01a2 100644 --- a/tests/test_litellm/test_ssl_verify_unit.py +++ b/tests/test_litellm/test_ssl_verify_unit.py @@ -17,7 +17,6 @@ sys.path.insert(0, str(Path(__file__).parent)) import litellm.proxy.guardrails.guardrail_hooks.aim.aim as _aim_module import litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks as _cato_networks_module from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM -from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail from litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks import CatoNetworksGuardrail @@ -87,23 +86,6 @@ class TestBaseAWSLLMSSLVerify: assert True # If we got here without error, parameter was accepted -class TestBedrockLLMSSLVerify: - """Test SSL verification parameter handling in BedrockLLM.""" - - def test_bedrock_llm_accepts_ssl_verify_in_optional_params(self): - """Test that BedrockLLM can receive ssl_verify in optional_params.""" - # This is a simple test to verify the parameter is accepted - # The actual propagation is tested in integration tests - bedrock_llm = BedrockLLM() - - # Verify the class exists and can be instantiated - assert bedrock_llm is not None - - # Verify _get_ssl_verify method exists and works - result = bedrock_llm._get_ssl_verify(ssl_verify="/path/to/cert.pem") - assert result == "/path/to/cert.pem" - - class TestAimGuardrailSSLVerify: """Test SSL verification parameter handling in AimGuardrail.""" diff --git a/type-discipline-budget.json b/type-discipline-budget.json index d56d5a6e305..bef3a4c98aa 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 23287 + "limit": 23253 }, "LIT002": { - "limit": 27473 + "limit": 27433 }, "LIT003": { "limit": 292 @@ -15,7 +15,7 @@ "limit": 0 }, "LIT006": { - "limit": 1109 + "limit": 1108 }, "LIT007": { "limit": 0 @@ -24,6 +24,6 @@ "limit": 1004 }, "LIT009": { - "limit": 2495 + "limit": 2474 } } diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 1c07f2d6247..6819b2851f5 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -151,14 +151,11 @@ "no-restricted-imports": { "count": 1 }, - "prefer-const": { - "count": 3 - }, "react-hooks/purity": { "count": 1 }, "react-hooks/set-state-in-effect": { - "count": 2 + "count": 1 } }, "src/app/(dashboard)/caching/_components/cache_health.tsx": { @@ -210,11 +207,6 @@ "count": 1 } }, - "src/app/(dashboard)/cost-optimization/_components/AutorouterTab.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx": { "no-restricted-imports": { "count": 1 @@ -1696,7 +1688,7 @@ "count": 1 }, "no-restricted-imports": { - "count": 3 + "count": 2 }, "prefer-const": { "count": 2 @@ -2550,17 +2542,12 @@ "count": 1 } }, - "src/components/add_model/add_auto_router_tab.test.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/add_model/add_auto_router_tab.tsx": { "local/filename-pascal-case": { "count": 1 }, "no-restricted-imports": { - "count": 3 + "count": 2 } }, "src/components/add_model/add_model_modes.tsx": { @@ -2568,19 +2555,6 @@ "count": 1 } }, - "src/components/add_model/add_model_tab.test.tsx": { - "no-restricted-imports": { - "count": 2 - } - }, - "src/components/add_model/add_model_tab.tsx": { - "local/filename-pascal-case": { - "count": 1 - }, - "no-restricted-imports": { - "count": 4 - } - }, "src/components/add_model/advanced_settings.tsx": { "local/filename-pascal-case": { "count": 1 @@ -3343,7 +3317,7 @@ "count": 5 }, "no-restricted-syntax": { - "count": 153 + "count": 152 }, "prefer-const": { "count": 32 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTable.test.tsx index 9a7a7bd2eb9..2b97bcbc072 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTable.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTable.test.tsx @@ -35,6 +35,15 @@ describe("BudgetTable", () => { expect(screen.getByText("10")).toBeInTheDocument(); }); + it("should render the budget id without a fixed character-count clamp", () => { + const budgetId = "ecc1869c-6231-4380-a56d-1a0be457477d"; + renderWithProviders(); + const idCell = screen.getByText(budgetId); + expect(idCell.className).not.toMatch(/max-w-\[\d+(ch|rem|px)\]/); + expect(idCell.className).toContain("max-w-full"); + expect(idCell.className).toContain("truncate"); + }); + it("should show n/a for missing rate limits and Unlimited for a missing max budget", () => { renderWithProviders( , diff --git a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTableColumns.tsx index 456ab9d6b68..e3fbc9dba08 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTableColumns.tsx @@ -75,7 +75,7 @@ export const getBudgetTableColumns = ({ header: "Budget ID", size: 220, enableSorting: false, - cell: ({ row }) => , + cell: ({ row }) => , }, { id: "max_budget", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx index e1d02e9352d..fd8dd011b05 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx @@ -4,36 +4,50 @@ import { screen, waitFor, within } from "@testing-library/react"; import { renderWithProviders } from "../../../../../tests/test-utils"; import CacheDashboard from "./cache_dashboard"; -const { adminGlobalCacheActivity, cachingHealthCheckCall } = vi.hoisted(() => ({ - adminGlobalCacheActivity: vi.fn(), +const { useCacheActivity, cachingHealthCheckCall } = vi.hoisted(() => ({ + useCacheActivity: vi.fn(), cachingHealthCheckCall: vi.fn(), })); vi.mock("@/components/networking", () => ({ - adminGlobalCacheActivity, cachingHealthCheckCall, })); -const cacheActivity = [ - { - api_key: "sk-1", - model: "gpt-5.1", - call_type: "acompletion", - total_rows: 1500, - cache_hit_true_rows: 300, - cached_completion_tokens: 12000, - generated_completion_tokens: 48000, +vi.mock("@/app/(dashboard)/hooks/caching/useCacheActivity", () => ({ + useCacheActivity, +})); + +const cacheActivity = { + groups: [ + { + call_type: "acompletion", + api_requests: 1000, + cache_hits: 300, + failed_requests: 200, + cached_completion_tokens: 12000, + generated_completion_tokens: 48000, + }, + { + call_type: "aembedding", + api_requests: 550, + cache_hits: 100, + failed_requests: 50, + cached_completion_tokens: 2000, + generated_completion_tokens: 9000, + }, + ], + totals: { + api_requests: 1550, + cache_hits: 400, + failed_requests: 250, + cached_completion_tokens: 14000, + cache_hit_ratio: (400 / 2200) * 100, }, - { - api_key: "sk-2", - model: "text-embedding-3-large", - call_type: "aembedding", - total_rows: 700, - cache_hit_true_rows: 100, - cached_completion_tokens: 2000, - generated_completion_tokens: 9000, + filter_options: { + key_aliases: ["my-key", "Unnamed Key"], + models: ["gpt-5.1", "text-embedding-3-large"], }, -]; +}; const renderDashboard = () => renderWithProviders( @@ -75,7 +89,7 @@ const legendFillByCategory = (card: HTMLElement) => describe("CacheDashboard cache analytics charts", () => { beforeEach(() => { vi.clearAllMocks(); - adminGlobalCacheActivity.mockResolvedValue(cacheActivity); + useCacheActivity.mockReturnValue({ data: cacheActivity, refetch: vi.fn() }); }); it("renders both chart card titles", async () => { @@ -108,8 +122,13 @@ describe("CacheDashboard cache analytics charts", () => { expect(legendFillByCategory(requestsCard)).toEqual({ "LLM API requests": "var(--color-sky-500, #0ea5e9)", "Cache hit": "var(--color-teal-500, #14b8a6)", + "Failed requests": "var(--color-red-500, #ef4444)", }); - expect(barFills(requestsCard)).toEqual(["var(--color-sky-500, #0ea5e9)", "var(--color-teal-500, #14b8a6)"]); + expect(barFills(requestsCard)).toEqual([ + "var(--color-sky-500, #0ea5e9)", + "var(--color-teal-500, #14b8a6)", + "var(--color-red-500, #ef4444)", + ]); }); it("renders the tokens chart with each category legend-bound to its fill and stacked in order", async () => { @@ -133,18 +152,39 @@ describe("CacheDashboard cache analytics charts", () => { } }); - it("stacks the two categories into one column per call_type", async () => { + it("stacks all categories into one column per call_type", async () => { renderDashboard(); const { requestsCard, tokensCard } = await findChartCards(); - for (const card of [requestsCard, tokensCard]) { + const expectedRects = { requests: 6, tokens: 4 }; + for (const [card, rectCount] of [ + [requestsCard, expectedRects.requests], + [tokensCard, expectedRects.tokens], + ] as const) { const rects = Array.from(card.querySelectorAll("path.recharts-rectangle")); - expect(rects).toHaveLength(4); + expect(rects).toHaveLength(rectCount); const xPositions = rects.map((rect) => rect.getAttribute("d")?.split(",")[0]); expect(new Set(xPositions).size).toBe(2); } }); + it("renders the server-computed cache hit ratio", async () => { + renderDashboard(); + + expect(await screen.findByText("18.18%")).toBeInTheDocument(); + }); + + it("passes the date range and selected filters to the activity query", () => { + renderDashboard(); + + expect(useCacheActivity).toHaveBeenCalledWith({ + startDate: expect.stringMatching(/^\d{4}-\d{2}-\d{2}$/), + endDate: expect.stringMatching(/^\d{4}-\d{2}-\d{2}$/), + keyAliases: [], + models: [], + }); + }); + it("formats y-axis ticks with compact notation", async () => { renderDashboard(); const { requestsCard, tokensCard } = await findChartCards(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx index 95b73d1aacb..47c266ceac0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx @@ -19,13 +19,29 @@ import { import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { RefreshCw } from "lucide-react"; -import { adminGlobalCacheActivity, cachingHealthCheckCall } from "@/components/networking"; +import { cachingHealthCheckCall } from "@/components/networking"; +import { useCacheActivity, type CacheActivityGroup } from "@/app/(dashboard)/hooks/caching/useCacheActivity"; // Import the new component import { CacheHealthTab } from "./cache_health"; import CacheSettings from "./cache_settings"; import CoordinationRedisSettings from "./coordination_redis_settings"; +const REQUEST_SERIES = { + apiRequests: "LLM API requests", + cacheHits: "Cache hit", + failed: "Failed requests", +} as const; + +const toChartDatum = (group: CacheActivityGroup) => ({ + name: group.call_type, + [REQUEST_SERIES.apiRequests]: group.api_requests, + [REQUEST_SERIES.cacheHits]: group.cache_hits, + [REQUEST_SERIES.failed]: group.failed_requests, + "Cached Completion Tokens": group.cached_completion_tokens, + "Generated Completion Tokens": group.generated_completion_tokens, +}); + const formatDateWithoutTZ = (date: Date | undefined) => { if (!date) return undefined; return date.toISOString().split("T")[0]; @@ -49,26 +65,6 @@ interface CachePageProps { premiumUser: boolean; } -interface cacheDataItem { - api_key: string; - model: string; - cache_hit_true_rows: number; - cached_completion_tokens: number; - total_rows: number; - generated_completion_tokens: number; - call_type: string; - - // Add other properties as needed -} - -type uiData = { - name: string; - "LLM API requests": number; - "Cache hit": number; - "Cached Completion Tokens": number; - "Generated Completion Tokens": number; -}; - interface CacheHealthResponse { status?: string; cache_type?: string; @@ -97,13 +93,8 @@ const deepParse = (input: any) => { }; const CacheDashboard: React.FC = ({ accessToken, token, userRole, userID, premiumUser }) => { - const [filteredData, setFilteredData] = useState([]); const [selectedApiKeys, setSelectedApiKeys] = useState([]); const [selectedModels, setSelectedModels] = useState([]); - const [data, setData] = useState([]); - const [cachedResponses, setCachedResponses] = useState("0"); - const [cachedTokens, setCachedTokens] = useState("0"); - const [cacheHitRatio, setCacheHitRatio] = useState("0"); const [dateValue, setDateValue] = useState({ from: new Date(Date.now() - 7 * 24 * 60 * 60 * 1000), @@ -113,120 +104,24 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole const [lastRefreshed, setLastRefreshed] = useState(""); const [healthCheckResponse, setHealthCheckResponse] = useState(""); - useEffect(() => { - if (!accessToken || !dateValue) { - return; - } - const fetchData = async () => { - const response = await adminGlobalCacheActivity( - accessToken, - formatDateWithoutTZ(dateValue.from), - formatDateWithoutTZ(dateValue.to), - ); - setData(response); - }; - fetchData(); - - const currentDate = new Date(); - setLastRefreshed(currentDate.toLocaleString()); - }, [accessToken]); - - const uniqueApiKeys = Array.from(new Set(data.map((item) => item?.api_key ?? ""))); - const uniqueModels = Array.from(new Set(data.map((item) => item?.model ?? ""))); - const uniqueCallTypes = Array.from(new Set(data.map((item) => item?.call_type ?? ""))); - - const updateCachingData = async (startTime: Date | undefined, endTime: Date | undefined) => { - if (!startTime || !endTime || !accessToken) { - return; - } - - let new_cache_data = await adminGlobalCacheActivity( - accessToken, - formatDateWithoutTZ(startTime), - formatDateWithoutTZ(endTime), - ); - - setData(new_cache_data); - }; + const { data: activity, refetch } = useCacheActivity({ + startDate: formatDateWithoutTZ(dateValue.from), + endDate: formatDateWithoutTZ(dateValue.to), + keyAliases: selectedApiKeys, + models: selectedModels, + }); useEffect(() => { - let newData: cacheDataItem[] = data; - if (selectedApiKeys.length > 0) { - newData = newData.filter((item) => selectedApiKeys.includes(item.api_key)); - } + setLastRefreshed(new Date().toLocaleString()); + }, []); - if (selectedModels.length > 0) { - newData = newData.filter((item) => selectedModels.includes(item.model)); - } - - /* - Data looks like this - [{"api_key":"sk-test-mock-key-001","call_type":"acompletion","model":"llama3-8b-8192","total_rows":13,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-002","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-123","call_type":"acompletion","model":"gpt-3.5-turbo","total_rows":19,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-123","call_type":"aimage_generation","model":"","total_rows":3,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-003","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-004","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-005","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - */ - - // What data we need for bar chat - // ui_data = [ - // { - // name: "Call Type", - // Cache hit: 20, - // LLM API requests: 10, - // } - // ] - - let llm_api_requests = 0; - let cache_hits = 0; - let cached_tokens = 0; - const processedData = newData.reduce((acc: uiData[], item) => { - if (!item.call_type) { - item.call_type = "Unknown"; - } - - llm_api_requests += (item.total_rows || 0) - (item.cache_hit_true_rows || 0); - cache_hits += item.cache_hit_true_rows || 0; - cached_tokens += item.cached_completion_tokens || 0; - - const existingItem = acc.find((i) => i.name === item.call_type); - if (existingItem) { - existingItem["LLM API requests"] += (item.total_rows || 0) - (item.cache_hit_true_rows || 0); - existingItem["Cache hit"] += item.cache_hit_true_rows || 0; - existingItem["Cached Completion Tokens"] += item.cached_completion_tokens || 0; - existingItem["Generated Completion Tokens"] += item.generated_completion_tokens || 0; - } else { - acc.push({ - name: item.call_type, - "LLM API requests": (item.total_rows || 0) - (item.cache_hit_true_rows || 0), - "Cache hit": item.cache_hit_true_rows || 0, - "Cached Completion Tokens": item.cached_completion_tokens || 0, - "Generated Completion Tokens": item.generated_completion_tokens || 0, - }); - } - return acc; - }, []); - - // set header cache statistics - setCachedResponses(valueFormatterNumbers(cache_hits)); - setCachedTokens(valueFormatterNumbers(cached_tokens)); - let allRequests = cache_hits + llm_api_requests; - if (allRequests > 0) { - let cache_hit_ratio = ((cache_hits / allRequests) * 100).toFixed(2); - setCacheHitRatio(cache_hit_ratio); - } else { - setCacheHitRatio("0"); - } - - setFilteredData(processedData); - }, [selectedApiKeys, selectedModels, dateValue, data]); + const uniqueApiKeys = activity?.filter_options.key_aliases ?? []; + const uniqueModels = activity?.filter_options.models ?? []; + const chartData = (activity?.groups ?? []).map(toChartDatum); const handleRefreshClick = () => { - // Update the 'lastRefreshed' state to the current date and time - const currentDate = new Date(); - setLastRefreshed(currentDate.toLocaleString()); + refetch(); + setLastRefreshed(new Date().toLocaleString()); }; const runCachingHealthCheck = async () => { @@ -257,10 +152,12 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole } }; + const totals = activity?.totals; + const hasRequests = totals != null && totals.api_requests + totals.cache_hits + totals.failed_requests > 0; const statCards = [ - { label: "Cache Hit Ratio", value: `${cacheHitRatio}%` }, - { label: "Cache Hits", value: cachedResponses }, - { label: "Cached Completion Tokens", value: cachedTokens }, + { label: "Cache Hit Ratio", value: `${hasRequests ? totals.cache_hit_ratio.toFixed(2) : "0"}%` }, + { label: "Cache Hits", value: valueFormatterNumbers(totals?.cache_hits ?? 0) }, + { label: "Cached Completion Tokens", value: valueFormatterNumbers(totals?.cached_completion_tokens ?? 0) }, ]; return ( @@ -380,7 +277,6 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole value={dateValue} onValueChange={(value) => { setDateValue(value); - updateCachingData(value.from, value.to); }} /> @@ -404,12 +300,12 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole @@ -423,7 +319,7 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole = ({ accessToken, userRole }) => { - const [form] = Form.useForm(); - - if (!accessToken) { - return null; - } - - return ( -
- form.resetFields()} accessToken={accessToken} userRole={userRole} /> -
- ); -}; - -export default AutorouterTab; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx index 363525c48af..ca7adf07941 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx @@ -27,7 +27,6 @@ vi.mock("@/app/(dashboard)/router-settings/_components/general_settings", () => })); vi.mock("./PromptCompressionTab", () => ({ __esModule: true, default: () =>
})); -vi.mock("./AutorouterTab", () => ({ __esModule: true, default: () =>
})); import CostOptimizationView from "./CostOptimizationView"; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx index 46aa23fcfc0..42dc7719144 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx @@ -3,7 +3,6 @@ import { describe, expect, it, vi } from "vitest"; vi.mock("./UsageTab", () => ({ __esModule: true, default: () =>
})); vi.mock("./PromptCompressionTab", () => ({ __esModule: true, default: () =>
})); -vi.mock("./AutorouterTab", () => ({ __esModule: true, default: () =>
})); vi.mock("./PromptCachingTab", () => ({ __esModule: true, default: () =>
})); import CostOptimizationView from "./CostOptimizationView"; @@ -11,13 +10,13 @@ import CostOptimizationView from "./CostOptimizationView"; const renderView = () => render(); describe("CostOptimizationView", () => { - it("renders all four cost-optimization tabs", () => { - const { getByText } = renderView(); + it("renders the three cost-optimization tabs and no autorouter tab", () => { + const { getByText, queryByText } = renderView(); expect(getByText("Usage")).toBeInTheDocument(); expect(getByText("Prompt Compression")).toBeInTheDocument(); - expect(getByText("Autorouter")).toBeInTheDocument(); expect(getByText("Prompt Caching")).toBeInTheDocument(); + expect(queryByText("Autorouter")).not.toBeInTheDocument(); }); it("defaults to the Usage tab and switches the active tab on click", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx index 3bab6afee57..f6593e80999 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx @@ -6,7 +6,6 @@ import { Alert, Tabs } from "antd"; import UsageTab from "./UsageTab"; import PromptCompressionTab from "./PromptCompressionTab"; -import AutorouterTab from "./AutorouterTab"; import PromptCachingTab from "./PromptCachingTab"; import { useDailyActivityRange } from "./useDailyActivityRange"; @@ -30,11 +29,6 @@ const CostOptimizationView: React.FC = ({ accessToken label: "Prompt Compression", children: , }, - { - key: "autorouter", - label: "Autorouter", - children: , - }, { key: "caching", label: "Prompt Caching", @@ -50,7 +44,8 @@ const CostOptimizationView: React.FC = ({ accessToken

Cost Optimization

- Track and configure the mechanisms that save you money: prompt compression, prompt caching, and auto routing + Track and configure the mechanisms that save you money: prompt compression and prompt caching. Auto routers + live under Models + Endpoints, on the Auto-Routers tab

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx index 17331014c57..251ce631beb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx @@ -41,6 +41,7 @@ const modeDescriptions = { logging_only: "Logging Only - Only runs on logging callbacks without affecting the LLM call", pre_mcp_call: "Before MCP Tool Call - Runs before MCP tool execution and validates tool calls", during_mcp_call: "During MCP Tool Call - Runs in parallel with MCP tool execution for monitoring", + post_mcp_call: "After MCP Tool Call - Runs after MCP tool execution and checks the tool result", }; interface GuardrailPreset { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.test.ts new file mode 100644 index 00000000000..1a3baee1ce3 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.test.ts @@ -0,0 +1,72 @@ +import { renderHook } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { useCacheActivity, type CacheActivityParams } from "./useCacheActivity"; + +const useQueryMock = vi.fn(); +vi.mock("@/lib/http/api", () => ({ + $api: { useQuery: (...args: unknown[]) => useQueryMock(...args) }, +})); + +const mockUseAuthorized = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => mockUseAuthorized(), +})); + +const params: CacheActivityParams = { + startDate: "2026-07-20", + endDate: "2026-07-27", + keyAliases: ["my-key"], + models: ["gpt-5.1"], +}; + +const lastCallOptions = (): { enabled: boolean } => { + const calls = useQueryMock.mock.calls; + return calls[calls.length - 1][3] as { enabled: boolean }; +}; + +describe("useCacheActivity", () => { + beforeEach(() => { + vi.clearAllMocks(); + useQueryMock.mockReturnValue({ data: undefined }); + mockUseAuthorized.mockReturnValue({ accessToken: "test-access-token" }); + }); + + it("queries GET /global/activity/cache_hits with dates and filters as query params", () => { + renderHook(() => useCacheActivity(params)); + + expect(useQueryMock).toHaveBeenCalledWith( + "get", + "/global/activity/cache_hits", + { + params: { + query: { + start_date: "2026-07-20", + end_date: "2026-07-27", + key_aliases: ["my-key"], + models: ["gpt-5.1"], + }, + }, + }, + expect.any(Object), + ); + }); + + it("enables the query when authorized and both dates are set", () => { + renderHook(() => useCacheActivity(params)); + + expect(lastCallOptions().enabled).toBe(true); + }); + + it("disables the query without an access token", () => { + mockUseAuthorized.mockReturnValue({ accessToken: null }); + renderHook(() => useCacheActivity(params)); + + expect(lastCallOptions().enabled).toBe(false); + }); + + it("disables the query while the date range is incomplete", () => { + renderHook(() => useCacheActivity({ ...params, endDate: undefined })); + + expect(lastCallOptions().enabled).toBe(false); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.ts new file mode 100644 index 00000000000..af4486ad33b --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.ts @@ -0,0 +1,32 @@ +import { $api } from "@/lib/http/api"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import type { components } from "@/lib/http/schema"; + +export type CacheActivityResponse = components["schemas"]["CacheActivityResponse"]; +export type CacheActivityGroup = components["schemas"]["CacheActivityGroup"]; + +export interface CacheActivityParams { + startDate: string | undefined; + endDate: string | undefined; + keyAliases: string[]; + models: string[]; +} + +export const useCacheActivity = ({ startDate, endDate, keyAliases, models }: CacheActivityParams) => { + const { accessToken } = useAuthorized(); + return $api.useQuery( + "get", + "/global/activity/cache_hits", + { + params: { + query: { + start_date: startDate ?? "", + end_date: endDate ?? "", + key_aliases: keyAliases, + models, + }, + }, + }, + { enabled: Boolean(accessToken && startDate && endDate) }, + ); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts index f83ebd2622a..f04c4b7bfcd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts @@ -7,6 +7,7 @@ import { selectAutoRouterModelGroups, useAllProxyModels, useAutoRouterModelGroups, + useAutoRouters, useInfiniteModelInfo, useModelHub, useModelsInfo, @@ -113,6 +114,9 @@ describe("useModelsInfo", () => { undefined, undefined, undefined, + // exclude_auto_routers defaults off: only the Models + Endpoints table opts in, so + // every other consumer of this hook keeps seeing auto-routers. + false, ); expect(modelInfoCall).toHaveBeenCalledTimes(1); }); @@ -137,6 +141,9 @@ describe("useModelsInfo", () => { undefined, undefined, undefined, + // exclude_auto_routers defaults off: only the Models + Endpoints table opts in, so + // every other consumer of this hook keeps seeing auto-routers. + false, ); }); @@ -1079,4 +1086,24 @@ describe("useAutoRouterModelGroups", () => { await waitFor(() => expect(modelInfoCall).toHaveBeenCalled()); expect(result.current.size).toBe(0); }); + + // The Auto-Routers tab and the models table read the same /v2/model/info data. Six call + // sites across the app invalidate ["models","list"] after a write; if the auto-router query + // sits in its own namespace, an edit through ModelInfoView leaves the tab stale until a full + // reload, and every future writer has to remember a second key. + describe("auto-router cache namespace", () => { + it("keys the auto-router list under models/list so existing invalidations reach it", async () => { + (modelInfoCall as any).mockResolvedValue(mockPaginatedModelInfoResponse); + const { result } = renderHook(() => useAutoRouters(), { wrapper }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + + const keys = queryClient + .getQueryCache() + .findAll({ queryKey: ["models", "list"] }) + .map((query) => query.queryKey); + + expect(keys.some((key) => JSON.stringify(key).includes("autoRouters"))).toBe(true); + }); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts index ad5e3c91ec3..52459d69b9a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts @@ -1,4 +1,4 @@ -import { useQuery, useInfiniteQuery, UseQueryResult } from "@tanstack/react-query"; +import { useQuery, useInfiniteQuery, useQueryClient, UseQueryResult } from "@tanstack/react-query"; import { createQueryKeys } from "../common/queryKeysFactory"; import { modelInfoCall, modelHubCall, modelAvailableCall } from "@/components/networking"; import useAuthorized from "../useAuthorized"; @@ -24,7 +24,6 @@ export interface PaginatedModelInfoResponse { const modelKeys = createQueryKeys("models"); const modelHubKeys = createQueryKeys("modelHub"); -const autoRouterKeys = createQueryKeys("autoRouterModelGroups"); const allProxyModelsKeys = createQueryKeys("allProxyModels"); const selectedTeamModelsKeys = createQueryKeys("selectedTeamModels"); const infiniteModelKeys = createQueryKeys("infiniteModels"); @@ -38,6 +37,7 @@ export const useModelsInfo = ( teamId?: string, sortBy?: string, sortOrder?: string, + excludeAutoRouters: boolean = false, ) => { const { accessToken, userId, userRole } = useAuthorized(); return useQuery({ @@ -52,10 +52,25 @@ export const useModelsInfo = ( ...(teamId && { teamId }), ...(sortBy && { sortBy }), ...(sortOrder && { sortOrder }), + // Part of the key: callers that exclude auto-routers must not share a cache entry + // with callers that keep them. + ...(excludeAutoRouters && { excludeAutoRouters: "true" }), }, }), queryFn: async () => - await modelInfoCall(accessToken!, userId!, userRole!, page, size, search, modelId, teamId, sortBy, sortOrder), + await modelInfoCall( + accessToken!, + userId!, + userRole!, + page, + size, + search, + modelId, + teamId, + sortBy, + sortOrder, + excludeAutoRouters, + ), enabled: Boolean(accessToken && userId && userRole), }); }; @@ -69,6 +84,30 @@ export interface AutoRouterCandidateDeployment { litellm_params?: { model?: string | null } | null; } +export interface AutoRouterDeployment extends AutoRouterCandidateDeployment { + litellm_params?: { + model?: string | null; + complexity_router_config?: unknown; + complexity_router_default_model?: string | null; + auto_router_config?: unknown; + auto_router_default_model?: string | null; + auto_router_embedding_model?: string | null; + adaptive_router_config?: unknown; + adaptive_router_default_model?: string | null; + quality_router_config?: unknown; + quality_router_default_model?: string | null; + } | null; + model_info?: { + id?: string | null; + /** False for config.yaml-defined deployments, which the update and delete routes refuse. */ + db_model?: boolean | null; + created_at?: string | null; + updated_at?: string | null; + team_id?: string | null; + created_by?: string | null; + } | null; +} + export const isAutoRouterDeployment = (deployment: AutoRouterCandidateDeployment): boolean => Boolean(deployment?.litellm_params?.model?.startsWith(AUTO_ROUTER_MODEL_PREFIX)); @@ -80,11 +119,14 @@ export const selectAutoRouterModelGroups = (deployments: AutoRouterCandidateDepl .filter((modelName): modelName is string => Boolean(modelName)), ); +export const selectAutoRouterDeployments = (deployments: AutoRouterDeployment[]): AutoRouterDeployment[] => + deployments.filter(isAutoRouterDeployment); + const fetchAllModelDeployments = async ( accessToken: string, userId: string, userRole: string, -): Promise => { +): Promise => { const firstPage: PaginatedModelInfoResponse = await modelInfoCall( accessToken, userId, @@ -100,18 +142,28 @@ const fetchAllModelDeployments = async ( ); return [firstPage, ...remainingPages].flatMap( (page: PaginatedModelInfoResponse) => page?.data ?? [], - ) as AutoRouterCandidateDeployment[]; + ) as AutoRouterDeployment[]; }; +/** + * Deliberately under the same `models/list` namespace as useModelsInfo: it is the same + * /v2/model/info data, and every writer in the app already invalidates ["models","list"]. + * A private namespace meant an edit through ModelInfoView left this list stale, and every + * future writer would have had to remember a second key. + */ +const autoRouterListKey = (userId: string | null, userRole: string | null) => + modelKeys.list({ + filters: { + scope: "autoRouters", + ...(userId && { userId }), + ...(userRole && { userRole }), + }, + }); + export const useAutoRouterModelGroups = (): ReadonlySet => { const { accessToken, userId, userRole } = useAuthorized(); - const { data } = useQuery>({ - queryKey: autoRouterKeys.list({ - filters: { - ...(userId && { userId }), - ...(userRole && { userRole }), - }, - }), + const { data } = useQuery>({ + queryKey: autoRouterListKey(userId, userRole), queryFn: async () => await fetchAllModelDeployments(accessToken!, userId!, userRole!), enabled: Boolean(accessToken && userId && userRole), select: selectAutoRouterModelGroups, @@ -119,6 +171,23 @@ export const useAutoRouterModelGroups = (): ReadonlySet => { return data ?? NO_AUTO_ROUTERS; }; +export const useAutoRouters = (): UseQueryResult => { + const { accessToken, userId, userRole } = useAuthorized(); + return useQuery({ + queryKey: autoRouterListKey(userId, userRole), + queryFn: async () => await fetchAllModelDeployments(accessToken!, userId!, userRole!), + enabled: Boolean(accessToken && userId && userRole), + select: selectAutoRouterDeployments, + }); +}; + +export const useInvalidateAutoRouters = (): (() => Promise) => { + const queryClient = useQueryClient(); + return async () => { + await queryClient.invalidateQueries({ queryKey: modelKeys.lists() }); + }; +}; + export const useModelHub = () => { const { accessToken } = useAuthorized(); return useQuery({ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx index 8dffc80a70e..5650bd1d7e4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx @@ -190,7 +190,7 @@ const OAuthFormFields: React.FC = ({ label={ } name="issuer" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx index 1dc7736d5ac..1d206f81030 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx @@ -104,6 +104,9 @@ const AllModelsTab = ({ teamIdForQuery, sortBy, sortOrder, + // Auto-routers are routing constructs, not deployments; the sibling Auto-Routers tab + // lists and manages them. Excluded server-side so total_count stays honest. + true, ); const isLoading = isLoadingModelsInfo || isLoadingModelCostMap; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersPanel.test.tsx new file mode 100644 index 00000000000..9ec551bc227 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersPanel.test.tsx @@ -0,0 +1,260 @@ +import userEvent from "@testing-library/user-event"; +import { describe, expect, it, vi } from "vitest"; + +import { renderWithProviders, screen, testQueryClient, waitFor } from "@/../tests/test-utils"; + +import { AutoRoutersPanel } from "./AutoRoutersPanel"; + +const { modelInfoCall, modelDeleteCall } = vi.hoisted(() => ({ + modelInfoCall: vi.fn(), + modelDeleteCall: vi.fn().mockResolvedValue({}), +})); + +vi.mock("@/components/networking", () => ({ + modelInfoCall, + modelDeleteCall, + modelHubCall: vi.fn(), + modelAvailableCall: vi.fn().mockResolvedValue({ data: [] }), +})); + +vi.mock("@/components/llm_calls/fetch_models", () => ({ + fetchAvailableModels: vi.fn().mockResolvedValue([]), +})); + +const { openModel } = vi.hoisted(() => ({ openModel: vi.fn() })); + +vi.mock("@/app/(dashboard)/models-and-endpoints/detailNavigation", () => ({ + useModelDetailRouting: () => ({ openModel, modelId: null, teamId: null, openTeam: vi.fn(), close: vi.fn() }), +})); + +vi.mock("@/components/edit_auto_router/edit_auto_router_modal", () => ({ + __esModule: true, + default: ({ modelData }: { modelData: { model_name?: string; model_info?: { id?: string } } }) => ( +
+ edit:{modelData.model_name}:{modelData.model_info?.id} +
+ ), +})); + +vi.mock("@/components/add_model/add_auto_router_tab", () => ({ + __esModule: true, + default: ({ handleOk }: { handleOk: () => void }) => ( + + ), +})); + +// A realistic /v2/model/info page: two auto-routers among ordinary deployments. The panel must +// render exactly the auto_router/* rows; a view that renders page.data unfiltered passes a +// "renders a table" assertion but fails this one. +const DEPLOYMENTS = [ + { + // DB-created adaptive router: no editor for its shape, but it must stay deletable, since + // auto-routers are excluded from Models + Endpoints and this tab is the only delete path. + model_name: "adaptive-router", + litellm_params: { model: "auto_router/adaptive_router" }, + model_info: { id: "auto-3", db_model: true }, + }, + { + // config.yaml row: the API refuses both update and delete, so neither control may appear. + model_name: "config-router", + litellm_params: { + model: "auto_router/complexity_router", + complexity_router_config: { tiers: {}, classifier_type: "llm" }, + }, + model_info: { id: "auto-4", db_model: false }, + }, + { + model_name: "gpt-4o-mini", + litellm_params: { model: "openai/gpt-4o-mini" }, + model_info: { id: "plain-1" }, + }, + { + model_name: "tri-tier-router", + litellm_params: { + model: "auto_router/complexity_router", + complexity_router_config: { tiers: { SIMPLE: ["gpt-4o-mini"] }, classifier_type: "heuristic" }, + complexity_router_default_model: "gpt-4o-mini", + }, + model_info: { id: "auto-1", db_model: true, created_at: "2026-07-28T21:40:09.900000+00:00" }, + }, + { + model_name: "anthropic-opus-4-6", + litellm_params: { model: "anthropic/claude-opus-4-6" }, + model_info: { id: "plain-2" }, + }, + { + model_name: "support-router", + litellm_params: { + model: "auto_router/support-router", + auto_router_config: JSON.stringify({ routes: [{ name: "gpt-4o-mini" }] }), + auto_router_default_model: "gpt-4o-mini", + }, + model_info: { id: "auto-2", db_model: true, created_at: "2026-07-27T10:00:00.000000+00:00" }, + }, +]; + +const pageOf = (data: typeof DEPLOYMENTS) => ({ + data, + total_count: data.length, + current_page: 1, + total_pages: 1, + size: 1000, +}); + +const mockDeploymentsPage = () => { + modelInfoCall.mockResolvedValue(pageOf(DEPLOYMENTS)); +}; + +const renderPanel = (canModify = true) => + renderWithProviders( + , + ); + +describe("AutoRoutersPanel", () => { + beforeEach(() => { + // The shared test client caches with staleTime: Infinity and refetchOnMount: false, so + // without this every test after the first reads the previous test's deployment page. + testQueryClient.clear(); + modelInfoCall.mockReset(); + modelDeleteCall.mockClear(); + openModel.mockClear(); + mockDeploymentsPage(); + }); + + it("lists only auto_router deployments, not every model on the proxy", async () => { + renderPanel(); + + expect(await screen.findByText("tri-tier-router")).toBeInTheDocument(); + expect(await screen.findByText("support-router")).toBeInTheDocument(); + expect(screen.queryByText("gpt-4o-mini", { selector: "span.text-sm.font-medium" })).not.toBeInTheDocument(); + expect(screen.queryByText("anthropic-opus-4-6", { selector: "span.text-sm.font-medium" })).not.toBeInTheDocument(); + }); + + it("labels Type by classifier rather than by router family", async () => { + renderPanel(); + + expect(await screen.findByText("Heuristic")).toBeInTheDocument(); + expect(await screen.findByText("Semantic")).toBeInTheDocument(); + }); + + // Reuses the models-page drill-in, so an auto router opens the full ModelInfoView with + // Model Settings and Edit Settings, not a parallel detail view that reimplements part of it. + it("opens the shared model detail view on row click", async () => { + const user = userEvent.setup(); + renderPanel(); + + await user.click(await screen.findByRole("button", { name: "support-router" })); + + expect(openModel).toHaveBeenCalledWith("auto-2"); + }); + + it("opens the create form in a dialog and refetches the list after a create", async () => { + const user = userEvent.setup(); + renderPanel(); + + await screen.findByText("tri-tier-router"); + const callsBeforeCreate = modelInfoCall.mock.calls.length; + + expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: "Add Auto Router" })); + + // A dialog, not a full-panel swap: the list stays mounted behind it. + const dialog = await screen.findByRole("dialog"); + expect(dialog).toHaveTextContent("Add Auto Router"); + expect(screen.getByText("tri-tier-router")).toBeInTheDocument(); + + await user.click(await screen.findByRole("button", { name: "Submit auto router" })); + + // Back on the list, and the deployment query was invalidated so a new router shows up + // without a manual page reload. + expect(await screen.findByText("tri-tier-router")).toBeInTheDocument(); + await waitFor(() => expect(modelInfoCall.mock.calls.length).toBeGreaterThan(callsBeforeCreate)); + }); + + // The page decides who may write (proxy admin or team admin); the panel just has to make + // every write affordance absent when told no, rather than let a submit 403 later. Reading + // stays open: a read-only caller can still drill into the detail view. + it("shows the list but no write affordances when canModify is false", async () => { + renderPanel(false); + + expect(await screen.findByText("tri-tier-router")).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Add Auto Router" })).not.toBeInTheDocument(); + expect(screen.queryByTestId("auto-router-actions-auto-1")).not.toBeInTheDocument(); + // Still navigable, because opening the detail view is a read. + expect(screen.getByRole("button", { name: "tri-tier-router" })).toBeInTheDocument(); + }); + + // Auto-routers are hidden from Models + Endpoints, which used to be the only route to the + // delete action, so this tab is now the only place an auto router can be removed. + it("deletes the chosen router by its model id and refetches", async () => { + const user = userEvent.setup(); + renderPanel(); + + await screen.findByText("support-router"); + const callsBeforeDelete = modelInfoCall.mock.calls.length; + + await user.click(screen.getByTestId("auto-router-actions-auto-2")); + await user.click(await screen.findByTestId("auto-router-action-delete")); + await user.click(await screen.findByRole("button", { name: /^delete$/i })); + + await waitFor(() => expect(modelDeleteCall).toHaveBeenCalledWith("token", "auto-2")); + await waitFor(() => expect(modelInfoCall.mock.calls.length).toBeGreaterThan(callsBeforeDelete)); + }); + + it("does not delete when the confirmation is dismissed", async () => { + const user = userEvent.setup(); + renderPanel(); + + await screen.findByText("support-router"); + + await user.click(screen.getByTestId("auto-router-actions-auto-2")); + await user.click(await screen.findByTestId("auto-router-action-delete")); + await user.click(await screen.findByRole("button", { name: /cancel/i })); + + expect(modelDeleteCall).not.toHaveBeenCalled(); + }); + + it("gives a read-only caller no delete affordance", async () => { + renderPanel(false); + + await screen.findByText("support-router"); + expect(screen.queryByTestId("auto-router-actions-auto-2")).not.toBeInTheDocument(); + }); + + it("renders an empty state when the proxy has models but no auto routers", async () => { + modelInfoCall.mockResolvedValue( + pageOf(DEPLOYMENTS.filter((d) => !d.litellm_params.model.startsWith("auto_router/"))), + ); + + renderPanel(); + + expect(await screen.findByText("No auto routers yet")).toBeInTheDocument(); + }); + + it("keeps delete available on a DB-created adaptive router that has no editor", async () => { + const user = userEvent.setup(); + renderPanel(); + + await screen.findByText("adaptive-router"); + await user.click(screen.getByTestId("auto-router-actions-auto-3")); + await user.click(await screen.findByTestId("auto-router-action-delete")); + await user.click(await screen.findByRole("button", { name: /^delete$/i })); + + await waitFor(() => expect(modelDeleteCall).toHaveBeenCalledWith("token", "auto-3")); + }); + + it("offers no delete on a config-defined router, which the API would refuse", async () => { + renderPanel(); + + await screen.findByText("config-router"); + expect(screen.queryByTestId("auto-router-actions-auto-4")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersPanel.tsx new file mode 100644 index 00000000000..f27c1e1c44a --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersPanel.tsx @@ -0,0 +1,130 @@ +"use client"; + +import { Plus } from "lucide-react"; +import { useMemo, useState } from "react"; + +import { useAutoRouters, useInvalidateAutoRouters } from "@/app/(dashboard)/hooks/models/useModels"; +import { useModelDetailRouting } from "@/app/(dashboard)/models-and-endpoints/detailNavigation"; +import AddAutoRouterTab from "@/components/add_model/add_auto_router_tab"; +import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; +import NotificationsManager from "@/components/molecules/notifications_manager"; +import { modelDeleteCall } from "@/components/networking"; +import { Button } from "@/components/ui/button"; +import { Dialog, DialogContent, DialogDescription, DialogHeader, DialogTitle } from "@/components/ui/dialog"; +import { type ModelWriteScope } from "@/utils/modelPermissions"; +import { Team } from "@/components/networking"; + +import { AutoRoutersTable } from "./AutoRoutersTable"; +import { AutoRouterRow, toAutoRouterRows } from "./autoRouterRows"; + +interface AutoRoutersPanelProps { + accessToken: string; + userRole: string; + userID: string | null; + teams: Team[] | null; + /** Owned by the page, which knows how this caller must scope what they create. */ + createScope: ModelWriteScope; +} + +export function AutoRoutersPanel({ accessToken, userRole, userID, teams, createScope }: AutoRoutersPanelProps) { + const canCreate = createScope !== "forbidden"; + const { data: deployments, isLoading } = useAutoRouters(); + const invalidateAutoRouters = useInvalidateAutoRouters(); + // Clicking a router opens the same ?model= drill-in the All Models table uses, so an auto + // router gets the full ModelInfoView: Model Settings, Edit Settings, Edit Auto Router and + // Delete. A separate detail view here would be a worse copy of it. + const { openModel } = useModelDetailRouting(); + const [isCreating, setIsCreating] = useState(false); + const [deletingRouter, setDeletingRouter] = useState(null); + const [isDeleting, setIsDeleting] = useState(false); + + const routers = useMemo( + () => toAutoRouterRows(deployments ?? [], { userRole, userID }, teams), + [deployments, userRole, userID, teams], + ); + + const handleCreated = () => { + setIsCreating(false); + void invalidateAutoRouters(); + }; + + const handleConfirmDelete = async () => { + if (!deletingRouter) return; + setIsDeleting(true); + try { + await modelDeleteCall(accessToken, deletingRouter.id); + NotificationsManager.success(`Deleted auto router: ${deletingRouter.name}`); + setDeletingRouter(null); + await invalidateAutoRouters(); + } catch (error) { + NotificationsManager.fromBackend(`Failed to delete auto router: ${error}`); + } finally { + setIsDeleting(false); + } + }; + + return ( +
+
+
+

Auto routers

+

+ Auto routers sit above your deployments and pick a model per request. They are called like any other model, + so clients keep using a single model name. +

+
+ {canCreate && ( + + )} +
+ + openModel(row.id)} + onDeleteClick={setDeletingRouter} + /> + + + {/* The form is long, so the dialog caps its height and scrolls its body rather than + growing past the viewport. */} + + + Add Auto Router + + Routes each request to a model by classifying its complexity. Called like any other model, so clients keep + using a single model name. + + + + + + + {deletingRouter && ( + setDeletingRouter(null)} + onOk={handleConfirmDelete} + confirmLoading={isDeleting} + /> + )} +
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersTable.tsx new file mode 100644 index 00000000000..943388f8535 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersTable.tsx @@ -0,0 +1,68 @@ +"use client"; + +import { SortingState } from "@tanstack/react-table"; +import { useMemo, useState } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; +import { AutoRouterIcon } from "@/components/shared/table_cells"; + +import { getAutoRoutersTableColumns } from "./AutoRoutersTableColumns"; +import { AutoRouterRow } from "./autoRouterRows"; + +interface AutoRoutersTableProps { + routers: AutoRouterRow[]; + isLoading: boolean; + canModify: boolean; + onRouterClick: (row: AutoRouterRow) => void; + onDeleteClick: (row: AutoRouterRow) => void; +} + +const PAGE_SIZE_OPTIONS = [10, 25, 50]; + +function EmptyState({ canModify }: { canModify: boolean }) { + return ( +
+
+ +
+
No auto routers yet
+
+ {canModify + ? "Create an auto router to pick the right model per request instead of pinning one." + : "An auto router picks the right model per request instead of pinning one."} +
+
+ ); +} + +export function AutoRoutersTable({ + routers, + isLoading, + canModify, + onRouterClick, + onDeleteClick, +}: AutoRoutersTableProps) { + const [sorting, setSorting] = useState([]); + + const columns = useMemo( + () => getAutoRoutersTableColumns({ canModify, onRouterClick, onDeleteClick }), + [canModify, onRouterClick, onDeleteClick], + ); + + return ( + router.id} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + paginationMode="client" + pageSizeOptions={PAGE_SIZE_OPTIONS} + isLoading={isLoading} + loadingMessage="Loading auto routers…" + noDataMessage={} + size="compact" + /> + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersTableColumns.tsx new file mode 100644 index 00000000000..995ba634c34 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/AutoRoutersTableColumns.tsx @@ -0,0 +1,173 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { useEffect, useMemo, useRef, useState } from "react"; +import { MoreHorizontal, Trash2 } from "lucide-react"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { DateCell, IdentityCell } from "@/components/shared/table_cells"; +import { Badge } from "@/components/ui/badge"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; + +import { AutoRouterRow } from "./autoRouterRows"; +import { fitPills } from "./fitPills"; + +function TypeCell({ row }: { row: AutoRouterRow }) { + return ( + + {row.typeLabel} + + ); +} + +function TargetsCell({ targets }: { targets: string[] }) { + const containerRef = useRef(null); + const [width, setWidth] = useState(0); + + useEffect(() => { + const node = containerRef.current; + if (!node || typeof ResizeObserver === "undefined") return; + const observer = new ResizeObserver((entries) => { + const measured = entries[0]?.contentRect.width; + if (typeof measured === "number") setWidth(measured); + }); + observer.observe(node); + return () => observer.disconnect(); + }, []); + + const { visible, overflow } = useMemo(() => fitPills(targets, width), [targets, width]); + + if (targets.length === 0) { + return -; + } + + return ( +
+ {visible.map((target) => ( + + {target} + + ))} + {overflow > 0 && ( + + +{overflow} + + )} +
+ ); +} + +function AutoRouterRowActions({ + row, + onDeleteClick, +}: { + row: AutoRouterRow; + onDeleteClick: (row: AutoRouterRow) => void; +}) { + return ( + + + + + + onDeleteClick(row)} + > + + Delete auto router + + + + ); +} + +interface AutoRoutersTableColumnsDeps { + canModify: boolean; + onRouterClick: (row: AutoRouterRow) => void; + onDeleteClick: (row: AutoRouterRow) => void; +} + +export const getAutoRoutersTableColumns = ({ + canModify, + onRouterClick, + onDeleteClick, +}: AutoRoutersTableColumnsDeps): ColumnDef[] => [ + { + id: "name", + accessorKey: "name", + meta: { title: "Name" }, + header: ({ column }) => , + size: 260, + enableSorting: true, + cell: ({ row }) => onRouterClick(row.original)} />, + }, + { + id: "kind", + accessorKey: "kind", + meta: { title: "Type" }, + header: "Type", + size: 180, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "targets", + meta: { title: "Routes to" }, + header: "Routes to", + size: 320, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "defaultModel", + accessorKey: "defaultModel", + meta: { title: "Default model" }, + header: "Default model", + size: 200, + enableSorting: false, + cell: ({ row }) => + row.original.defaultModel ? ( + + {row.original.defaultModel} + + ) : ( + - + ), + }, + { + id: "createdAt", + accessorKey: "createdAt", + meta: { title: "Created" }, + header: ({ column }) => , + size: 150, + enableSorting: true, + sortingFn: "datetime", + cell: ({ row }) => , + }, + ...(canModify + ? [ + { + id: "actions", + meta: { title: "" }, + header: "", + size: 60, + enableSorting: false, + cell: ({ row }) => + row.original.canDelete ? : null, + } satisfies ColumnDef, + ] + : []), +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts new file mode 100644 index 00000000000..9944653b638 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts @@ -0,0 +1,262 @@ +import { describe, expect, it } from "vitest"; + +import { autoRouterStrategy, isComplexityRouter } from "@/components/add_model/auto_router_strategies"; +import { toAutoRouterRow, toAutoRouterRows } from "./autoRouterRows"; + +// Existing cases assert resource classification, so they run as a proxy admin: the actor +// gate is then a pass-through and canEdit/canDelete still reflect the row itself. +const ADMIN = { userRole: "Admin", userID: "u-admin" }; +const TEAM_ADMIN = { userRole: "Internal User", userID: "u-team-admin" }; + +const complexityDeployment = { + model_name: "tri-tier-router", + litellm_params: { + model: "auto_router/complexity_router", + complexity_router_config: { + tiers: { + SIMPLE: ["gpt-4o-mini"], + MEDIUM: ["anthropic-sonnet-4-6"], + COMPLEX: ["anthropic-opus-4-6", "gpt-4o-mini"], + REASONING: [], + }, + classifier_type: "heuristic", + }, + complexity_router_default_model: "gpt-4o-mini", + }, + model_info: { id: "cid-1", db_model: true, created_at: "2026-07-28T21:40:09.900000+00:00" }, +}; + +const semanticDeployment = { + model_name: "support-router", + litellm_params: { + model: "auto_router/support-router", + auto_router_config: JSON.stringify({ + routes: [ + { name: "gpt-4o-mini", utterances: ["reset my password"] }, + { name: "anthropic-opus-4-6", utterances: ["design a distributed system"] }, + ], + }), + auto_router_default_model: "gpt-4o-mini", + }, + model_info: { id: "sid-1", db_model: true, created_at: "2026-07-27T10:00:00.000000+00:00" }, +}; + +describe("autoRouterRows", () => { + it("classifies a complexity router and unions its tier models as targets", () => { + const row = toAutoRouterRow(complexityDeployment, 0, ADMIN, null); + + expect(row.kind).toBe("complexity"); + expect(row.typeLabel).toBe("Heuristic"); + // Union across tiers, de-duplicated: gpt-4o-mini appears in both SIMPLE and COMPLEX. + expect(row.targets).toEqual(["gpt-4o-mini", "anthropic-sonnet-4-6", "anthropic-opus-4-6"]); + expect(row.defaultModel).toBe("gpt-4o-mini"); + expect(row.id).toBe("cid-1"); + }); + + it("parses a semantic router whose config arrives as a JSON string", () => { + const row = toAutoRouterRow(semanticDeployment, 0, ADMIN, null); + + expect(row.kind).toBe("semantic"); + expect(row.typeLabel).toBe("Semantic"); + expect(row.targets).toEqual(["gpt-4o-mini", "anthropic-opus-4-6"]); + expect(row.defaultModel).toBe("gpt-4o-mini"); + }); + + it("shows a tier pinned as a bare string, which the backend accepts as `str | list[str]`", () => { + const row = toAutoRouterRow( + { + ...complexityDeployment, + litellm_params: { + ...complexityDeployment.litellm_params, + complexity_router_config: { + tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: ["anthropic-sonnet-4-6"], COMPLEX: "", REASONING: [] }, + classifier_type: "heuristic", + }, + }, + }, + 0, + ADMIN, + null, + ); + + expect(row.targets).toEqual(["gpt-4o-mini", "anthropic-sonnet-4-6"]); + }); + + it("labels a router using the LLM classifier", () => { + const row = toAutoRouterRow( + { + ...complexityDeployment, + litellm_params: { + ...complexityDeployment.litellm_params, + complexity_router_config: { tiers: {}, classifier_type: "llm", adaptive: true }, + }, + }, + 0, + ADMIN, + null, + ); + + expect(row.typeLabel).toBe("LLM Classifier"); + }); + + it("treats a deployment carrying complexity_router_config as complexity even off the canonical model string", () => { + expect(isComplexityRouter({ model: "auto_router/legacy", complexity_router_config: { tiers: {} } })).toBe(true); + }); + + it("survives an unparseable config instead of throwing", () => { + const row = toAutoRouterRow( + { + model_name: "broken", + litellm_params: { model: "auto_router/broken", auto_router_config: "{not json" }, + model_info: { id: "bid-1" }, + }, + 0, + ADMIN, + null, + ); + + expect(row.kind).toBe("semantic"); + expect(row.targets).toEqual([]); + }); + + it("falls back to a stable synthetic id when the deployment has no model_info id", () => { + const rows = toAutoRouterRows( + [ + { model_name: "a", litellm_params: { model: "auto_router/a" } }, + { model_name: "b", litellm_params: { model: "auto_router/b" } }, + ], + ADMIN, + null, + ); + + expect(rows.map((row) => row.id)).toEqual(["a-0", "b-1"]); + }); + // Regression: adaptive and quality routers used to fall through to the semantic branch, + // which read the wrong config key and reported an empty route list and a null default. + it("classifies an adaptive router as adaptive, not semantic", () => { + const row = toAutoRouterRow( + { + model_name: "smart-router", + litellm_params: { + model: "auto_router/adaptive_router", + adaptive_router_default_model: "gpt-4o-mini", + adaptive_router_config: { available_models: ["gpt-4o", "gpt-4o-mini"] }, + }, + model_info: { id: "ad-1" }, + }, + 0, + ADMIN, + null, + ); + + expect(row.kind).toBe("adaptive"); + expect(row.typeLabel).toBe("Adaptive"); + expect(row.targets).toEqual(["gpt-4o", "gpt-4o-mini"]); + expect(row.defaultModel).toBe("gpt-4o-mini"); + }); + + it("classifies a quality router as quality, not semantic", () => { + const row = toAutoRouterRow( + { + model_name: "quality-router", + litellm_params: { + model: "auto_router/quality_router", + quality_router_default_model: "gpt-4o", + quality_router_config: { available_models: ["gpt-4o"] }, + }, + model_info: { id: "q-1" }, + }, + 0, + ADMIN, + null, + ); + + expect(row.kind).toBe("quality"); + expect(row.typeLabel).toBe("Quality"); + expect(row.targets).toEqual(["gpt-4o"]); + }); + + it("mirrors the backend prefix ordering, so a named strategy never reads as semantic", () => { + const kindOf = (model: string) => autoRouterStrategy({ model }).kind; + expect(kindOf("auto_router/complexity_router")).toBe("complexity"); + expect(kindOf("auto_router/adaptive_router")).toBe("adaptive"); + expect(kindOf("auto_router/quality_router")).toBe("quality"); + expect(kindOf("auto_router/my-own-router")).toBe("semantic"); + }); + + // The capability matrix. Origin and strategy constrain DIFFERENT capabilities, and + // collapsing them into one "editable" flag is what stranded DB-created adaptive routers + // with no delete control. Live-verified: for a config row PATCH /model/{id}/update 404s + // and POST /model/delete 400s. + const rowFor = (model: string, dbModel: boolean) => + toAutoRouterRow( + { model_name: "r", litellm_params: { model }, model_info: { id: "x", db_model: dbModel } }, + 0, + ADMIN, + null, + ); + + it.each([ + { model: "auto_router/complexity_router", db: true, canEdit: true, canDelete: true, reason: null }, + { model: "auto_router/my-semantic", db: true, canEdit: true, canDelete: true, reason: null }, + // No editor for its shape, but deleting never reads the config, so delete stays. + { model: "auto_router/adaptive_router", db: true, canEdit: false, canDelete: true, reason: "no-editor" }, + { model: "auto_router/quality_router", db: true, canEdit: false, canDelete: true, reason: "no-editor" }, + // config.yaml rows: the API refuses both, whatever the strategy. + { model: "auto_router/complexity_router", db: false, canEdit: false, canDelete: false, reason: "config-managed" }, + { model: "auto_router/adaptive_router", db: false, canEdit: false, canDelete: false, reason: "config-managed" }, + ])("$model (db_model=$db) -> canEdit=$canEdit canDelete=$canDelete", (spec) => { + const row = rowFor(spec.model, spec.db); + expect(row.canEdit).toBe(spec.canEdit); + expect(row.canDelete).toBe(spec.canDelete); + expect(row.editBlockedReason).toBe(spec.reason); + }); + + it("treats a missing db_model as config-defined rather than assuming it is writable", () => { + const row = toAutoRouterRow({ ...complexityDeployment, model_info: { id: "unknown-1" } }, 0, ADMIN, null); + expect(row.canEdit).toBe(false); + expect(row.canDelete).toBe(false); + }); +}); + +describe("autoRouterRows actor gating", () => { + const TEAMS = [ + { team_id: "team-1", members_with_roles: [{ user_id: "u-team-admin", user_email: "t@t", role: "admin" }] }, + ] as never; + + const rowIn = (actor: { userRole: string; userID: string }, teamId: string | null) => + toAutoRouterRow( + { ...complexityDeployment, model_info: { id: "cid-1", db_model: true, team_id: teamId } }, + 0, + actor, + TEAMS, + ); + + // Opening the tab to team admins puts rows they cannot act on in the same list: other + // teams' routers, and the proxy-level unscoped ones. PATCH and DELETE both 403 those, so + // the affordance has to be per row rather than per tab. + it("hides write affordances on another team's router", () => { + const row = rowIn(TEAM_ADMIN, "other-team"); + expect(row.canEdit).toBe(false); + expect(row.canDelete).toBe(false); + }); + + it("hides them on an unscoped router a proxy admin owns", () => { + const row = rowIn(TEAM_ADMIN, null); + expect(row.canEdit).toBe(false); + expect(row.canDelete).toBe(false); + }); + + // Authorizing on created_by would fail this: the API lets any admin of the owning team act. + it("keeps them on the team's router regardless of who created it", () => { + const row = rowIn(TEAM_ADMIN, "team-1"); + expect(row.canEdit).toBe(true); + expect(row.canDelete).toBe(true); + }); + + it("lets a proxy admin act on any team's router", () => { + const row = rowIn(ADMIN, "other-team"); + expect(row.canEdit).toBe(true); + expect(row.canDelete).toBe(true); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts new file mode 100644 index 00000000000..35172d67e84 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts @@ -0,0 +1,121 @@ +import { AutoRouterDeployment } from "@/app/(dashboard)/hooks/models/useModels"; +import { + AutoRouterKind, + EditBlockedReason, + autoRouterCapabilities, + autoRouterStrategy, +} from "@/components/add_model/auto_router_strategies"; +import { normalizeTierModels } from "@/components/add_model/complexity_router_tiers"; +import { Team } from "@/components/networking"; +import { type ModelActor, canModifyModel } from "@/utils/modelPermissions"; + +export type { AutoRouterKind }; + +/** Who is looking at the list; decides which rows offer write affordances. */ +export type AutoRouterActor = ModelActor; + +export interface AutoRouterRow { + id: string; + name: string; + kind: AutoRouterKind; + typeLabel: string; + /** Edit needs an API-created row AND a strategy the dashboard has a form for. */ + canEdit: boolean; + /** + * Resource capability ANDed with the caller's standing on this specific row. A team admin + * sees rows they cannot delete (another team's, or one a teammate created), and the API + * would 403 those, so the affordance has to be per row rather than per tab. + */ + canDelete: boolean; + editBlockedReason: EditBlockedReason | null; + targets: string[]; + defaultModel: string | null; + createdAt: string | null; + deployment: AutoRouterDeployment; +} + +const safeParse = (value: string): unknown => { + try { + return JSON.parse(value); + } catch { + return null; + } +}; + +const asRecord = (value: unknown): Record => { + const parsed: unknown = typeof value === "string" ? safeParse(value) : value; + return typeof parsed === "object" && parsed !== null && !Array.isArray(parsed) + ? (parsed as Record) + : {}; +}; + +const asStringArray = (value: unknown): string[] => + Array.isArray(value) ? value.filter((entry): entry is string => typeof entry === "string") : []; + +const dedupe = (models: string[]): string[] => Array.from(new Set(models)); + +export const complexityTypeLabel = (config: Record): string => + config.classifier_type === "llm" ? "LLM Classifier" : "Heuristic"; + +interface Presentation { + typeLabel: string; + targets: string[]; +} + +// Adaptive and quality both declare a flat pool and have no editor here, so the row reports +// what is configured rather than interpreting it. +const configManaged = (label: string, config: Record): Presentation => ({ + typeLabel: label, + targets: asStringArray(config.available_models), +}); + +/** How each strategy renders itself, given its own config object. */ +const PRESENTERS: Record) => Presentation> = { + complexity: (config) => ({ + typeLabel: complexityTypeLabel(config), + targets: dedupe(Object.values(asRecord(config.tiers)).flatMap(normalizeTierModels)), + }), + semantic: (config) => { + const routes = dedupe( + (Array.isArray(config.routes) ? config.routes : []) + .map((route) => asRecord(route).name) + .filter((name): name is string => typeof name === "string" && name.length > 0), + ); + return { typeLabel: "Semantic", targets: routes }; + }, + adaptive: (config) => configManaged("Adaptive", config), + quality: (config) => configManaged("Quality", config), +}; + +export const toAutoRouterRow = ( + deployment: AutoRouterDeployment, + index: number, + actor: AutoRouterActor, + teams: Team[] | null, +): AutoRouterRow => { + const params = deployment.litellm_params ?? {}; + const info = deployment.model_info ?? {}; + const name = deployment.model_name ?? ""; + const strategy = autoRouterStrategy(params); + const { canEdit, canDelete, editBlockedReason } = autoRouterCapabilities(params, info); + const mayActOnRow = canModifyModel(actor, teams, { teamId: info.team_id, isDbModel: info.db_model === true }); + + return { + id: info.id ?? `${name}-${index}`, + name, + kind: strategy.kind, + canEdit: canEdit && mayActOnRow, + canDelete: canDelete && mayActOnRow, + editBlockedReason, + createdAt: info.created_at ?? null, + defaultModel: (params[strategy.defaultModelKey] as string | null | undefined) ?? null, + deployment, + ...PRESENTERS[strategy.kind](asRecord(params[strategy.configKey])), + }; +}; + +export const toAutoRouterRows = ( + deployments: AutoRouterDeployment[], + actor: AutoRouterActor, + teams: Team[] | null, +): AutoRouterRow[] => deployments.map((deployment, index) => toAutoRouterRow(deployment, index, actor, teams)); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/fitPills.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/fitPills.test.ts new file mode 100644 index 00000000000..9d702810c35 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/fitPills.test.ts @@ -0,0 +1,43 @@ +import { describe, expect, it } from "vitest"; + +import { fitPills, pillWidth } from "./fitPills"; + +const TARGETS = ["anthropic-sonnet-4-6", "gpt-4o-mini", "anthropic-opus-4-6", "voyage-4-large"]; + +describe("fitPills", () => { + it("keeps everything on one row when it all fits", () => { + const wide = TARGETS.reduce((total, label) => total + pillWidth(label) + 4, 0) + 40; + expect(fitPills(TARGETS, wide)).toEqual({ visible: TARGETS, overflow: 0 }); + }); + + it("shows more pills as the column gets wider", () => { + const narrow = fitPills(TARGETS, 200); + const wider = fitPills(TARGETS, 420); + + expect(narrow.visible.length).toBeLessThan(wider.visible.length); + expect(narrow.visible.length + narrow.overflow).toBe(TARGETS.length); + expect(wider.visible.length + wider.overflow).toBe(TARGETS.length); + }); + + it("reserves room for the +N counter so the row never overflows", () => { + const { visible } = fitPills(TARGETS, 220); + const used = visible.reduce((total, label, index) => total + pillWidth(label) + (index === 0 ? 0 : 4), 0); + // 28px counter + its 4px gap must still fit alongside the visible pills. + expect(used + 32).toBeLessThanOrEqual(220); + }); + + it("always shows at least one pill, even when a single name is wider than the column", () => { + expect(fitPills(["an-extremely-long-deployment-name-that-never-fits"], 40)).toEqual({ + visible: ["an-extremely-long-deployment-name-that-never-fits"], + overflow: 0, + }); + }); + + it("shows one pill before the first measurement rather than flashing every pill", () => { + expect(fitPills(TARGETS, 0)).toEqual({ visible: [TARGETS[0]], overflow: 3 }); + }); + + it("handles an empty target list", () => { + expect(fitPills([], 300)).toEqual({ visible: [], overflow: 0 }); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/fitPills.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/fitPills.ts new file mode 100644 index 00000000000..fba116af0fd --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/fitPills.ts @@ -0,0 +1,52 @@ +/** + * How many pills fit on ONE row of a given width, leaving room for a "+N" counter. + * + * jsdom reports no layout and the shared ResizeObserver mock only fires inside chart + * subtrees, so this stays a pure width-in / count-out function: the component measures and + * this decides, which keeps the overflow rule unit-testable. + */ + +const CHAR_WIDTH = 6.5; +const PILL_PADDING = 18; +const PILL_GAP = 4; +const OVERFLOW_WIDTH = 28; + +export const pillWidth = (label: string): number => label.length * CHAR_WIDTH + PILL_PADDING; + +export interface FittedPills { + visible: string[]; + overflow: number; +} + +export const fitPills = (labels: string[], availableWidth: number): FittedPills => { + if (labels.length === 0) return { visible: [], overflow: 0 }; + + // Unmeasured (0 or negative) means the first paint before ResizeObserver reports. Show one + // pill rather than all of them, so the row never flashes multi-line and then collapses. + if (availableWidth <= 0) { + return { visible: labels.slice(0, 1), overflow: labels.length - 1 }; + } + + const fitted: string[] = []; + let used = 0; + + for (const [index, label] of labels.entries()) { + const remaining = labels.length - index - 1; + const gap = fitted.length === 0 ? 0 : PILL_GAP; + // Anything still queued after this pill needs room for the "+N" counter beside it. + const reserve = remaining > 0 ? PILL_GAP + OVERFLOW_WIDTH : 0; + + if (used + gap + pillWidth(label) + reserve > availableWidth) break; + + used += gap + pillWidth(label); + fitted.push(label); + } + + // Always show at least one pill; a single over-long name truncates via CSS instead of + // collapsing the cell to a bare "+N". + if (fitted.length === 0) { + return { visible: labels.slice(0, 1), overflow: labels.length - 1 }; + } + + return { visible: fitted, overflow: labels.length - fitted.length }; +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.test.tsx index b9dd09d5a71..222a2e04117 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.test.tsx @@ -7,6 +7,7 @@ import ModelsAndEndpointsPage from "./page"; vi.mock("./panels/AllModelsPanel", () => ({ default: () =>
})); vi.mock("./panels/AddModelPanel", () => ({ default: () =>
})); +vi.mock("./panels/AutoRoutersTabPanel", () => ({ default: () =>
})); vi.mock("./panels/LlmCredentialsPanel", () => ({ default: () =>
})); vi.mock("./panels/PassThroughPanel", () => ({ default: () =>
})); vi.mock("./panels/HealthStatusPanel", () => ({ default: () =>
})); @@ -97,4 +98,34 @@ describe("ModelsAndEndpointsPage", () => { expect(queryByRole("tab", { name: "LLM Credentials" })).toBeNull(); expect(queryByRole("tab", { name: "Health Status" })).toBeNull(); }); + + // Auto-routers are excluded from the All Models table, so this tab is their home: the only + // place in the product to list, create, edit or delete one. + describe("Auto-Routers tab", () => { + it("sits third, after All Models and Add Model", () => { + const { getAllByRole } = renderPage(); + + const tabs = getAllByRole("tab").map((tab) => tab.textContent); + expect(tabs[0]).toContain("All Models"); + expect(tabs[1]).toBe("Add Model"); + expect(tabs[2]).toContain("Auto-Routers"); + // Badged Beta while the tab settles; BetaBadge renders the label text. + expect(tabs[2]).toContain("Beta"); + }); + + it("renders its panel when selected", async () => { + const user = userEvent.setup(); + const { getByRole, getByTestId } = renderPage(); + + await user.click(getByRole("tab", { name: /Auto-Routers/ })); + expect(getByTestId("panel-auto-routers")).toBeInTheDocument(); + }); + + it("is hidden from non-admins, who cannot write models", () => { + mockUseAuthorized.mockReturnValue(NON_ADMIN); + const { queryByRole } = renderPage(); + + expect(queryByRole("tab", { name: /Auto-Routers/ })).toBeNull(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx index cb173367459..22ffb6d6cc8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx @@ -7,13 +7,16 @@ import { useQueryClient } from "@tanstack/react-query"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; -import { all_admin_roles, internalUserRoles, isProxyAdminRole, isUserTeamAdminForAnyTeam } from "@/utils/roles"; +import { all_admin_roles, internalUserRoles } from "@/utils/roles"; +import { canCreateModels } from "@/utils/modelPermissions"; +import BetaBadge from "@/components/BetaBadge"; import CostOptimizationFeedbackBanner from "@/components/molecules/cost_optimization_feedback_banner"; import ModelInfoView from "@/components/model_info_view"; import TeamInfoView from "@/components/team/TeamInfo"; import { useModelDetailRouting } from "@/app/(dashboard)/models-and-endpoints/detailNavigation"; import { useModelDashboardData } from "@/app/(dashboard)/models-and-endpoints/useModelDashboardData"; import AllModelsPanel from "@/app/(dashboard)/models-and-endpoints/panels/AllModelsPanel"; +import AutoRoutersTabPanel from "@/app/(dashboard)/models-and-endpoints/panels/AutoRoutersTabPanel"; import AddModelPanel from "@/app/(dashboard)/models-and-endpoints/panels/AddModelPanel"; import LlmCredentialsPanel from "@/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel"; import PassThroughPanel from "@/app/(dashboard)/models-and-endpoints/panels/PassThroughPanel"; @@ -24,6 +27,7 @@ import PriceDataPanel from "@/app/(dashboard)/models-and-endpoints/panels/PriceD type ModelTabSlug = | "add" + | "auto-routers" | "llm-credentials" | "pass-through" | "health" @@ -35,6 +39,7 @@ const BASE_TAB_KEY = "all-models"; const TAB_LABELS: Record = { add: "Add Model", + "auto-routers": "Auto-Routers", "llm-credentials": "LLM Credentials", "pass-through": "Pass-Through Endpoints", health: "Health Status", @@ -47,6 +52,8 @@ const renderPanel = (key: string) => { switch (key) { case BASE_TAB_KEY: return ; + case "auto-routers": + return ; case "add": return ; case "llm-credentials": @@ -77,31 +84,48 @@ export default function ModelsAndEndpointsPage() { const [activeKey, setActiveKey] = useState(BASE_TAB_KEY); const [lastRefreshed, setLastRefreshed] = useState(""); - const isProxyAdmin = userRole && isProxyAdminRole(userRole); const isInternalUser = userRole && internalUserRoles.includes(userRole); - const isUserTeamAdmin = userID && isUserTeamAdminForAnyTeam(teams ?? null, userID); - const addModelDisabledForInternalUsers = - isInternalUser && uiSettings?.values?.disable_model_add_for_internal_users === true; - const shouldHideAddModelTab = !isProxyAdmin && (addModelDisabledForInternalUsers || !isUserTeamAdmin); + const canCreate = canCreateModels( + { userRole, userID }, + { + teams: teams ?? null, + disabledForInternalUsers: + isInternalUser === true && uiSettings?.values?.disable_model_add_for_internal_users === true, + }, + ); const isAdmin = all_admin_roles.includes(userRole); const visibleSlugs = useMemo>( () => [ "", - ...(shouldHideAddModelTab ? [] : (["add"] as const)), + ...(canCreate ? (["add"] as const) : []), + ...(isAdmin || canCreate ? (["auto-routers"] as const) : []), ...(isAdmin ? (["llm-credentials", "pass-through", "health", "retry-settings", "model-group-alias", "price-data"] as const) : []), ], - [shouldHideAddModelTab, isAdmin], + [canCreate, isAdmin], ); const allModelsLabel = isAdmin ? "All Models" : "Your Models"; + // Auto-Routers carries a Beta badge; BetaBadge honours the admin setting that hides these. + const tabLabel = (slug: "" | ModelTabSlug): React.ReactNode => { + if (!slug) return allModelsLabel; + if (slug === "auto-routers") { + return ( + + {TAB_LABELS[slug]} + + ); + } + return TAB_LABELS[slug]; + }; + const tabItems = visibleSlugs.map((slug) => { const key = slug || BASE_TAB_KEY; return { key, - label: slug ? TAB_LABELS[slug] : allModelsLabel, + label: tabLabel(slug), children: key === activeKey ? renderPanel(key) : null, }; }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx index 26dcc60d717..4d7bcc921af 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx @@ -3,7 +3,7 @@ import { Form } from "antd"; import { useState } from "react"; import { useQueryClient } from "@tanstack/react-query"; -import AddModelTab from "@/components/add_model/add_model_tab"; +import AddModelForm from "@/components/add_model/AddModelForm"; import { handleAddModelSubmit } from "@/components/add_model/handle_add_model_submit"; import { Providers, getPlaceholder, getProviderModels } from "@/components/provider_info_helpers"; import NotificationsManager from "@/components/molecules/notifications_manager"; @@ -14,7 +14,7 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { vertexCredentialsUploadProps } from "@/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload"; export default function AddModelPanel() { - const { accessToken, userRole } = useAuthorized(); + const { accessToken } = useAuthorized(); const [form] = Form.useForm(); const queryClient = useQueryClient(); const { data: modelCostMapData } = useModelCostMap(); @@ -39,7 +39,7 @@ export default function AddModelPanel() { }; return ( - ); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AutoRoutersTabPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AutoRoutersTabPanel.tsx new file mode 100644 index 00000000000..5f7d56e8e33 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AutoRoutersTabPanel.tsx @@ -0,0 +1,40 @@ +"use client"; + +import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { internalUserRoles } from "@/utils/roles"; +import { modelCreationScope } from "@/utils/modelPermissions"; + +import { AutoRoutersPanel } from "../components/AutoRouters/AutoRoutersPanel"; + +/** + * Owns the permission decision for the Auto-Routers tab so the panel stays a renderer. + * Creating an auto router is a POST /model/new, the same endpoint Add Model posts to, so it + * takes the same audience rule: a proxy admin, or a team admin who scopes it to a team. + * Viewer roles reach the list without write affordances. + */ +export default function AutoRoutersTabPanel() { + const { accessToken, userRole, userId: userID } = useAuthorized(); + const { data: teams } = useTeams(); + const { data: uiSettings } = useUISettings(); + + const isInternalUser = userRole != null && internalUserRoles.includes(userRole); + const scope = modelCreationScope( + { userRole, userID }, + { + teams: teams ?? null, + disabledForInternalUsers: isInternalUser && uiSettings?.values?.disable_model_add_for_internal_users === true, + }, + ); + + return ( + + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts b/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts index c49c73c3578..5acf444a359 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts @@ -1,7 +1,11 @@ -export function navigateWithParams(mutate: (params: URLSearchParams) => void): void { +export function navigateWithParams(mutate: (params: URLSearchParams) => void, mode: "push" | "replace" = "push"): void { const params = new URLSearchParams(window.location.search); mutate(params); const qs = params.toString(); const url = qs ? `${window.location.pathname}?${qs}` : window.location.pathname; - window.history.pushState(null, "", url); + if (mode === "replace") { + window.history.replaceState(null, "", url); + } else { + window.history.pushState(null, "", url); + } } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx index d381e5e65ca..3f9de478069 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx @@ -1,7 +1,9 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { render, screen } from "@testing-library/react"; +import { act, render, screen } from "@testing-library/react"; import React from "react"; -import { describe, expect, it, vi } from "vitest"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import type OrganizationsTableComponent from "./OrganizationsTable"; +import type OrganizationInfoViewComponent from "@/components/organization/organization_view"; vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({ __esModule: true, @@ -18,12 +20,50 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ userRole: null, }), })); +type OrganizationsTableProps = React.ComponentProps; +type OrganizationInfoViewProps = React.ComponentProps; + +let capturedTableProps: OrganizationsTableProps | null = null; vi.mock("./OrganizationsTable", () => ({ __esModule: true, - default: (props: { isLoading: boolean }) => ( -
isLoading:{String(props.isLoading)}
- ), + default: (props: OrganizationsTableProps) => { + capturedTableProps = props; + return
isLoading:{String(props.isLoading)}
; + }, })); +const mockOrgInfoView = vi.fn<(props: OrganizationInfoViewProps) => void>(); +vi.mock("@/components/organization/organization_view", () => ({ + __esModule: true, + default: (props: OrganizationInfoViewProps) => { + mockOrgInfoView(props); + return
; + }, +})); + +// The selected org is URL-derived (?org=) via useOrgDetailRouting. Next's real useSearchParams +// re-renders subscribers on history.pushState/replaceState; mirror that so URL changes propagate. +vi.mock("next/navigation", async () => { + const { useSyncExternalStore } = await import("react"); + const LOCATION_CHANGE_EVENT = "test-locationchange"; + for (const method of ["pushState", "replaceState"] as const) { + const original = window.history[method].bind(window.history); + window.history[method] = (...args: Parameters) => { + original(...args); + window.dispatchEvent(new Event(LOCATION_CHANGE_EVENT)); + }; + } + const subscribe = (onChange: () => void) => { + window.addEventListener(LOCATION_CHANGE_EVENT, onChange); + window.addEventListener("popstate", onChange); + return () => { + window.removeEventListener(LOCATION_CHANGE_EVENT, onChange); + window.removeEventListener("popstate", onChange); + }; + }; + return { + useSearchParams: () => new URLSearchParams(useSyncExternalStore(subscribe, () => window.location.search)), + }; +}); import OrganizationsPanel from "./OrganizationsPanel"; @@ -34,6 +74,12 @@ const renderWithQueryClient = (ui: React.ReactElement) => { return render({ui}); }; +beforeEach(() => { + capturedTableProps = null; + mockOrgInfoView.mockClear(); + window.history.replaceState(null, "", "/organizations/"); +}); + describe("OrganizationsPanel", () => { it("gates non-premium users behind the enterprise notice", () => { renderWithQueryClient(); @@ -55,3 +101,60 @@ describe("OrganizationsPanel", () => { expect(screen.getByTestId("organizations-table")).toHaveTextContent("isLoading:false"); }); }); + +describe("OrganizationsPanel - org detail deep link (?org=)", () => { + it("clicking an organization pushes ?org= and opens the detail view", () => { + renderWithQueryClient(); + + act(() => capturedTableProps?.onOrganizationClick("org-deep-link")); + + expect(window.location.search).toContain("org=org-deep-link"); + expect(mockOrgInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ organizationId: "org-deep-link" })); + }); + + it("opens the org detail directly from a ?org= deep link", () => { + window.history.replaceState(null, "", "/organizations/?org=org-from-url"); + renderWithQueryClient(); + + expect(mockOrgInfoView).toHaveBeenLastCalledWith( + expect.objectContaining({ organizationId: "org-from-url", editOrg: false }), + ); + expect(screen.queryByTestId("organizations-table")).not.toBeInTheDocument(); + }); + + it("closing the org detail removes ?org= and returns to the list", () => { + window.history.replaceState(null, "", "/organizations/?org=org-from-url"); + renderWithQueryClient(); + + act(() => mockOrgInfoView.mock.calls.at(-1)?.[0].onClose()); + + expect(window.location.search).not.toContain("org="); + expect(screen.queryByTestId("organization-info-view")).not.toBeInTheDocument(); + expect(screen.getByTestId("organizations-table")).toBeInTheDocument(); + }); + + it("the edit action opens the detail in edit mode with ?org= set", () => { + renderWithQueryClient(); + + act(() => capturedTableProps?.onEditClick("org-edit")); + + expect(window.location.search).toContain("org=org-edit"); + expect(mockOrgInfoView).toHaveBeenLastCalledWith( + expect.objectContaining({ organizationId: "org-edit", editOrg: true }), + ); + }); + + it("a plain row click after leaving an edit view via browser history does not reopen in edit mode", () => { + renderWithQueryClient(); + + act(() => capturedTableProps?.onEditClick("org-edit")); + expect(mockOrgInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ editOrg: true })); + + act(() => window.history.pushState(null, "", "/organizations/")); + act(() => capturedTableProps?.onOrganizationClick("org-plain")); + + expect(mockOrgInfoView).toHaveBeenLastCalledWith( + expect.objectContaining({ organizationId: "org-plain", editOrg: false }), + ); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx index b1c026d3904..a21c0669677 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx @@ -1,5 +1,6 @@ import { organizationKeys, useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; import { useUserModels } from "@/app/(dashboard)/hooks/models/useModels"; +import { useOrgDetailRouting } from "@/app/(dashboard)/organizations/detailNavigation"; import OrganizationFilters, { FilterState } from "@/app/(dashboard)/organizations/OrganizationFilters"; import { useQueryClient } from "@tanstack/react-query"; import React, { useState } from "react"; @@ -19,7 +20,7 @@ interface OrganizationsPanelProps { } const OrganizationsPanel: React.FC = ({ userRole, accessToken, premiumUser }) => { - const [selectedOrgId, setSelectedOrgId] = useState(null); + const { orgId: selectedOrgId, openOrg, close: closeOrgDetail } = useOrgDetailRouting(); const [editOrg, setEditOrg] = useState(false); const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); const [orgToDelete, setOrgToDelete] = useState(null); @@ -108,7 +109,7 @@ const OrganizationsPanel: React.FC = ({ userRole, acces { - setSelectedOrgId(null); + closeOrgDetail(); setEditOrg(false); }} accessToken={accessToken} @@ -132,9 +133,12 @@ const OrganizationsPanel: React.FC = ({ userRole, acces isLoading={isLoading} userRole={userRole} searchActive={searchActive} - onOrganizationClick={setSelectedOrgId} + onOrganizationClick={(organizationId) => { + setEditOrg(false); + openOrg(organizationId); + }} onEditClick={(organizationId) => { - setSelectedOrgId(organizationId); + openOrg(organizationId); setEditOrg(true); }} onDeleteClick={handleDelete} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.test.ts new file mode 100644 index 00000000000..46b7c4313ea --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.test.ts @@ -0,0 +1,53 @@ +/* @vitest-environment jsdom */ +import { act, renderHook } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { useOrgDetailRouting } from "./detailNavigation"; + +vi.mock("next/navigation", () => ({ useSearchParams: () => new URLSearchParams(window.location.search) })); + +describe("useOrgDetailRouting", () => { + beforeEach(() => { + window.history.pushState(null, "", "/organizations/"); + }); + + it("openOrg sets ?org= via history.pushState (no full navigation)", () => { + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useOrgDetailRouting()); + act(() => result.current.openOrg("org-abc123")); + expect(spy).toHaveBeenCalledWith(null, "", expect.stringContaining("org=org-abc123")); + spy.mockRestore(); + }); + + it("openOrg preserves unrelated query params", () => { + window.history.pushState(null, "", "/organizations/?foo=bar"); + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useOrgDetailRouting()); + act(() => result.current.openOrg("org-abc123")); + const url = spy.mock.calls.at(-1)?.[2] as string; + expect(url).toContain("foo=bar"); + expect(url).toContain("org=org-abc123"); + spy.mockRestore(); + }); + + it("close removes only the org param", () => { + window.history.pushState(null, "", "/organizations/?foo=bar&org=org-abc123"); + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useOrgDetailRouting()); + act(() => result.current.close()); + const url = spy.mock.calls.at(-1)?.[2] as string; + expect(url).toContain("foo=bar"); + expect(url).not.toContain("org="); + spy.mockRestore(); + }); + + it("exposes orgId from ?org=", () => { + window.history.pushState(null, "", "/organizations/?org=org-abc123"); + const { result } = renderHook(() => useOrgDetailRouting()); + expect(result.current.orgId).toBe("org-abc123"); + }); + + it("orgId is null when no org param is present", () => { + const { result } = renderHook(() => useOrgDetailRouting()); + expect(result.current.orgId).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.ts b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.ts new file mode 100644 index 00000000000..8c55c7b750c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.ts @@ -0,0 +1,32 @@ +import { useSearchParams } from "next/navigation"; +import { useCallback } from "react"; + +import { navigateWithParams } from "../navigateWithParams"; + +export interface OrgDetailRouting { + orgId: string | null; + openOrg: (id: string) => void; + close: () => void; +} + +export function useOrgDetailRouting(): OrgDetailRouting { + const searchParams = useSearchParams(); + + const openOrg = useCallback((id: string) => { + navigateWithParams((params) => { + params.set("org", id); + }); + }, []); + + const close = useCallback(() => { + navigateWithParams((params) => { + params.delete("org"); + }); + }, []); + + return { + orgId: searchParams?.get("org") ?? null, + openOrg, + close, + }; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.test.tsx index c4cfa98b2a1..9ffcbfc9975 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.test.tsx @@ -13,7 +13,6 @@ vi.mock("@/components/networking", () => ({ vi.mock("@/components/router_settings", () => ({ default: () => null })); vi.mock("@/components/Settings/RouterSettings/Fallbacks/Fallbacks", () => ({ default: () => null })); vi.mock("@/components/routing_groups", () => ({ default: () => null })); - // Mirrors the /config/list ordering: the two prompt-caching rows sit between the // General-tab rows in the unfiltered response but are filtered out of the General // tab's table, so any index-based lookup into the unfiltered array reads the wrong @@ -99,3 +98,19 @@ describe("GeneralSettings General tab", () => { expect(within(row).getByRole("spinbutton")).toHaveValue("1.00"); }); }); + +// The five tabs here are proxy-wide settings. Auto-routers moved to Models + Endpoints. +describe("GeneralSettings tabs", () => { + beforeEach(() => { + vi.mocked(getGeneralSettingsCall).mockResolvedValue([]); + }); + + it("renders the proxy-wide tabs and no auto-router tab", async () => { + renderWithProviders(); + + for (const name of ["Loadbalancing", "Routing Groups", "Fallbacks", "Prompt Caching", "General"]) { + expect(await screen.findByRole("tab", { name })).toBeInTheDocument(); + } + expect(screen.queryByRole("tab", { name: /auto.?router/i })).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx index fa3447e0cbf..ed7b17067d5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx @@ -13,7 +13,7 @@ import { Icon, Switch, } from "@tremor/react"; -import { TabPanel, TabPanels, TabGroup, TabList, Tab } from "@tremor/react"; +import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { getGeneralSettingsCall, updateConfigFieldSetting, deleteConfigFieldSetting } from "@/components/networking"; import { InputNumber, Select as AntdSelect } from "antd"; import { TrashIcon } from "@heroicons/react/outline"; @@ -232,82 +232,80 @@ const GeneralSettings: React.FC = ({ accessToken, user return (
- - - Loadbalancing - Routing Groups - Fallbacks - Prompt Caching - General - - - - - - - - - - - - - - - - - - - - Setting - Value - Status - Action - - - - {generalSettings - .filter((value) => value.field_type !== "TypedDictionary" && value.field_tab !== PROMPT_CACHING_TAB) - .map((value, index) => ( - - - {value.field_name} -

- {value.field_description} -

-
- - - - - {value.stored_in_db == true ? ( - - ) : value.stored_in_db == false ? ( - - ) : ( - - )} - - - - handleResetField(value.field_name)}> - Reset - - -
- ))} -
-
-
-
-
-
+ + + Loadbalancing + Routing Groups + Fallbacks + Prompt Caching + General + + + + + + + + + + + + + + + + + + + Setting + Value + Status + Action + + + + {generalSettings + .filter((value) => value.field_type !== "TypedDictionary" && value.field_tab !== PROMPT_CACHING_TAB) + .map((value, index) => ( + + + {value.field_name} +

+ {value.field_description} +

+
+ + + + + {value.stored_in_db == true ? ( + + ) : value.stored_in_db == false ? ( + + ) : ( + + )} + + + + handleResetField(value.field_name)}> + Reset + + +
+ ))} +
+
+
+
+
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.test.ts new file mode 100644 index 00000000000..e5d5b1a4073 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.test.ts @@ -0,0 +1,53 @@ +/* @vitest-environment jsdom */ +import { act, renderHook } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { useTeamDetailRouting } from "./detailNavigation"; + +vi.mock("next/navigation", () => ({ useSearchParams: () => new URLSearchParams(window.location.search) })); + +describe("useTeamDetailRouting", () => { + beforeEach(() => { + window.history.pushState(null, "", "/teams/"); + }); + + it("openTeam sets ?team= via history.pushState (no full navigation)", () => { + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useTeamDetailRouting()); + act(() => result.current.openTeam("team-abc123")); + expect(spy).toHaveBeenCalledWith(null, "", expect.stringContaining("team=team-abc123")); + spy.mockRestore(); + }); + + it("openTeam preserves unrelated query params", () => { + window.history.pushState(null, "", "/teams/?foo=bar"); + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useTeamDetailRouting()); + act(() => result.current.openTeam("team-abc123")); + const url = spy.mock.calls.at(-1)?.[2] as string; + expect(url).toContain("foo=bar"); + expect(url).toContain("team=team-abc123"); + spy.mockRestore(); + }); + + it("close removes only the team param", () => { + window.history.pushState(null, "", "/teams/?foo=bar&team=team-abc123"); + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useTeamDetailRouting()); + act(() => result.current.close()); + const url = spy.mock.calls.at(-1)?.[2] as string; + expect(url).toContain("foo=bar"); + expect(url).not.toContain("team="); + spy.mockRestore(); + }); + + it("exposes teamId from ?team=", () => { + window.history.pushState(null, "", "/teams/?team=team-abc123"); + const { result } = renderHook(() => useTeamDetailRouting()); + expect(result.current.teamId).toBe("team-abc123"); + }); + + it("teamId is null when no team param is present", () => { + const { result } = renderHook(() => useTeamDetailRouting()); + expect(result.current.teamId).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.ts b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.ts new file mode 100644 index 00000000000..d5208f094cb --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.ts @@ -0,0 +1,32 @@ +import { useSearchParams } from "next/navigation"; +import { useCallback } from "react"; + +import { navigateWithParams } from "../navigateWithParams"; + +export interface TeamDetailRouting { + teamId: string | null; + openTeam: (id: string) => void; + close: () => void; +} + +export function useTeamDetailRouting(): TeamDetailRouting { + const searchParams = useSearchParams(); + + const openTeam = useCallback((id: string) => { + navigateWithParams((params) => { + params.set("team", id); + }); + }, []); + + const close = useCallback(() => { + navigateWithParams((params) => { + params.delete("team"); + }); + }, []); + + return { + teamId: searchParams?.get("team") ?? null, + openTeam, + close, + }; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index 89c38c6274f..82ca66b10c0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx @@ -1,4 +1,4 @@ -import { act, fireEvent, render, screen, waitFor } from "@testing-library/react"; +import { act, fireEvent, render, screen, waitFor, within } from "@testing-library/react"; import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest"; import * as networking from "@/components/networking"; import EntityUsage from "./EntityUsage"; @@ -497,7 +497,7 @@ describe("EntityUsage", () => { it.each([ ["Cost", "Tag Spend Overview"], - ["Model Activity", "metrics-source:models"], + ["Model Activity", "metrics-source:model_groups"], ["Key Activity", "metrics-source:api_keys"], ["Endpoint Activity", "Endpoint Usage Panel"], ])("shows only the %s panel for a non-team entity type", async (tabLabel, marker) => { @@ -518,7 +518,7 @@ describe("EntityUsage", () => { it.each([ ["Cost", "Team Spend Overview"], - ["Model Activity", "metrics-source:models"], + ["Model Activity", "metrics-source:model_groups"], ["Agent Activity", "metrics-source:entities"], ["Key Activity", "metrics-source:api_keys"], ["Endpoint Activity", "Endpoint Usage Panel"], @@ -584,15 +584,41 @@ describe("EntityUsage", () => { expect(screen.getByText("Request / Token Consumption")).toBeInTheDocument(); }); - it("should display Top Models title for non-agent entity types", async () => { + it("should display Top Public Model Names title for non-agent entity types", async () => { render(); await waitFor(() => { expect(mockTagDailyActivityCall).toHaveBeenCalled(); }); - const topModelsElements = screen.getAllByText("Top Models"); - expect(topModelsElements.length).toBeGreaterThan(0); + expect(screen.getByText("Top Public Model Names")).toBeInTheDocument(); + }); + + it("defaults Model Activity to public model names and toggles to litellm models", async () => { + const { container } = render(); + + await waitFor(() => { + expect(mockTagDailyActivityCall).toHaveBeenCalled(); + }); + + act(() => { + fireEvent.click(screen.getByText("Model Activity")); + }); + + const modelActivityPanel = () => selectedPanels(container)[0] as HTMLElement; + expect(modelActivityPanel().textContent).toContain("metrics-source:model_groups"); + + act(() => { + fireEvent.click(within(modelActivityPanel()).getByText("Litellm Model Name")); + }); + + expect(modelActivityPanel().textContent).toContain("metrics-source:models"); + + act(() => { + fireEvent.click(within(modelActivityPanel()).getByText("Public Model Name")); + }); + + expect(modelActivityPanel().textContent).toContain("metrics-source:model_groups"); }); it("should display Top Agents title for agent entity type", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index e330983b6f9..4d44791d1a9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -49,6 +49,7 @@ import { } from "@/components/UsagePage/types"; import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters"; import EndpointUsage from "../EndpointUsage/EndpointUsage"; +import ModelViewToggle, { ModelViewType } from "../ModelViewToggle"; import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView"; import TopModelView from "./TopModelView"; @@ -110,6 +111,7 @@ const ENTITY_FETCH_FNS: Record Promise> = { const EntityUsage: React.FC = ({ accessToken, entityType, entityId, entityList, dateValue }) => { const { teams } = useTeams(); const [selectedTags, setSelectedTags] = useState([]); + const [modelViewType, setModelViewType] = useState("groups"); const [topKeysLimit, setTopKeysLimit] = useState(5); const [topModelsLimit, setTopModelsLimit] = useState(5); const [topAgentsLimit, setTopAgentsLimit] = useState(5); @@ -153,14 +155,15 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti const agentSpendData = agentSpendDataRaw as unknown as EntitySpendData; - const modelMetrics = processActivityData(spendData, "models", teams || []); + const modelBreakdownKey = modelViewType === "groups" ? "model_groups" : "models"; + const modelMetrics = processActivityData(spendData, modelBreakdownKey, teams || []); const keyMetrics = processActivityData(spendData, "api_keys", teams || []); const agentMetrics = entityType === "team" ? processActivityData(agentSpendData, "entities", teams || []) : {}; const getTopModels = () => { const modelSpend: { [key: string]: any } = {}; spendData.results.forEach((day) => { - Object.entries(day.breakdown.models || {}).forEach(([model, metrics]) => { + Object.entries(day.breakdown[modelBreakdownKey] || {}).forEach(([model, metrics]) => { if (!modelSpend[model]) { modelSpend[model] = { spend: 0, @@ -406,6 +409,8 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti const capitalizedEntityLabel = entityType.charAt(0).toUpperCase() + entityType.slice(1); + const modelViewTitle = modelViewType === "groups" ? "Top Public Model Names" : "Top Litellm Models"; + const costPanel = ( {/* Total Spend Card */} @@ -604,7 +609,10 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti {/* Top Models */} - {entityType === "agent" ? "Top Agents" : "Top Models"} +
+ {entityType === "agent" ? "Top Agents" : modelViewTitle} + +
= ({ accessToken, entityType, enti { key: "models", label: entityType === "agent" ? "Request / Token Consumption" : "Model Activity", - content: , + content: ( + <> +
+ +
+ + + ), }, ...(entityType === "team" ? [{ key: "agents", label: "Agent Activity", content: }] diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/ModelViewToggle.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/ModelViewToggle.tsx new file mode 100644 index 00000000000..0ee6dd3b19c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/ModelViewToggle.tsx @@ -0,0 +1,29 @@ +export type ModelViewType = "groups" | "individual"; + +const MODEL_VIEW_OPTIONS: readonly { value: ModelViewType; label: string }[] = [ + { value: "groups", label: "Public Model Name" }, + { value: "individual", label: "Litellm Model Name" }, +]; + +interface ModelViewToggleProps { + value: ModelViewType; + onChange: (value: ModelViewType) => void; +} + +export default function ModelViewToggle({ value, onChange }: ModelViewToggleProps) { + return ( +
+ {MODEL_VIEW_OPTIONS.map((option) => ( + + ))} +
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx index cf122137f91..98dae51fa37 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx @@ -30,8 +30,10 @@ vi.mock("@/components/networking", () => ({ // Mock child components to simplify testing vi.mock("@/components/activity_metrics", () => ({ - ActivityMetrics: () =>
Activity Metrics
, - processActivityData: () => ({ data: [], metadata: {} }), + ActivityMetrics: ({ modelMetrics }: { modelMetrics?: { __source?: string } }) => ( +
{`activity-source:${modelMetrics?.__source ?? "none"}`}
+ ), + processActivityData: (_data: unknown, key: string) => ({ __source: key }), })); vi.mock("@/components/view_user_spend", () => ({ @@ -1043,8 +1045,8 @@ describe("UsagePage", () => { // Default should be "groups" view showing "Top Public Model Names" expect(screen.getByText("Top Public Model Names")).toBeInTheDocument(); - expect(screen.getByText("Public Model Name")).toBeInTheDocument(); - expect(screen.getByText("Litellm Model Name")).toBeInTheDocument(); + expect(screen.getAllByText("Public Model Name").length).toBeGreaterThan(0); + expect(screen.getAllByText("Litellm Model Name").length).toBeGreaterThan(0); }); it("should switch to Litellm Model Name view on toggle click", async () => { @@ -1055,7 +1057,7 @@ describe("UsagePage", () => { }); // Click the "Litellm Model Name" toggle - const litellmToggle = screen.getByText("Litellm Model Name"); + const litellmToggle = screen.getAllByText("Litellm Model Name")[0]; act(() => { fireEvent.click(litellmToggle); }); @@ -1074,7 +1076,7 @@ describe("UsagePage", () => { }); // Switch to individual first - const litellmToggle = screen.getByText("Litellm Model Name"); + const litellmToggle = screen.getAllByText("Litellm Model Name")[0]; act(() => { fireEvent.click(litellmToggle); }); @@ -1084,7 +1086,7 @@ describe("UsagePage", () => { }); // Switch back to groups - const publicToggle = screen.getByText("Public Model Name"); + const publicToggle = screen.getAllByText("Public Model Name")[0]; act(() => { fireEvent.click(publicToggle); }); @@ -1093,6 +1095,34 @@ describe("UsagePage", () => { expect(screen.getByText("Top Public Model Names")).toBeInTheDocument(); }); }); + + it("should feed the Model Activity tab from the model_groups breakdown by default", async () => { + renderWithProviders(); + + await waitFor(() => { + expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled(); + }); + + expect(screen.getByText("activity-source:model_groups")).toBeInTheDocument(); + expect(screen.queryByText("activity-source:models")).not.toBeInTheDocument(); + }); + + it("should switch the Model Activity tab to the litellm models breakdown on toggle click", async () => { + renderWithProviders(); + + await waitFor(() => { + expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled(); + }); + + act(() => { + fireEvent.click(screen.getAllByText("Litellm Model Name")[0]); + }); + + await waitFor(() => { + expect(screen.getByText("activity-source:models")).toBeInTheDocument(); + }); + expect(screen.queryByText("activity-source:model_groups")).not.toBeInTheDocument(); + }); }); describe("customer usage banner", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx index d2f75609d18..46a17017d39 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx @@ -55,6 +55,7 @@ import { DailyData, KeyMetricWithMetadata, MetricWithMetadata } from "@/componen import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters"; import EndpointUsage from "./EndpointUsage/EndpointUsage"; import EntityUsage, { EntityList } from "./EntityUsage/EntityUsage"; +import ModelViewToggle, { ModelViewType } from "./ModelViewToggle"; import SpendByProvider from "./EntityUsage/SpendByProvider"; import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView"; import UsageAIChatPanel from "./UsageAIChatPanel"; @@ -143,7 +144,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { // For admins: null means global view (all users), a string means filter by that user // For non-admins: always set to their own user ID const [selectedUserId, setSelectedUserId] = useState(isAdmin ? null : userID || null); - const [modelViewType, setModelViewType] = useState<"groups" | "individual">("groups"); + const [modelViewType, setModelViewType] = useState("groups"); const [isCloudZeroModalOpen, setIsCloudZeroModalOpen] = useState(false); const [isGlobalExportModalOpen, setIsGlobalExportModalOpen] = useState(false); const [isAiChatOpen, setIsAiChatOpen] = useState(false); @@ -438,7 +439,10 @@ const UsagePage: React.FC = ({ teams, organizations }) => { () => [...userSpendData.results].sort((a, b) => new Date(a.date).getTime() - new Date(b.date).getTime()), [userSpendData.results], ); - const modelMetrics = useMemo(() => processActivityData(userSpendData, "models", teams), [userSpendData, teams]); + const modelMetrics = useMemo( + () => processActivityData(userSpendData, modelViewType === "groups" ? "model_groups" : "models", teams), + [userSpendData, modelViewType, teams], + ); const keyMetrics = useMemo(() => processActivityData(userSpendData, "api_keys", teams), [userSpendData, teams]); const mcpServerMetrics = useMemo( () => processActivityData(userSpendData, "mcp_servers", teams), @@ -753,28 +757,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { value={topModelsLimit} onChange={(value) => setTopModelsLimit(value as number)} /> -
- - -
+
{loading ? ( @@ -839,6 +822,9 @@ const UsagePage: React.FC = ({ teams, organizations }) => { {/* Activity Panel */} +
+ +
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/user_edit_view.tsx b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/user_edit_view.tsx index 54392adf885..c83a8e48e43 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/user_edit_view.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/user_edit_view.tsx @@ -1,11 +1,14 @@ import { InfoCircleOutlined } from "@ant-design/icons"; import { Button, SelectItem, TextInput, Textarea } from "@tremor/react"; -import { Checkbox, Form, Select, Tooltip } from "antd"; +import { Checkbox, Form, Input, Select, Tooltip } from "antd"; import React, { useState } from "react"; import { all_admin_roles } from "@/utils/roles"; import BudgetDurationDropdown from "@/components/common_components/budget_duration_dropdown"; import { getModelDisplayName } from "@/components/key_team_helpers/fetch_available_models_team_key"; import NumericalInput from "@/components/shared/numerical_input"; +import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector"; +import MCPToolPermissions from "@/components/mcp_server_management/MCPToolPermissions"; +import type { ObjectPermission } from "@/components/object_permission_types"; interface UserEditViewProps { userData: any; @@ -18,8 +21,18 @@ interface UserEditViewProps { userModels: string[]; possibleUIRoles: Record> | null; isBulkEdit?: boolean; + objectPermission?: ObjectPermission | null; } +const buildMcpFieldValues = (objectPermission: ObjectPermission | null | undefined) => ({ + mcp_servers_and_groups: { + servers: objectPermission?.mcp_servers ?? [], + accessGroups: objectPermission?.mcp_access_groups ?? [], + toolsets: objectPermission?.mcp_toolsets ?? [], + }, + mcp_tool_permissions: objectPermission?.mcp_tool_permissions ?? {}, +}); + export function UserEditView({ userData, onCancel, @@ -31,9 +44,11 @@ export function UserEditView({ userModels, possibleUIRoles, isBulkEdit = false, + objectPermission, }: UserEditViewProps) { const [form] = Form.useForm(); const [unlimitedBudget, setUnlimitedBudget] = useState(false); + const canEditMcpPermissions = !isBulkEdit && all_admin_roles.includes(userRole || ""); // Set initial form values React.useEffect(() => { @@ -50,8 +65,9 @@ export function UserEditView({ max_budget: isUnlimited ? "" : maxBudget, budget_duration: userData.user_info?.budget_duration, metadata: userData.user_info?.metadata ? JSON.stringify(userData.user_info.metadata, null, 2) : undefined, + ...(canEditMcpPermissions ? buildMcpFieldValues(objectPermission) : {}), }); - }, [userData, form]); + }, [userData, objectPermission, canEditMcpPermissions, form]); const handleUnlimitedBudgetChange = (e: any) => { const checked = e.target.checked; @@ -186,6 +202,52 @@ export function UserEditView({