diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index f212dd9d15e..d06b9a16e6d 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -14,11 +14,15 @@ permissions: jobs: lint: runs-on: ubuntu-latest - timeout-minutes: 5 + timeout-minutes: 10 steps: - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + # Check out the PR head, not the default refs/pull/N/merge: the merge ref + # folds in newer base commits, which the diff-based gates (ruff delta, + # Any-discipline) would otherwise blame on this branch. with: + ref: ${{ github.event.pull_request.head.sha }} fetch-depth: 0 clean: true persist-credentials: false @@ -73,6 +77,12 @@ jobs: run: | uv run --no-sync python scripts/ruff_strict_gate.py --base "$BASE_SHA" + - name: Check type-discipline budget (mutable collections / casts / type guards / kwargs / unexplained suppressions, delta vs base) + env: + BASE_SHA: ${{ github.event.pull_request.base.sha }} + run: | + uv run --no-sync python scripts/type_discipline_gate.py --base "$BASE_SHA" + - name: Print OpenAI version run: | uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')" @@ -80,8 +90,11 @@ jobs: - name: Run MyPy type checking run: | cd litellm - uv run --no-sync mypy . - cd .. + (uv run --no-sync mypy . || true) | uv run --no-sync python ../scripts/type_check_gate.py --tool mypy + + - name: Run basedpyright type checking + run: | + (uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py --tool basedpyright - name: Check for circular imports run: | @@ -93,6 +106,83 @@ jobs: run: | uv run --no-sync python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1) + # Intentionally NON-GATING. This job turns red when a *-budget.json ceiling is + # raised (or a rule/budget is dropped) so a loosening is obvious in review, but it + # must be kept OUT of the branch-protection required-checks list so a justified + # bump can still be merged by a human who has seen and accepted the red. + budget-ratchet: + runs-on: ubuntu-latest + timeout-minutes: 5 + permissions: + contents: read + + steps: + - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + fetch-depth: 0 + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Ratchet check (budgets may only decrease; non-gating) + env: + BASE_SHA: ${{ github.event.pull_request.base.sha }} + run: | + python scripts/budget_ratchet_check.py --base "$BASE_SHA" + + any-discipline: + # Separate job: the first run cold-builds litellm's type cache (~2 min, ~3 GB), + # so keep it off the main lint job's time budget. Subsequent runs reuse the + # cached .mypy_cache_any and only re-type-check the changed files. + runs-on: ubuntu-latest + timeout-minutes: 10 + + steps: + - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + # Check out the PR head, not the default refs/pull/N/merge: the merge ref + # folds in newer base commits, which the diff-based gates (ruff delta, + # Any-discipline) would otherwise blame on this branch. + with: + ref: ${{ github.event.pull_request.head.sha }} + fetch-depth: 0 + clean: true + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Set up uv + uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 + with: + version: "0.10.9" + + - name: Install dependencies + run: | + uv sync --frozen + + # Keyed on deps + mypy config (which fix the type cache's validity), not on + # source content, so changed files always differ from the restored cache. + # The gate also defensively invalidates each target's cache entry, so + # correctness never depends on cache freshness -- this is purely for speed. + - name: Restore Any-gate type cache + uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 + with: + path: .mypy_cache_any + key: any-mypy-cache-${{ runner.os }}-py3.12-${{ hashFiles('uv.lock', 'litellm/mypy.ini') }} + restore-keys: | + any-mypy-cache-${{ runner.os }}-py3.12- + + - name: Check Any discipline on changed lines + env: + BASE_SHA: ${{ github.event.pull_request.base.sha }} + run: | + uv run --no-sync python scripts/check_any_discipline.py --changed --base "$BASE_SHA" + secret-scan: runs-on: ubuntu-latest timeout-minutes: 5 diff --git a/.gitignore b/.gitignore index a303591c635..1be40c0f863 100644 --- a/.gitignore +++ b/.gitignore @@ -77,6 +77,7 @@ tests/local_testing/log.txt litellm/proxy/_new_new_secret_config.yaml litellm/proxy/custom_guardrail.py **/.mypy_cache/ +**/.mypy_cache_any/ litellm/proxy/application.log tests/llm_translation/vertex_test_account.json tests/llm_translation/test_vertex_key.json diff --git a/CLAUDE.md b/CLAUDE.md index 32fd0aadddb..a81ee1f3b91 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -36,7 +36,13 @@ Don't hesitate to use values in .env to get needed API keys and other secrets, a Run tests, format your code, and lint your code before each commit -When you fix strict-rule violations gated by `ruff-strict-budget.json`, run `make lint-strict-budget-update` and commit the lowered baselines so the ceilings ratchet down instead of leaving stale headroom +When you fix violations gated by `ruff-strict-budget.json`, `mypy-code-budget.json`, or `basedpyright-code-budget.json`, run `make lint-budget-update` and commit the lowered baselines so the ceilings ratchet down instead of leaving stale headroom + +If you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and bringing it closer to the max, just validate it in the caller with Pydantic (a model or `TypeAdapter` that returns the typed thing or raises will do) and then pass the now typed variable in + +The Any-discipline gate (`make lint-any`, also a CI job) fails when a line you changed under `litellm/` holds a value typed `Any`, including the `X | Any`. Ideally `# any-ok: ` is never used; treat it as a last resort for a genuine typed/untyped boundary that Pydantic truly can't model + +If you get an LIT001 or LIT002 fail, refactor the code to follow functional programming best practices rather than introducing mutable data structures. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `frozenset()` instead of seeding an empty `list`/`dict`/`set` and mutating it over time. Ideally `# mutable-ok` is never used; reach for it only as a genuine last resort when an immutable rewrite is truly impossible, and always pair it with a real reason Ask to commit and push your work when you're done (or if you're confident that your code is good and works, just do it) @@ -67,8 +73,6 @@ Follow these coding conventions for new/updated code (a three-line fix in a lega - No file sprawl: deliberate file and folder structure - Standard over hand-rolled: use the official SDK or a library where one exists; where none does, follow industry standards instead of inventing local conventions -if you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and bringing it closer to the max, just validate it in the caller (a simple function that returns the typed thing or raises will do) and then pass the now typed variable in - Follow conventional commits for commit names and PR titles ## Think Before Coding diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 2177c764806..97a8d53f831 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -155,6 +155,7 @@ Individual linting commands: make format-check # Check Black formatting make lint-ruff # Run Ruff linting make lint-mypy # Run MyPy type checking +make lint-any # Fail on Any-typed values on changed lines make check-circular-imports # Check for circular imports make check-import-safety # Check import safety ``` diff --git a/Makefile b/Makefile index e89932ac896..25fe88b9057 100644 --- a/Makefile +++ b/Makefile @@ -5,7 +5,8 @@ test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \ test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \ info lint lint-dev format \ - lint-strict-budget lint-strict-budget-update \ + lint-mypy lint-mypy-budget-update lint-basedpyright lint-basedpyright-budget-update \ + lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-any \ install-dev install-proxy-dev install-test-deps install-hooks \ install-helm-unittest check-circular-imports check-import-safety @@ -23,10 +24,15 @@ help: @echo " make format-check - Check Black code formatting (matches CI)" @echo " make lint - Run all linting (Ruff, MyPy, Black check, circular imports, import safety)" @echo " make lint-ruff - Run Ruff linting only" - @echo " make lint-mypy - Run MyPy type checking only" + @echo " make lint-mypy - Run MyPy (disallow_untyped_defs), gated by per-rule error counts" + @echo " make lint-mypy-budget-update - Re-capture the MyPy per-rule budget (ratchet)" + @echo " make lint-basedpyright - Run basedpyright strict, gated by per-rule error counts" + @echo " make lint-basedpyright-budget-update - Re-capture the basedpyright per-rule budget (ratchet)" @echo " make lint-black - Check Black formatting (matches CI)" - @echo " make lint-strict-budget - Gate the codebase total of each strict ruff rule against its ceiling" - @echo " make lint-strict-budget-update - Re-capture per-rule baselines in ruff-strict-budget.json (ratchet)" + @echo " make lint-ruff-budget - Gate the codebase total of each strict ruff rule against its ceiling" + @echo " make lint-ruff-budget-update - Re-capture per-rule baselines in ruff-strict-budget.json (ratchet)" + @echo " make lint-budget-update - Re-capture all three ratchet budgets (ruff + mypy + basedpyright)" + @echo " make lint-any - Fail if changed lines under litellm/ hold an Any-typed value" @echo " make check-circular-imports - Check for circular imports" @echo " make check-import-safety - Check import safety" @echo " make test - Run all tests" @@ -121,7 +127,16 @@ lint-ruff-FULL-dev: install-dev else echo "No changed .py files to check."; fi lint-mypy: install-dev - cd litellm && $(UV_RUN) mypy . --ignore-missing-imports && cd .. + cd litellm && ($(UV_RUN) mypy . || true) | $(UV_RUN) python ../scripts/type_check_gate.py --tool mypy + +lint-mypy-budget-update: install-dev + cd litellm && ($(UV_RUN) mypy . || true) | $(UV_RUN) python ../scripts/type_check_gate.py --tool mypy --update + +lint-basedpyright: install-dev + ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --tool basedpyright + +lint-basedpyright-budget-update: install-dev + ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --tool basedpyright --update lint-black: format-check @@ -129,9 +144,18 @@ lint-strict-budget: install-dev $(UV_RUN) python scripts/ruff_strict_gate.py \ --base $$(git rev-parse --verify origin/litellm_internal_staging 2>/dev/null && echo origin/litellm_internal_staging || echo upstream/litellm_internal_staging) -lint-strict-budget-update: install-dev +lint-ruff-budget: install-dev + $(UV_RUN) python scripts/ruff_strict_gate.py + +lint-ruff-budget-update: install-dev $(UV_RUN) python scripts/ruff_strict_gate.py --update +# Ratchet all three budgets in one shot (ruff strict + mypy + basedpyright) +lint-budget-update: lint-ruff-budget-update lint-mypy-budget-update lint-basedpyright-budget-update + +lint-any: install-dev + $(UV_RUN) python scripts/check_any_discipline.py --changed + check-circular-imports: install-dev cd litellm && $(UV_RUN) python ../tests/documentation_tests/test_circular_imports.py && cd .. @@ -139,10 +163,10 @@ check-import-safety: install-dev @$(UV_RUN) python -c "from litellm import *; print('[from litellm import *] OK! no issues!');" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1) # Combined linting (matches test-linting.yml workflow) -lint: format-check lint-ruff lint-mypy check-circular-imports check-import-safety lint-strict-budget +lint: format-check lint-ruff lint-mypy lint-basedpyright check-circular-imports check-import-safety lint-ruff-budget lint-any # Faster linting for local development (only checks changed code) -lint-dev: lint-format-changed lint-mypy check-circular-imports check-import-safety +lint-dev: lint-format-changed lint-mypy lint-any check-circular-imports check-import-safety # Testing targets test: install-test-deps diff --git a/README.md b/README.md index d600f3952c6..d7dc665dcec 100644 --- a/README.md +++ b/README.md @@ -327,6 +327,7 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \ | [Maritalk (`maritalk`)](https://docs.litellm.ai/docs/providers/maritalk) | ✅ | ✅ | ✅ | | | | | | | | | [Meta - Llama API (`meta_llama`)](https://docs.litellm.ai/docs/providers/meta_llama) | ✅ | ✅ | ✅ | | | | | | | | | [Mistral AI API (`mistral`)](https://docs.litellm.ai/docs/providers/mistral) | ✅ | ✅ | ✅ | ✅ | | | | | | | +| [ModelScope (`modelscope`)](https://docs.litellm.ai/docs/providers/modelscope) | ✅ | ✅ | ✅ | | ✅ | | | | | | | [Moonshot (`moonshot`)](https://docs.litellm.ai/docs/providers/moonshot) | ✅ | ✅ | ✅ | | | | | | | | | [Morph (`morph`)](https://docs.litellm.ai/docs/providers/morph) | ✅ | ✅ | ✅ | | | | | | | | | [Nebius AI Studio (`nebius`)](https://docs.litellm.ai/docs/providers/nebius) | ✅ | ✅ | ✅ | ✅ | | | | | | | diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json new file mode 100644 index 00000000000..d8ea65d47c6 --- /dev/null +++ b/basedpyright-code-budget.json @@ -0,0 +1,194 @@ +{ + "reportAny": { + "baseline": 24954, + "slack": 2500 + }, + "reportArgumentType": { + "baseline": 1863, + "slack": 3 + }, + "reportAssignmentType": { + "baseline": 220, + "slack": 3 + }, + "reportAttributeAccessIssue": { + "baseline": 335, + "slack": 3 + }, + "reportCallIssue": { + "baseline": 77, + "slack": 10 + }, + "reportConstantRedefinition": { + "baseline": 39, + "slack": 3 + }, + "reportDeprecated": { + "baseline": 217, + "slack": 10 + }, + "reportDuplicateImport": { + "baseline": 28, + "slack": 3 + }, + "reportExplicitAny": { + "baseline": 6931, + "slack": 700 + }, + "reportFunctionMemberAccess": { + "baseline": 7, + "slack": 3 + }, + "reportGeneralTypeIssues": { + "baseline": 151, + "slack": 3 + }, + "reportIncompatibleMethodOverride": { + "baseline": 52, + "slack": 10 + }, + "reportIncompatibleVariableOverride": { + "baseline": 8, + "slack": 3 + }, + "reportInconsistentOverload": { + "baseline": 12, + "slack": 3 + }, + "reportIndexIssue": { + "baseline": 26, + "slack": 3 + }, + "reportInvalidTypeForm": { + "baseline": 23, + "slack": 3 + }, + "reportInvalidTypeVarUse": { + "baseline": 2, + "slack": 3 + }, + "reportMatchNotExhaustive": { + "baseline": 1, + "slack": 3 + }, + "reportMissingParameterType": { + "baseline": 3933, + "slack": 10 + }, + "reportMissingTypeArgument": { + "baseline": 10612, + "slack": 1000 + }, + "reportMissingTypeStubs": { + "baseline": 27, + "slack": 10 + }, + "reportOperatorIssue": { + "baseline": 6, + "slack": 3 + }, + "reportOptionalCall": { + "baseline": 4, + "slack": 3 + }, + "reportOptionalIterable": { + "baseline": 3, + "slack": 3 + }, + "reportOptionalMemberAccess": { + "baseline": 724, + "slack": 10 + }, + "reportOptionalOperand": { + "baseline": 3, + "slack": 3 + }, + "reportOptionalSubscript": { + "baseline": 11, + "slack": 3 + }, + "reportPossiblyUnboundVariable": { + "baseline": 52, + "slack": 10 + }, + "reportPrivateUsage": { + "baseline": 1625, + "slack": 10 + }, + "reportRedeclaration": { + "baseline": 8, + "slack": 3 + }, + "reportReturnType": { + "baseline": 118, + "slack": 10 + }, + "reportTypedDictNotRequiredAccess": { + "baseline": 20, + "slack": 3 + }, + "reportUndefinedVariable": { + "baseline": 2, + "slack": 3 + }, + "reportUnknownArgumentType": { + "baseline": 30603, + "slack": 3000 + }, + "reportUnknownLambdaType": { + "baseline": 76, + "slack": 10 + }, + "reportUnknownMemberType": { + "baseline": 27322, + "slack": 2500 + }, + "reportUnknownParameterType": { + "baseline": 13636, + "slack": 1000 + }, + "reportUnknownVariableType": { + "baseline": 21776, + "slack": 2000 + }, + "reportUnnecessaryCast": { + "baseline": 118, + "slack": 10 + }, + "reportUnnecessaryComparison": { + "baseline": 680, + "slack": 10 + }, + "reportUnnecessaryContains": { + "baseline": 4, + "slack": 3 + }, + "reportUnnecessaryIsInstance": { + "baseline": 807, + "slack": 10 + }, + "reportUntypedBaseClass": { + "baseline": 110, + "slack": 3 + }, + "reportUntypedFunctionDecorator": { + "baseline": 22, + "slack": 3 + }, + "reportUnusedClass": { + "baseline": 22, + "slack": 3 + }, + "reportUnusedFunction": { + "baseline": 137, + "slack": 10 + }, + "reportUnusedImport": { + "baseline": 670, + "slack": 10 + }, + "reportUnusedVariable": { + "baseline": 865, + "slack": 10 + } +} diff --git a/codecov.yaml b/codecov.yaml index 58681b884d0..3baea13e2d3 100644 --- a/codecov.yaml +++ b/codecov.yaml @@ -35,6 +35,22 @@ component_management: - component_id: "Enterprise" paths: - "enterprise/**" + - component_id: "Batches" + paths: + - "*/proxy/batches_endpoints/**" + - "litellm/batches/**" + - "*/llms/*/batches/**" + - component_id: "Videos" + paths: + - "litellm/videos/**" + - "*/proxy/video_endpoints/**" + - "*/llms/*/videos/**" + - component_id: "Realtime" + paths: + - "litellm/realtime_api/**" + - "*/proxy/realtime_endpoints/**" + - "*/llms/*/realtime/**" + - "litellm/litellm_core_utils/realtime_streaming.py" comment: layout: "header, diff, flags, components" # show component info in the PR comment diff --git a/db_scripts/create_views.py b/db_scripts/create_views.py index 3027b38958d..2b34664452d 100644 --- a/db_scripts/create_views.py +++ b/db_scripts/create_views.py @@ -15,7 +15,7 @@ db = Prisma( ) -async def check_view_exists(): # noqa: PLR0915 +async def check_view_exists(): """ Checks if the LiteLLM_VerificationTokenView and MonthlyGlobalSpend exists in the user's db. @@ -34,8 +34,7 @@ async def check_view_exists(): # noqa: PLR0915 print("LiteLLM_VerificationTokenView Exists!") # noqa except Exception: # If an error occurs, the view does not exist, so create it - await db.execute_raw( - """ + await db.execute_raw(""" CREATE VIEW "LiteLLM_VerificationTokenView" AS SELECT v.*, @@ -45,8 +44,7 @@ async def check_view_exists(): # noqa: PLR0915 t.rpm_limit AS team_rpm_limit FROM "LiteLLM_VerificationToken" v LEFT JOIN "LiteLLM_TeamTable" t ON v.team_id = t.team_id; - """ - ) + """) print("LiteLLM_VerificationTokenView Created!") # noqa diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index a1f63f388b4..6830147116d 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -412,7 +412,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): detail=f"User {user_api_key_dict.user_id} does not have access to the file {file_id}", ) - async def async_pre_call_hook( # noqa: PLR0915 + async def async_pre_call_hook( self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 75229bacc8f..a057df65500 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -483,7 +483,7 @@ async def new_project( response_model=LiteLLM_ProjectTable, ) @management_endpoint_wrapper -async def update_project( # noqa: PLR0915 +async def update_project( data: UpdateProjectRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), diff --git a/litellm/__init__.py b/litellm/__init__.py index d5fbb41c462..0d6a788e368 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -73,6 +73,7 @@ from litellm.constants import ( replicate_models, clarifai_models, huggingface_models, + modelscope_models, empower_models, together_ai_models, baseten_models, @@ -900,6 +901,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None): heroku_models.add(key) elif value.get("litellm_provider") == "dashscope": dashscope_models.add(key) + elif value.get("litellm_provider") == "modelscope": + modelscope_models.add(key) elif value.get("litellm_provider") == "moonshot": moonshot_models.add(key) elif value.get("litellm_provider") == "publicai": @@ -1019,6 +1022,7 @@ model_list = list( | zai_models | fal_ai_models | deepseek_models + | modelscope_models | azure_ai_models | voyage_models | infinity_models @@ -1152,6 +1156,7 @@ models_by_provider: dict = { "elevenlabs": elevenlabs_models, "heroku": heroku_models, "dashscope": dashscope_models, + "modelscope": modelscope_models, "moonshot": moonshot_models, "publicai": publicai_models, "v0": v0_models, @@ -1975,6 +1980,9 @@ if TYPE_CHECKING: from .llms.dashscope.rerank.transformation import ( DashScopeRerankConfig as DashScopeRerankConfig, ) + from .llms.modelscope.chat.transformation import ( + ModelScopeChatConfig as ModelScopeChatConfig, + ) from .llms.moonshot.chat.transformation import ( MoonshotChatConfig as MoonshotChatConfig, ) diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 6073b6b2833..e653b40fd04 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -306,6 +306,7 @@ LLM_CONFIG_NAMES = ( "GigaChatConfig", "GigaChatEmbeddingConfig", "DashScopeChatConfig", + "ModelScopeChatConfig", "MoonshotChatConfig", "DockerModelRunnerChatConfig", "V0ChatConfig", @@ -1161,6 +1162,10 @@ _LLM_CONFIGS_IMPORT_MAP = { ".llms.dashscope.chat.transformation", "DashScopeChatConfig", ), + "ModelScopeChatConfig": ( + ".llms.modelscope.chat.transformation", + "ModelScopeChatConfig", + ), "MoonshotChatConfig": (".llms.moonshot.chat.transformation", "MoonshotChatConfig"), "DockerModelRunnerChatConfig": ( ".llms.docker_model_runner.chat.transformation", diff --git a/litellm/_logging.py b/litellm/_logging.py index 6b99f50e014..bb743c32878 100644 --- a/litellm/_logging.py +++ b/litellm/_logging.py @@ -419,7 +419,7 @@ def _enable_debugging(): def print_verbose(print_statement): try: if set_verbose: - print(redact_secrets(str(print_statement))) # noqa + print(redact_secrets(str(print_statement))) # noqa: T201 except Exception: pass diff --git a/litellm/_redis.py b/litellm/_redis.py index 5ab551453bb..1b6e1a5e4b0 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -311,7 +311,7 @@ def get_redis_url_from_environment(): return f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}" -def _get_redis_client_logic(**env_overrides): # noqa: PLR0915 +def _get_redis_client_logic(**env_overrides): """ Common functionality across sync + async redis client implementations """ @@ -567,7 +567,7 @@ def get_redis_client(**env_overrides): return redis.Redis(**redis_kwargs) -def get_redis_async_client( # noqa: PLR0915 +def get_redis_async_client( connection_pool: Optional[async_redis.BlockingConnectionPool] = None, **env_overrides, ) -> Union[async_redis.Redis, async_redis.RedisCluster]: diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index dcb5cb74ec4..2b6f2cd12b4 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -436,7 +436,7 @@ def _build_streaming_logging_obj( return logging_obj -async def asend_message_streaming( # noqa: PLR0915 +async def asend_message_streaming( a2a_client: Optional["A2AClientType"] = None, request: Optional["SendStreamingMessageRequest"] = None, api_base: Optional[str] = None, diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 15ee9303969..f124882b5a4 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -157,7 +157,7 @@ async def acreate_batch( @client -def create_batch( # noqa: PLR0915 +def create_batch( completion_window: Literal["24h"], endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"], input_file_id: str, diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index b6cfc8e7907..997ad10bc33 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -27,7 +27,7 @@ from litellm.types.utils import EmbeddingResponse, all_litellm_params from .azure_blob_cache import AzureBlobCache from .base_cache import BaseCache from .disk_cache import DiskCache -from .dual_cache import DualCache # noqa +from .dual_cache import DualCache # noqa: F401 from .gcs_cache import GCSCache from .in_memory_cache import InMemoryCache from .qdrant_semantic_cache import QdrantSemanticCache @@ -41,7 +41,7 @@ def print_verbose(print_statement): try: verbose_logger.debug(print_statement) if litellm.set_verbose: - print(print_statement) # noqa + print(print_statement) # noqa: T201 except Exception: pass diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 48691335b40..2a8bd856040 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -394,7 +394,7 @@ class LLMCachingHandler: return cr["model"] return None - def _process_async_embedding_cached_response( # noqa: PLR0915 + def _process_async_embedding_cached_response( self, final_embedding_cached_response: Optional[EmbeddingResponse], cached_result: List[Optional[CachedEmbedding]], diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index cb521efca05..68d3b8c20b3 100644 --- a/litellm/caching/qdrant_semantic_cache.py +++ b/litellm/caching/qdrant_semantic_cache.py @@ -28,7 +28,7 @@ from .base_cache import BaseCache class QdrantSemanticCache(BaseCache): CACHE_KEY_FIELD_NAME = "litellm_cache_key" - def __init__( # noqa: PLR0915 + def __init__( self, qdrant_api_base=None, qdrant_api_key=None, diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 6d8b5cf8a57..3fa6b983e5f 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -693,7 +693,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): original_response = model_call_details.get("original_response") return cls._recover_output_items_from_raw_sse(original_response) - def transform_response( # noqa: PLR0915 + def transform_response( self, model: str, raw_response: "BaseModel", @@ -1211,7 +1211,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): return self.chunk_parser(json.loads(str_line)) @staticmethod - def translate_responses_chunk_to_openai_stream( # noqa: PLR0915 + def translate_responses_chunk_to_openai_stream( parsed_chunk: Union[dict, BaseModel], ) -> "ModelResponseStream": """ @@ -1293,9 +1293,15 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): provider_specific_fields ) + from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, + ) + tool_call_index = parsed_chunk.get("output_index", 0) tool_call_chunk = ChatCompletionToolCallChunk( - id=output_item.get("call_id"), + id=LiteLLMCompletionResponsesConfig._tool_call_id_from_responses_item( + output_item.get("id"), output_item.get("call_id") + ), index=tool_call_index, type="function", function=function_chunk, diff --git a/litellm/constants.py b/litellm/constants.py index 663afb87fb5..b51d15b6d25 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -190,6 +190,10 @@ DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE = int( # Override with LITELLM_MAX_CALLBACKS env var for large deployments (e.g., many teams with guardrails) MAX_CALLBACKS = get_env_int("LITELLM_MAX_CALLBACKS", 100) +# Metadata key recording which pre_call guardrails the proxy loop already ran, +# so the deployment-level hook does not re-run them for the same request +PRE_CALL_EXECUTED_GUARDRAILS_KEY = "_pre_call_executed_guardrails" + # Generic fallback for unknown models DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET = int( os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET", 128) @@ -618,6 +622,7 @@ LITELLM_CHAT_PROVIDERS = [ "nscale", "nebius", "dashscope", + "modelscope", "moonshot", "publicai", "v0", @@ -776,6 +781,7 @@ openai_compatible_endpoints: List = [ "inference.api.nscale.com/v1", "api.studio.nebius.ai/v1", "https://dashscope-intl.aliyuncs.com/compatible-mode/v1", + "https://api-inference.modelscope.cn/v1", "https://api.moonshot.ai/v1", "https://api.publicai.co/v1", "https://api.synthetic.new/openai/v1", @@ -793,6 +799,7 @@ openai_compatible_endpoints: List = [ "https://ai-gateway.vercel.sh/v1", "https://api.inference.wandb.ai/v1", "https://api.clarifai.com/v2/ext/openai/v1", + "https://api.libertai.io/v1", ] @@ -836,10 +843,12 @@ openai_compatible_providers: List = [ "poe", # Poe - JSON-configured provider "chutes", # Chutes - JSON-configured provider "parasail", # Parasail - JSON-configured provider + "libertai", # LibertAI - JSON-configured provider "featherless_ai", "nscale", "nebius", "dashscope", + "modelscope", "moonshot", "v0", "helicone", @@ -865,6 +874,7 @@ openai_text_completion_compatible_providers: List = ( "featherless_ai", "nebius", "dashscope", + "modelscope", "moonshot", "publicai", "synthetic", @@ -1125,6 +1135,48 @@ WANDB_MODELS: set = set( ] ) +modelscope_models: set = set( + [ + # Qwen series models + "Qwen/Qwen3-0.6B", + "Qwen/Qwen3-1.7B", + "Qwen/Qwen3-4B", + "Qwen/Qwen3-8B", + "Qwen/Qwen3-14B", + "Qwen/Qwen3-30B-A3B", + "Qwen/Qwen3-32B", + "Qwen/Qwen3-235B-A22B", + "Qwen/Qwen3-235B-A22B-Instruct-2507", + "Qwen/Qwen3-235B-A22B-Thinking-2507", + "Qwen/Qwen3-30B-A3B-Thinking-2507", + "Qwen/Qwen3-Coder-30B-A3B-Instruct", + "Qwen/Qwen3-Coder-480B-A35B-Instruct", + "Qwen/Qwen3-Next-80B-A3B-Instruct", + "Qwen/Qwen3-Next-80B-A3B-Thinking", + "Qwen/Qwen3-VL-235B-A22B-Instruct", + "Qwen/Qwen3-VL-8B-Instruct", + "Qwen/Qwen3-VL-8B-Thinking", + "Qwen/Qwen3.5-122B-A10B", + "Qwen/Qwen3.5-27B", + "Qwen/Qwen3.5-35B-A3B", + "Qwen/Qwen3.5-397B-A17B", + "Qwen/QwQ-32B", + "Qwen/QwQ-32B-Preview", + "Qwen/QVQ-72B-Preview", + "Qwen/Qwen-Image-Edit", + # DeepSeek series models + "deepseek-ai/DeepSeek-R1-0528", + "deepseek-ai/DeepSeek-R1-Distill-Llama-70B", + "deepseek-ai/DeepSeek-R1-Distill-Llama-8B", + "deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B", + "deepseek-ai/DeepSeek-R1-Distill-Qwen-14B", + "deepseek-ai/DeepSeek-R1-Distill-Qwen-32B", + "deepseek-ai/DeepSeek-R1-Distill-Qwen-7B", + "deepseek-ai/DeepSeek-V3.2", + "deepseek-ai/DeepSeek-V4-Flash", + ] +) + BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[ "cohere", "anthropic", diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index e934c6a6f83..5c77400651b 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -288,7 +288,7 @@ def _transcription_usage_has_token_details( return (prompt_tokens_val > 0) or (completion_tokens_val > 0) -def cost_per_token( # noqa: PLR0915 +def cost_per_token( model: str = "", prompt_tokens: int = 0, completion_tokens: int = 0, @@ -1136,7 +1136,7 @@ def _store_cost_breakdown_in_logging_obj( pass -def completion_cost( # noqa: PLR0915 +def completion_cost( completion_response=None, model: Optional[str] = None, prompt="", diff --git a/litellm/images/main.py b/litellm/images/main.py index d95b7287d20..8b108ded4c9 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -195,7 +195,7 @@ def image_generation( @client -def image_generation( # noqa: PLR0915 +def image_generation( prompt: str, model: Optional[str] = None, n: Optional[int] = None, @@ -738,7 +738,7 @@ def image_variation( @client -def image_edit( # noqa: PLR0915 +def image_edit( image: Optional[Union[FileTypes, List[FileTypes]]] = None, prompt: Optional[str] = None, model: Optional[str] = None, diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 390af2cb6e6..2108ebae312 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -351,7 +351,7 @@ class SlackAlerting(CustomBatchLogger): except Exception: return 0 - async def send_daily_reports(self, router) -> bool: # noqa: PLR0915 + async def send_daily_reports(self, router) -> bool: """ Send a daily report on: - Top 5 deployments with most failed requests @@ -1179,7 +1179,7 @@ Model Info: if response.status_code == 200: return True else: - print("Error sending webhook alert. Error=", response.text) # noqa + print("Error sending webhook alert. Error=", response.text) # noqa: T201 return False @@ -1373,7 +1373,7 @@ Model Info: return False - async def send_alert( # noqa: PLR0915 + async def send_alert( self, message: str, level: Literal["Low", "Medium", "High"], diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 213622cb43a..296bfb6fc85 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -27,6 +27,11 @@ else: LiteLLMLoggingObj = Any +# Anthropic (and Bedrock Claude) reject requests with more than 4 cache_control +# breakpoints: "A maximum of 4 blocks with cache_control may be provided." +MAX_CACHE_CONTROL_BLOCKS = 4 + + class AnthropicCacheControlHook(CustomPromptManagement): def get_chat_completion_prompt( self, @@ -61,16 +66,30 @@ class AnthropicCacheControlHook(CustomPromptManagement): processed_messages = copy.deepcopy(messages) # Separate message-level and non-message-level injection points - remaining_points = [] + message_points: List[CacheControlMessageInjectionPoint] = [] + remaining_points: List[CacheControlInjectionPoint] = [] for point in injection_points: if point.get("location") == "message": - point = cast(CacheControlMessageInjectionPoint, point) - processed_messages = self._process_message_injection( - point=point, messages=processed_messages - ) + message_points.append(cast(CacheControlMessageInjectionPoint, point)) else: remaining_points.append(point) + # Non-message points (currently Bedrock tool_config) are handled in the + # provider transform, where each tool_config point appends at most one + # cachePoint to the tools. That block also counts toward Anthropic's + # limit, so reserve a slot for it here to leave room. + reserved_blocks = ( + 1 + if any(p.get("location") == "tool_config" for p in remaining_points) + else 0 + ) + + processed_messages = self._apply_message_injections( + points=message_points, + messages=processed_messages, + max_blocks=MAX_CACHE_CONTROL_BLOCKS - reserved_blocks, + ) + # Pass through non-message injection points for provider-specific handling if remaining_points: non_default_params["cache_control_injection_points"] = remaining_points @@ -78,14 +97,71 @@ class AnthropicCacheControlHook(CustomPromptManagement): return model, processed_messages, non_default_params @staticmethod - def _process_message_injection( - point: CacheControlMessageInjectionPoint, messages: List[AllMessageValues] + def _apply_message_injections( + points: List[CacheControlMessageInjectionPoint], + messages: List[AllMessageValues], + max_blocks: int, ) -> List[AllMessageValues]: - """Process message-level cache control injection.""" - control: ChatCompletionCachedContent = point.get( - "control", None - ) or ChatCompletionCachedContent(type="ephemeral") + """Apply message-level cache control injection points in order. + Anthropic allows at most ``MAX_CACHE_CONTROL_BLOCKS`` cache_control + breakpoints per request. Client-supplied breakpoints count toward that + limit, so we never inject onto a message that already carries + cache_control (preserving the client's TTL) and we stop injecting once + ``max_blocks`` is reached. Injection points are honored in config order, + so earlier points win when slots are scarce. + """ + used_blocks = sum( + AnthropicCacheControlHook._count_cache_control_blocks(msg) + for msg in messages + ) + + limit_reached = False + for point in points: + if used_blocks >= max_blocks: + limit_reached = True + break + + control: ChatCompletionCachedContent = point.get( + "control", None + ) or ChatCompletionCachedContent(type="ephemeral") + + for target_index in AnthropicCacheControlHook._resolve_target_indices( + point=point, messages=messages + ): + if used_blocks >= max_blocks: + limit_reached = True + break + + if AnthropicCacheControlHook._message_has_cache_control( + messages[target_index] + ): + # Client already marked this message; don't overwrite it. + continue + + messages[target_index] = ( + AnthropicCacheControlHook._safe_insert_cache_control_in_message( + messages[target_index], control + ) + ) + used_blocks += 1 + + if limit_reached: + break + + if limit_reached: + verbose_logger.warning( + f"AnthropicCacheControlHook: Reached the Anthropic limit of " + f"{MAX_CACHE_CONTROL_BLOCKS} cache_control blocks. Skipping further injection." + ) + + return messages + + @staticmethod + def _resolve_target_indices( + point: CacheControlMessageInjectionPoint, messages: List[AllMessageValues] + ) -> List[int]: + """Resolve which message indices an injection point targets.""" _targetted_index: Optional[Union[int, str]] = point.get("index", None) targetted_index: Optional[int] = None if isinstance(_targetted_index, str): @@ -96,36 +172,49 @@ class AnthropicCacheControlHook(CustomPromptManagement): else: targetted_index = _targetted_index - targetted_role = point.get("role", None) - # Case 1: Target by specific index if targetted_index is not None: original_index = targetted_index - # Handle negative indices (convert to positive) if targetted_index < 0: targetted_index += len(messages) if 0 <= targetted_index < len(messages): - messages[targetted_index] = ( - AnthropicCacheControlHook._safe_insert_cache_control_in_message( - messages[targetted_index], control - ) - ) - else: - verbose_logger.warning( - f"AnthropicCacheControlHook: Provided index {original_index} is out of bounds for message list of length {len(messages)}. " - f"Targeted index was {targetted_index}. Skipping cache control injection for this point." - ) + return [targetted_index] + + verbose_logger.warning( + f"AnthropicCacheControlHook: Provided index {original_index} is out of bounds for message list of length {len(messages)}. " + f"Targeted index was {targetted_index}. Skipping cache control injection for this point." + ) + return [] + # Case 2: Target by role - elif targetted_role is not None: - for msg in messages: - if msg.get("role") == targetted_role: - msg = ( - AnthropicCacheControlHook._safe_insert_cache_control_in_message( - message=msg, control=control - ) - ) - return messages + targetted_role = point.get("role", None) + if targetted_role is not None: + return [ + idx + for idx, msg in enumerate(messages) + if msg.get("role") == targetted_role + ] + + return [] + + @staticmethod + def _count_cache_control_blocks(message: AllMessageValues) -> int: + """Count cache_control breakpoints on a message (message + content level).""" + count = 0 + if message.get("cache_control") is not None: + count += 1 + content = message.get("content") + if isinstance(content, list): + for block in content: + if isinstance(block, dict) and block.get("cache_control") is not None: + count += 1 + return count + + @staticmethod + def _message_has_cache_control(message: AllMessageValues) -> bool: + """Return True if the message already carries any cache_control.""" + return AnthropicCacheControlHook._count_cache_control_blocks(message) > 0 @staticmethod def _safe_insert_cache_control_in_message( diff --git a/litellm/integrations/braintrust_logging.py b/litellm/integrations/braintrust_logging.py index 9b1c5077882..6a6313f72e1 100644 --- a/litellm/integrations/braintrust_logging.py +++ b/litellm/integrations/braintrust_logging.py @@ -133,9 +133,7 @@ class BraintrustLogger(CustomLogger): self.default_project_id = project_dict["id"] - def log_success_event( # noqa: PLR0915 - self, kwargs, response_obj, start_time, end_time - ): + def log_success_event(self, kwargs, response_obj, start_time, end_time): verbose_logger.debug("REACHES BRAINTRUST SUCCESS") try: litellm_call_id = kwargs.get("litellm_call_id") @@ -271,9 +269,7 @@ class BraintrustLogger(CustomLogger): except Exception as e: raise e # don't use verbose_logger.exception, if exception is raised - async def async_log_success_event( # noqa: PLR0915 - self, kwargs, response_obj, start_time, end_time - ): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): verbose_logger.debug("REACHES BRAINTRUST SUCCESS") try: litellm_call_id = kwargs.get("litellm_call_id") diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index fc5f0429b63..38245a2e5ba 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -1,3 +1,4 @@ +import secrets from datetime import datetime from typing import ( TYPE_CHECKING, @@ -43,6 +44,7 @@ if TYPE_CHECKING: dc = DualCache() +from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY from litellm.exceptions import ( BlockedPiiEntityError, GuardrailRaisedException, @@ -50,6 +52,12 @@ from litellm.exceptions import ( SensitiveDataRouteException, ) +# Per-process secret tagging each recorded marker. The deployment hook only +# honors markers carrying this token, so a caller cannot forge the metadata +# field to suppress a guardrail on the direct-SDK path that never reaches the +# proxy's metadata sanitizer. +_PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16) + def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]: """Extract session_id from request data (litellm_session_id or metadata).""" @@ -458,6 +466,49 @@ class CustomGuardrail(CustomLogger): return False + def _pre_call_marker(self) -> Optional[str]: + name = self.guardrail_name + if not name: + return None + return f"{_PRE_CALL_EXECUTED_TOKEN}:{name}" + + def mark_pre_call_hook_ran(self, data: Dict[str, Any]) -> None: + """ + Record that this guardrail's ``async_pre_call_hook`` already ran for this + request, so the deployment-level hook does not run it a second time. + + The proxy runs pre-call guardrails in ``ProxyLogging.pre_call_hook``. The + router later spreads a deployment's model-level ``guardrails`` into the + top-level request kwargs, which would otherwise re-trigger the same hook + from ``async_pre_call_deployment_hook``. + """ + marker = self._pre_call_marker() + if marker is None: + return + for meta_key in ("metadata", "litellm_metadata"): + meta = data.get(meta_key) + if isinstance(meta, dict): + executed = meta.get(PRE_CALL_EXECUTED_GUARDRAILS_KEY) + if isinstance(executed, list): + if marker not in executed: + executed.append(marker) + else: + meta[PRE_CALL_EXECUTED_GUARDRAILS_KEY] = [marker] + return + data["metadata"] = {PRE_CALL_EXECUTED_GUARDRAILS_KEY: [marker]} + + def _pre_call_hook_already_ran(self, data: Dict[str, Any]) -> bool: + marker = self._pre_call_marker() + if marker is None: + return False + for meta_key in ("metadata", "litellm_metadata"): + meta = data.get(meta_key) + if isinstance(meta, dict): + executed = meta.get(PRE_CALL_EXECUTED_GUARDRAILS_KEY) + if isinstance(executed, list) and marker in executed: + return True + return False + async def async_pre_call_deployment_hook( self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] ) -> Optional[dict]: @@ -468,6 +519,9 @@ class CustomGuardrail(CustomLogger): if litellm_guardrails is None or not isinstance(litellm_guardrails, list): return kwargs + if self._pre_call_hook_already_ran(kwargs): + return kwargs + if ( self.should_run_guardrail( data=kwargs, event_type=GuardrailEventHooks.pre_call @@ -567,6 +621,9 @@ class CustomGuardrail(CustomLogger): ): return False + if self.default_on is True and disable_global_guardrail is True: + return False + if self.default_on is True and disable_global_guardrail is not True: if self._event_hook_is_event_type(event_type): if isinstance(self.event_hook, Mode): diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index 0efc7d66876..b1c6956a16c 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -549,7 +549,7 @@ class LangFuseLogger: ) ) - def _log_langfuse_v2( # noqa: PLR0915 + def _log_langfuse_v2( self, user_id: Optional[str], metadata: dict, diff --git a/litellm/integrations/lunary.py b/litellm/integrations/lunary.py index b24a24e0881..7b1cbc32d43 100644 --- a/litellm/integrations/lunary.py +++ b/litellm/integrations/lunary.py @@ -75,16 +75,16 @@ class LunaryLogger: version = importlib.metadata.version("lunary") # type: ignore # if version < 0.1.43 then raise ImportError if packaging.version.Version(version) < packaging.version.Version("0.1.43"): # type: ignore - print( # noqa + print( # noqa: T201 "Lunary version outdated. Required: >= 0.1.43. Upgrade via 'pip install lunary --upgrade'" ) raise ImportError self.lunary_client = lunary except ImportError: - print( # noqa + print( # noqa: T201 "Lunary not installed. Please install it using 'pip install lunary'" - ) # noqa + ) raise ImportError def log_event( diff --git a/litellm/integrations/mock_client_factory.py b/litellm/integrations/mock_client_factory.py index 02a927fe64f..9b912ce70c8 100644 --- a/litellm/integrations/mock_client_factory.py +++ b/litellm/integrations/mock_client_factory.py @@ -107,7 +107,7 @@ def _is_url_match(url, matchers: List[str]) -> bool: return False -def create_mock_client_factory(config: MockClientConfig): # noqa: PLR0915 +def create_mock_client_factory(config: MockClientConfig): """ Factory function that creates mock client functions based on configuration. diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index fc37b6a34d8..6b50ef49b49 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -2198,9 +2198,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return kv_pairs - def set_attributes( # noqa: PLR0915 - self, span: Span, kwargs, response_obj: Optional[Any] - ): + def set_attributes(self, span: Span, kwargs, response_obj: Optional[Any]): try: if self.callback_name == "langtrace": from litellm.integrations.langtrace import LangtraceAttributes diff --git a/litellm/integrations/otel/README.md b/litellm/integrations/otel/README.md index 3edb96ed8d9..17011bb8db7 100644 --- a/litellm/integrations/otel/README.md +++ b/litellm/integrations/otel/README.md @@ -216,7 +216,13 @@ lives in [`plumbing/`](./plumbing): `TracerProvider` so one logger serves many tenants. The cache is a bounded LRU that flushes + shuts down evicted providers, since the key derives from request-supplied credentials and must not grow (or leak threads) without limit. -- [`metrics.py`](./plumbing/metrics.py) — GenAI client metric instruments. +- [`metrics.py`](./plumbing/metrics.py) — GenAI client metric instruments. The + six `gen_ai.client.*` histograms are recorded through the meter resolved by + `providers.resolve_meter_provider`: an injected provider wins (tests/DI), + 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). ### Adapter diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index 5e683ce7b99..1869e9ca388 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -10,6 +10,7 @@ from opentelemetry.sdk.trace import TracerProvider from opentelemetry.trace import Span, Tracer, get_current_span, use_span import litellm +from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.otel.model.baggage import promoted_baggage from litellm.integrations.otel.model.config import OpenTelemetryV2Config @@ -36,9 +37,15 @@ from litellm.integrations.otel.model.payloads import ( SpanError, is_mcp_tool_call, ) +from litellm.integrations.otel.plumbing.metrics import ( + GenAIMetricRecorder, + create_genai_metrics, +) from litellm.integrations.otel.plumbing.providers import ( build_tracer_provider, + get_meter, get_tracer, + resolve_meter_provider, ) from litellm.integrations.otel.plumbing.routing import TenantTracerCache from litellm.integrations.otel.model.spans import SpanRole, span_role_for_service @@ -95,7 +102,7 @@ class OpenTelemetryV2(CustomLogger): callback_name: str | None = None, tracer_provider: TracerProvider | None = None, logger_provider: Any | None = None, # reserved for OTel logs - meter_provider: Any | None = None, # reserved for metrics + meter_provider: Any | None = None, **kwargs: Any, ) -> None: super().__init__(**kwargs) @@ -107,6 +114,8 @@ class OpenTelemetryV2(CustomLogger): else build_tracer_provider(self.config) ) self.tracer: Tracer = get_tracer(self._tracer_provider, LITELLM_TRACER_NAME) + self._metrics_recorder = self._init_metrics(meter_provider) + self._metric_filter_error_logged = False self._emitter = SpanEmitter( self.tracer, self.config, mappers=resolve_mappers(self.config.mapper_names) ) @@ -116,6 +125,20 @@ class OpenTelemetryV2(CustomLogger): self._open_llm_calls: "OrderedDict[str, _LLMCallSpan]" = OrderedDict() self._init_otel_logger_on_litellm_proxy() + def _init_metrics(self, meter_provider: Any | None) -> "GenAIMetricRecorder | None": + """Create the six GenAI histograms when metrics are enabled, else ``None``. + + ``meter_provider`` is an explicit override (tests inject one); otherwise the + provider is resolved from the OTel global so the operator's configured + readers/exporters receive the metrics, building and registering one only + when no global provider is set. + """ + if not self.config.enable_metrics: + return None + provider = resolve_meter_provider(self.config, meter_provider) + meter = get_meter(provider, LITELLM_TRACER_NAME) + return GenAIMetricRecorder(create_genai_metrics(meter), self.callback_name) + # ====================================================================== # # Proxy global registration # ====================================================================== # @@ -208,6 +231,25 @@ class OpenTelemetryV2(CustomLogger): if self._emit_mcp_tool_call(kwargs, start_time, end_time): return self._close_llm_call(kwargs, start_time, end_time) + 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.""" + if self._metrics_recorder is None: + return + try: + self._metrics_recorder.record(kwargs, response_obj, start_time, end_time) + except ValueError as exc: + if not self._metric_filter_error_logged: + verbose_logger.error( + "OpenTelemetryV2: invalid otel.attributes metric filter, metrics disabled: %s", + exc, + ) + self._metric_filter_error_logged = True + except Exception as exc: + verbose_logger.debug("OpenTelemetryV2: metric recording failed: %s", exc) async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): if self._emit_mcp_tool_call(kwargs, start_time, end_time): diff --git a/litellm/integrations/otel/mappers/genai.py b/litellm/integrations/otel/mappers/genai.py index d4f14e97a7a..d9be68a06c2 100644 --- a/litellm/integrations/otel/mappers/genai.py +++ b/litellm/integrations/otel/mappers/genai.py @@ -10,7 +10,12 @@ table: one lambda per mapping operation, applied against the typed span data. from typing import Callable 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 ( + collect, + drop_none, + output_messages, + serialize_messages, +) from litellm.integrations.otel.model.payloads import ( GuardrailSpanData, LLMCallSpanData, @@ -47,6 +52,8 @@ class GenAIMapper: else None ), GenAI.REQUEST_SEED: lambda d: d.request_params.seed, + GenAI.INPUT_MESSAGES: lambda d: serialize_messages(d.messages_in), + GenAI.OUTPUT_MESSAGES: lambda d: serialize_messages(output_messages(d)), GenAI.RESPONSE_MODEL: lambda d: d.response_model, GenAI.RESPONSE_ID: lambda d: d.response_id, GenAI.RESPONSE_FINISH_REASONS: lambda d: ( diff --git a/litellm/integrations/otel/model/config.py b/litellm/integrations/otel/model/config.py index ca46182bc66..4f7c3277ebb 100644 --- a/litellm/integrations/otel/model/config.py +++ b/litellm/integrations/otel/model/config.py @@ -184,6 +184,20 @@ class OpenTelemetryV2Config(BaseSettings): ), ) + @field_validator("capture_message_content", mode="before") + @classmethod + def _normalize_capture_message_content(cls, value: object) -> object: + """Fold the capture mode to its canonical lower_snake_case form. + + V1 read this env var case-insensitively, so operators set the + UPPER_SNAKE_CASE form (e.g. ``SPAN_AND_EVENT``). The canonical values + here are lower_snake_case; normalizing at the boundary keeps both + spellings working and lets every downstream comparison stay exact. + """ + if isinstance(value, str): + return value.lower() + return value + @field_validator( "baggage_promoted_keys", "baggage_metadata_keys", diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index bb93a357516..6315a5a4a89 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -230,6 +230,10 @@ class Metric: 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" + RESPONSE_DURATION: Final = "gen_ai.client.response.duration" # litellm ``custom_llm_provider`` -> ``gen_ai.provider.name`` value. diff --git a/litellm/integrations/otel/plumbing/metrics.py b/litellm/integrations/otel/plumbing/metrics.py index edd120f91e6..95ac939ff7f 100644 --- a/litellm/integrations/otel/plumbing/metrics.py +++ b/litellm/integrations/otel/plumbing/metrics.py @@ -1,28 +1,265 @@ -"""GenAI client metrics (token usage + operation duration histograms).""" +"""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. + +The instrument names/units/descriptions and the recording + timing math mirror +the v1 :mod:`litellm.integrations.opentelemetry` integration so both engines emit +identical metrics. The attribute cardinality filter is reused from v1 by import +(no duplication of the valid-name set or its validation). +""" from dataclasses import dataclass +from datetime import datetime +from typing import Any, FrozenSet, Mapping, Optional from opentelemetry.metrics import Histogram, Meter -from litellm.integrations.otel.model.semconv import Metric +import litellm +from litellm.integrations.opentelemetry import ( + METRIC_METADATA_KEYS, + TOKEN_TYPE_ATTRIBUTE, + _build_metric_attribute_filter, + _resolve_metric_attribute_filter, +) +from litellm.integrations.otel.model.semconv import Metric, resolve_operation +from litellm.integrations.otel.model.utils import to_seconds +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps @dataclass(frozen=True) class GenAIMetrics: - token_usage: Histogram operation_duration: Histogram + token_usage: Histogram + token_cost: Histogram + time_to_first_token: Histogram + time_per_output_token: Histogram + response_duration: Histogram def create_genai_metrics(meter: Meter) -> GenAIMetrics: return GenAIMetrics( - token_usage=meter.create_histogram( - name=Metric.TOKEN_USAGE, - unit="{token}", - description="Number of tokens used per GenAI request.", - ), operation_duration=meter.create_histogram( name=Metric.OPERATION_DURATION, unit="s", - description="GenAI operation duration.", + description="GenAI operation duration", + ), + token_usage=meter.create_histogram( + name=Metric.TOKEN_USAGE, + unit="{token}", + description="GenAI token usage", + ), + token_cost=meter.create_histogram( + name=Metric.TOKEN_COST, + unit="USD", + description="GenAI request cost", + ), + time_to_first_token=meter.create_histogram( + name=Metric.TIME_TO_FIRST_TOKEN, + unit="s", + description="Time to first token for streaming requests", + ), + time_per_output_token=meter.create_histogram( + name=Metric.TIME_PER_OUTPUT_TOKEN, + unit="s", + description="Average time per output token (generation time / completion tokens)", + ), + response_duration=meter.create_histogram( + name=Metric.RESPONSE_DURATION, + unit="s", + description="Total LLM API generation time (excludes LiteLLM overhead)", ), ) + + +class GenAIMetricRecorder: + """Records the six GenAI histograms for one successful LLM call. + + The cardinality filter is resolved lazily on the first record: the proxy + populates ``callback_settings.otel.attributes`` after the logger is built, so + reading it at construction time would miss it. ``gen_ai.token.type`` is added + to the token-usage attributes after filtering so the input/output split always + survives. + """ + + def __init__( + self, metrics: GenAIMetrics, callback_name: Optional[str] = None + ) -> None: + self._metrics = metrics + self._callback_name = callback_name + self._include: Optional[FrozenSet[str]] = None + self._exclude: Optional[FrozenSet[str]] = None + self._filter_resolved = False + + def record( + self, + kwargs: Mapping[str, Any], + response_obj: Any, + start_time: datetime, + end_time: datetime, + ) -> None: + common_attrs = self._filter_attributes(self._common_attributes(kwargs)) + duration_s = (end_time - start_time).total_seconds() + + self._metrics.operation_duration.record(duration_s, attributes=common_attrs) + self._record_token_usage(response_obj, common_attrs) + + cost = kwargs.get("response_cost") + if cost: + self._metrics.token_cost.record(cost, attributes=common_attrs) + + self._record_time_to_first_token(kwargs, common_attrs) + self._record_time_per_output_token( + kwargs, response_obj, end_time, duration_s, common_attrs + ) + self._record_response_duration(kwargs, end_time, common_attrs) + + # ------------------------------------------------------------------ # + # 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"), + "gen_ai.framework": "litellm", + } + + std_log = kwargs.get("standard_logging_object") + md = getattr(std_log, "metadata", None) or (std_log or {}).get("metadata", {}) + for key in METRIC_METADATA_KEYS: + value = md.get(key) + if value is None: + continue + if isinstance(value, (dict, list)): + common_attrs[f"metadata.{key}"] = safe_dumps(value) + else: + 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) + + return common_attrs + + def _ensure_filter(self) -> None: + if self._filter_resolved: + return + attributes = None + if self._callback_name in (None, "otel"): + otel_settings = (litellm.callback_settings or {}).get("otel") or {} + raw = ( + otel_settings.get("attributes") + if isinstance(otel_settings, dict) + else None + ) + if raw is not None: + attributes = _build_metric_attribute_filter(raw) + # A bad filter (include_list + exclude_list both set, an unfilterable name) + # raises here; the caller (logger._record_metrics) surfaces it once at ERROR + # so the operator-fixable config error is visible. Not cached on the raise + # path -- _filter_resolved stays False -- so a corrected config takes effect + # without reconstructing the recorder. + self._include, self._exclude = _resolve_metric_attribute_filter(attributes) + self._filter_resolved = True + + def _filter_attributes(self, attrs: dict) -> dict: + self._ensure_filter() + if self._include is not None: + return {k: v for k, v in attrs.items() if k in self._include} + if self._exclude is not None: + return {k: v for k, v in attrs.items() if k not in self._exclude} + return attrs + + # ------------------------------------------------------------------ # + # Per-metric recording + # ------------------------------------------------------------------ # + + def _record_token_usage(self, response_obj: Any, common_attrs: dict) -> None: + if not response_obj: + return + usage = response_obj.get("usage") + if not usage: + return + in_attrs = {**common_attrs, TOKEN_TYPE_ATTRIBUTE: "input"} + out_attrs = {**common_attrs, TOKEN_TYPE_ATTRIBUTE: "output"} + self._metrics.token_usage.record( + usage.get("prompt_tokens", 0), attributes=in_attrs + ) + self._metrics.token_usage.record( + usage.get("completion_tokens", 0), attributes=out_attrs + ) + + def _record_time_to_first_token( + self, kwargs: Mapping[str, Any], common_attrs: dict + ) -> None: + if not kwargs.get("optional_params", {}).get("stream", False): + return + api_call_start = to_seconds(kwargs.get("api_call_start_time")) + completion_start = to_seconds(kwargs.get("completion_start_time")) + if api_call_start is None or completion_start is None: + return + self._metrics.time_to_first_token.record( + completion_start - api_call_start, attributes=common_attrs + ) + + def _record_time_per_output_token( + self, + kwargs: Mapping[str, Any], + response_obj: Any, + end_time: datetime, + duration_s: float, + common_attrs: dict, + ) -> None: + completion_tokens = None + if response_obj and (usage := response_obj.get("usage")): + completion_tokens = usage.get("completion_tokens") + if completion_tokens is None or completion_tokens <= 0: + return + + end_ts = to_seconds(end_time) + if end_ts is None: + generation_time = duration_s + else: + completion_start_time = kwargs.get("completion_start_time") + api_call_start_time = kwargs.get("api_call_start_time") + if completion_start_time is not None: + completion_start = to_seconds(completion_start_time) + generation_time = ( + duration_s + if completion_start is None + else end_ts - completion_start + ) + elif api_call_start_time is not None: + api_call_start = to_seconds(api_call_start_time) + generation_time = ( + duration_s if api_call_start is None else end_ts - api_call_start + ) + else: + generation_time = duration_s + + if generation_time > 0: + self._metrics.time_per_output_token.record( + generation_time / completion_tokens, attributes=common_attrs + ) + + def _record_response_duration( + self, kwargs: Mapping[str, Any], end_time: datetime, common_attrs: dict + ) -> None: + api_call_start_time = kwargs.get("api_call_start_time") + if api_call_start_time is None: + return + _end_time = kwargs.get("end_time") or end_time + if _end_time is None: + _end_time = datetime.now() + api_call_start = to_seconds(api_call_start_time) + end_ts = to_seconds(_end_time) + if api_call_start is None or end_ts is None: + return + duration = end_ts - api_call_start + if duration > 0: + self._metrics.response_duration.record(duration, attributes=common_attrs) diff --git a/litellm/integrations/otel/plumbing/providers.py b/litellm/integrations/otel/plumbing/providers.py index 4c98802479a..6d0710397a3 100644 --- a/litellm/integrations/otel/plumbing/providers.py +++ b/litellm/integrations/otel/plumbing/providers.py @@ -1,9 +1,11 @@ """Provider / exporter factory + the Baggage span processor.""" -from typing import Callable, Iterable +from typing import TYPE_CHECKING, Any, Callable, Iterable -from opentelemetry import baggage +from opentelemetry import baggage, metrics from opentelemetry.context import Context +from opentelemetry.metrics import MeterProvider, NoOpMeterProvider +from opentelemetry.sdk.metrics import MeterProvider as SDKMeterProvider from opentelemetry.sdk.resources import Resource from opentelemetry.sdk.trace import ReadableSpan, SpanProcessor, TracerProvider from opentelemetry.sdk.trace.export import ( @@ -25,6 +27,10 @@ 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 + _SPAN_KIND_BY_ROLE_KIND: dict[LiteLLMSpanKind, SpanKind] = { LiteLLMSpanKind.SERVER: SpanKind.SERVER, LiteLLMSpanKind.CLIENT: SpanKind.CLIENT, @@ -157,6 +163,120 @@ def build_span_exporter(config: OpenTelemetryV2Config) -> SpanExporter: ) +def _otlp_metrics_endpoint(endpoint: str | None) -> str | None: + """Point an OTLP/HTTP base endpoint at the ``/v1/metrics`` signal path. + + The OTLP/HTTP exporter only appends ``/v1/metrics`` when it reads + ``OTEL_EXPORTER_OTLP_ENDPOINT`` itself; an explicitly passed endpoint is used + verbatim, so a base URL would POST to the root. Mirror ``_otlp_traces_endpoint`` + for the metrics signal (rewriting a sibling signal path when present). + """ + if not endpoint: + return endpoint + endpoint = endpoint.rstrip("/") + if endpoint.endswith("/v1/metrics"): + return endpoint + for other_signal in ("/v1/traces", "/v1/logs"): + if endpoint.endswith(other_signal): + return endpoint[: -len(other_signal)] + "/v1/metrics" + return endpoint + "/v1/metrics" + + +def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader": + """Build a metric reader mirroring v1's exporter selection. + + ``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. + """ + from opentelemetry.sdk.metrics.export import ( + ConsoleMetricExporter, + PeriodicExportingMetricReader, + ) + + kind = (config.exporter or "console").lower() + if kind in ("otlp_http", "http", "http/protobuf", "http/json"): + 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, + ) + except ImportError as exc: + raise ImportError( + "OpenTelemetry OTLP gRPC metric exporter is not available. Install " + "`opentelemetry-exporter-otlp` and `grpcio` (or `litellm[grpc]`)." + ) from exc + + exporter = GRPCMetricExporter( + endpoint=config.endpoint, + headers=parse_headers(config.headers), + preferred_temporality={Histogram: AggregationTemporality.DELTA}, + ) + else: + exporter = ConsoleMetricExporter() + + return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) + + +def build_meter_provider( + config: OpenTelemetryV2Config, + metric_reader: "MetricReader | None" = None, +) -> SDKMeterProvider: + """Build the :class:`MeterProvider` for GenAI metrics. + + ``metric_reader`` is an explicit override (tests inject an + ``InMemoryMetricReader``); otherwise the reader is selected from the config's + exporter kind via :func:`build_metric_reader`. + """ + reader = metric_reader if metric_reader is not None else build_metric_reader(config) + return SDKMeterProvider(metric_readers=[reader], resource=build_resource(config)) + + +def resolve_meter_provider( + config: OpenTelemetryV2Config, + meter_provider: MeterProvider | None = None, +) -> MeterProvider: + """Resolve the :class:`MeterProvider` GenAI metrics record through. + + An injected provider wins (DI/tests). Otherwise reuse whatever the operator has + configured as the global, whether a real SDK provider or an explicit + ``NoOpMeterProvider``, so the GenAI histograms ride the operator's + readers/exporters and an explicit opt-out is honored. Only when the global is + still the default proxy placeholder does V2 build one from the config and + publish it as the global, mirroring how V2 owns trace export. The built + provider is the one returned, so its reader thread is always live, never + orphaned. + """ + if meter_provider is not None: + return meter_provider + + existing = metrics.get_meter_provider() + if isinstance(existing, (SDKMeterProvider, NoOpMeterProvider)): + return existing + + provider = build_meter_provider(config) + metrics.set_meter_provider(provider) + return provider + + +def get_meter(provider: MeterProvider, name: str = "litellm") -> "Meter": + return provider.get_meter(name, litellm_version) + + def build_resource(config: OpenTelemetryV2Config) -> Resource: attributes: dict[str, str] = {"service.name": config.service_name} if config.deployment_environment: diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 2119527a8e5..c63f114514a 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -75,7 +75,7 @@ class PrometheusLogger(CustomLogger): return cb return None - def __init__( # noqa: PLR0915 + def __init__( self, **kwargs, ): @@ -2255,7 +2255,7 @@ class PrometheusLogger(CustomLogger): or _litellm_params_metadata.get("user_agent"), } - def set_llm_deployment_failure_metrics(self, request_kwargs: dict): # noqa: PLR0915 + def set_llm_deployment_failure_metrics(self, request_kwargs: dict): """ Sets Failure metrics when an LLM API call fails diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 79f9b16bba0..f29b378fcde 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -1128,7 +1128,7 @@ class WebSearchInterceptionLogger(CustomLogger): ) raise - async def _execute_chat_completion_agentic_loop( # noqa: PLR0915 + async def _execute_chat_completion_agentic_loop( self, model: str, messages: List[Dict], @@ -1159,7 +1159,7 @@ class WebSearchInterceptionLogger(CustomLogger): **request_patch.kwargs, ) - async def _build_chat_completion_request_patch( # noqa: PLR0915 + async def _build_chat_completion_request_patch( self, model: str, messages: List[Dict], diff --git a/litellm/integrations/weights_biases.py b/litellm/integrations/weights_biases.py index e9539d27e97..5f087fe219a 100644 --- a/litellm/integrations/weights_biases.py +++ b/litellm/integrations/weights_biases.py @@ -21,10 +21,11 @@ try: # contains a (known) object attribute object: Literal["chat.completion", "edit", "text_completion"] - def __getitem__(self, key: K) -> V: ... # noqa + def __getitem__(self, key: K) -> V: ... - def get(self, key: K, default: Optional[V] = None) -> Optional[V]: # noqa - ... # pragma: no cover + def get( + self, key: K, default: Optional[V] = None + ) -> Optional[V]: ... # pragma: no cover class OpenAIRequestResponseResolver: def __call__( diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index e984df82140..98b792efa59 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -242,9 +242,28 @@ def _get_parent_otel_span_from_kwargs( return None -def process_response_headers(response_headers: Union[httpx.Headers, dict]) -> dict: +def process_response_headers( + response_headers: Union[httpx.Headers, dict], + preserve_litellm_internal_headers: bool = False, +) -> dict: + """ + `preserve_litellm_internal_headers` must only be True when the input is a + LiteLLM-owned dict (e.g. `_hidden_params["additional_headers"]` that has + already been through one round of processing). For raw upstream provider + headers — whether passed as `httpx.Headers` or a plain dict — it must + remain False, otherwise a malicious provider returning `x-litellm-*` could + spoof LiteLLM-internal markers (e.g. `x-litellm-attempted-fallbacks`). + + When the input is an `httpx.Headers` object the flag is always treated as + False regardless of what the caller requested, because `httpx.Headers` is + always a raw provider response and can never be LiteLLM-owned. + """ from litellm.types.utils import OPENAI_RESPONSE_HEADERS + # Raw httpx.Headers objects come directly from provider HTTP responses and + # must never be treated as LiteLLM-owned, regardless of caller intent. + _preserve = preserve_litellm_internal_headers and isinstance(response_headers, dict) + openai_headers = {} processed_headers = {} additional_headers = {} @@ -256,6 +275,12 @@ def process_response_headers(response_headers: Union[httpx.Headers, dict]) -> di "llm_provider-" ): # return raw provider headers (incl. openai-compatible ones) processed_headers[k] = v + elif _preserve and k.startswith("x-litellm-"): + # LiteLLM's own internal headers (e.g. x-litellm-attempted-fallbacks, + # x-litellm-model-group) are not LLM provider headers and must not be + # prefixed. Downstream consumers (proxy override, callers checking + # whether a fallback happened) look up the bare key. + processed_headers[k] = v else: additional_headers["{}-{}".format("llm_provider", k)] = v diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index ffaa5140916..6087e55b136 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -234,7 +234,7 @@ def extract_and_raise_litellm_exception( ) -def exception_type( # type: ignore # noqa: PLR0915 +def exception_type( # type: ignore model, original_exception, custom_llm_provider, @@ -250,14 +250,14 @@ def exception_type( # type: ignore # noqa: PLR0915 exception_mapping_worked = False exception_provider = custom_llm_provider if litellm.suppress_debug_info is False: - print() # noqa - print( # noqa - "\033[1;31mGive Feedback / Get Help: https://github.com/BerriAI/litellm/issues/new\033[0m" # noqa - ) # noqa - print( # noqa - "LiteLLM.Info: If you need to debug this error, use `litellm._turn_on_debug()'." # noqa - ) # noqa - print() # noqa + print() # noqa: T201 + print( # noqa: T201 + "\033[1;31mGive Feedback / Get Help: https://github.com/BerriAI/litellm/issues/new\033[0m" + ) + print( # noqa: T201 + "LiteLLM.Info: If you need to debug this error, use `litellm._turn_on_debug()'." + ) + print() # noqa: T201 litellm_response_headers = _get_response_headers( original_exception=original_exception diff --git a/litellm/litellm_core_utils/fallback_utils.py b/litellm/litellm_core_utils/fallback_utils.py index daacca85c8a..1606b53e1f9 100644 --- a/litellm/litellm_core_utils/fallback_utils.py +++ b/litellm/litellm_core_utils/fallback_utils.py @@ -7,6 +7,9 @@ from litellm.litellm_core_utils.core_helpers import ( safe_deep_copy, filter_internal_params, ) +from litellm.router_utils.add_retry_fallback_headers import ( + add_fallback_headers_to_response, +) from .asyncify import run_async_function @@ -42,7 +45,7 @@ async def async_completion_with_fallbacks(**kwargs): # Try each fallback model most_recent_exception_str: Optional[str] = None - for fallback in fallbacks: + for attempted_fallbacks, fallback in enumerate(fallbacks): try: completion_kwargs = safe_deep_copy(base_kwargs) # Handle dictionary fallback configurations @@ -63,7 +66,10 @@ async def async_completion_with_fallbacks(**kwargs): ) if response is not None: - return response + return add_fallback_headers_to_response( + response=response, + attempted_fallbacks=attempted_fallbacks, + ) except Exception as e: verbose_logger.exception( diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index 5dc3f5c6868..4941d52d7d6 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -154,7 +154,7 @@ def handle_anthropic_text_model_custom_llm_provider( return model, custom_llm_provider -def get_llm_provider( # noqa: PLR0915 +def get_llm_provider( model: str, custom_llm_provider: Optional[str] = None, api_base: Optional[str] = None, @@ -334,6 +334,9 @@ def get_llm_provider( # noqa: PLR0915 elif endpoint == "dashscope-intl.aliyuncs.com/compatible-mode/v1": custom_llm_provider = "dashscope" dynamic_api_key = get_secret_str("DASHSCOPE_API_KEY") + elif endpoint == "https://api-inference.modelscope.cn/v1": + custom_llm_provider = "modelscope" + dynamic_api_key = get_secret_str("MODELSCOPE_API_KEY") elif endpoint == "api.moonshot.ai/v1": custom_llm_provider = "moonshot" dynamic_api_key = get_secret_str("MOONSHOT_API_KEY") @@ -526,11 +529,11 @@ def get_llm_provider( # noqa: PLR0915 custom_llm_provider = "sap" if not custom_llm_provider: if litellm.suppress_debug_info is False: - print() # noqa - print( # noqa - "\033[1;31mProvider List: https://docs.litellm.ai/docs/providers\033[0m" # noqa - ) # noqa - print() # noqa + print() # noqa: T201 + print( # noqa: T201 + "\033[1;31mProvider List: https://docs.litellm.ai/docs/providers\033[0m" + ) + print() # noqa: T201 error_str = f"LLM Provider NOT provided. Pass in the LLM provider you are trying to call. You passed model={model}\n Pass model as E.g. For 'Huggingface' inference endpoints pass in `completion(model='huggingface/starcoder',..)` Learn more: https://docs.litellm.ai/docs/providers" # maps to openai.NotFoundError, this is raised when openai does not recognize the llm raise litellm.exceptions.BadRequestError( # type: ignore @@ -565,7 +568,7 @@ def get_llm_provider( # noqa: PLR0915 ) -def _get_openai_compatible_provider_info( # noqa: PLR0915 +def _get_openai_compatible_provider_info( model: str, api_base: Optional[str], api_key: Optional[str], @@ -927,6 +930,13 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915 ) = litellm.DashScopeChatConfig()._get_openai_compatible_provider_info( api_base, api_key ) + elif custom_llm_provider == "modelscope": + ( + api_base, + dynamic_api_key, + ) = litellm.ModelScopeChatConfig()._get_openai_compatible_provider_info( + api_base, api_key + ) elif custom_llm_provider == "moonshot": ( api_base, diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index 65c238344e9..e87042b9101 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -5,7 +5,7 @@ from litellm.exceptions import BadRequestError from litellm.types.utils import LlmProviders, LlmProvidersSet -def get_supported_openai_params( # noqa: PLR0915 +def get_supported_openai_params( model: str, custom_llm_provider: Optional[str] = None, request_type: Literal[ diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 2cc8e794d40..0e9c3783316 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -986,7 +986,7 @@ class Logging(LiteLLMLoggingBaseClass): self._get_masked_api_base(additional_args.get("api_base", "")) ) - def pre_call(self, input, api_key, model=None, additional_args={}): # noqa: PLR0915 + def pre_call(self, input, api_key, model=None, additional_args={}): # Log the exact input to the LLM API litellm.error_logs["PRE_CALL"] = locals() try: @@ -2119,7 +2119,7 @@ class Logging(LiteLLMLoggingBaseClass): await self.async_success_handler(result=complete_streaming_response) return - def success_handler( # noqa: PLR0915 + def success_handler( self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs ): verbose_logger.debug( @@ -2584,7 +2584,7 @@ class Logging(LiteLLMLoggingBaseClass): ), ) - async def async_success_handler( # noqa: PLR0915 + async def async_success_handler( self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs ): """ @@ -3036,7 +3036,7 @@ class Logging(LiteLLMLoggingBaseClass): kwargs=self.model_call_details, ) # type: ignore - def failure_handler( # noqa: PLR0915 + def failure_handler( self, exception, traceback_exception, start_time=None, end_time=None ): verbose_logger.debug( @@ -3753,7 +3753,7 @@ def _get_masked_values( } -def set_callbacks(callback_list, function_id=None): # noqa: PLR0915 +def set_callbacks(callback_list, function_id=None): """ Globally sets the callback client """ @@ -3854,7 +3854,7 @@ def set_callbacks(callback_list, function_id=None): # noqa: PLR0915 return None -def _init_custom_logger_compatible_class( # noqa: PLR0915 +def _init_custom_logger_compatible_class( logging_integration: _custom_logger_compatible_callbacks_literal, internal_usage_cache: Optional[DualCache], llm_router: Optional[ @@ -4611,7 +4611,7 @@ def _maybe_auto_initialize_arize_phoenix(_in_memory_loggers: list) -> None: ) -def get_custom_logger_compatible_class( # noqa: PLR0915 +def get_custom_logger_compatible_class( logging_integration: _custom_logger_compatible_callbacks_literal, ) -> Optional[CustomLogger]: try: @@ -5893,7 +5893,7 @@ def get_standard_logging_object_payload( def emit_standard_logging_payload(payload: StandardLoggingPayload): if os.getenv("LITELLM_PRINT_STANDARD_LOGGING_PAYLOAD"): - print(json.dumps(payload, indent=4)) # noqa + print(json.dumps(payload, indent=4)) # noqa: T201 def get_standard_logging_metadata( diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index d75850984a9..a7ac5b53349 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -683,7 +683,7 @@ def _get_regional_uplift_multiplier( return 1.0 -def generic_cost_per_token( # noqa: PLR0915 +def generic_cost_per_token( model: str, usage: Usage, custom_llm_provider: str, diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index 4e5b53a13d7..016bb6b1e22 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -471,7 +471,7 @@ def _should_convert_tool_call_to_json_mode( return False -def convert_to_model_response_object( # noqa: PLR0915 +def convert_to_model_response_object( response_object: Optional[dict] = None, model_response_object: Optional[ Union[ diff --git a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py index 06933a6fbcb..ba870eb9459 100644 --- a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py +++ b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py @@ -49,7 +49,8 @@ class ResponseMetadata: result=self.result, litellm_model_name=model, router_model_id=model_id ), "additional_headers": process_response_headers( - self._get_value_from_hidden_params("additional_headers") or {} + self._get_value_from_hidden_params("additional_headers") or {}, + preserve_litellm_internal_headers=True, ), "litellm_model_name": model, } diff --git a/litellm/litellm_core_utils/logging_callback_manager.py b/litellm/litellm_core_utils/logging_callback_manager.py index 6c749118dec..b7adda3a9a4 100644 --- a/litellm/litellm_core_utils/logging_callback_manager.py +++ b/litellm/litellm_core_utils/logging_callback_manager.py @@ -394,6 +394,22 @@ class LoggingCallbackManager: + litellm._async_failure_callback ) + def remove_callback_from_all_lists(self, obj, require_self=False) -> None: + """ + Remove a callback object from every callback list it may have been + promoted into, so a re-initialized callback leaves no stale instance behind. + """ + for callback_list in ( + litellm.callbacks, + litellm.success_callback, + litellm.failure_callback, + litellm._async_success_callback, + litellm._async_failure_callback, + ): + self.remove_callback_from_list_by_object( + callback_list, obj, require_self=require_self + ) + def get_active_additional_logging_utils_from_custom_logger( self, ) -> Set[AdditionalLoggingUtils]: diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 5059e612f2f..b95b73398ac 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1475,7 +1475,7 @@ def convert_to_gemini_tool_call_invoke( ) -def convert_to_gemini_tool_call_result( # noqa: PLR0915 +def convert_to_gemini_tool_call_result( message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], last_message_with_tool_calls: Optional[dict], model: Optional[str] = None, @@ -2227,7 +2227,7 @@ def _sanitize_empty_text_content( return message -def _add_missing_tool_results( # noqa: PLR0915 +def _add_missing_tool_results( current_message: AllMessageValues, messages: List[AllMessageValues], current_index: int, @@ -2484,7 +2484,7 @@ def sanitize_messages_for_tool_calling( return sanitized_messages -def anthropic_messages_pt( # noqa: PLR0915 +def anthropic_messages_pt( messages: List[AllMessageValues], model: str, llm_provider: str, @@ -3278,7 +3278,7 @@ def convert_to_cohere_tool_invoke(tool_calls: list) -> List[ToolCallObject]: return cohere_tool_invoke -def cohere_messages_pt_v2( # noqa: PLR0915 +def cohere_messages_pt_v2( messages: List, model: str, llm_provider: str, @@ -4703,7 +4703,7 @@ class BedrockConverseMessagesProcessor: return messages @staticmethod - async def _bedrock_converse_messages_pt_async( # noqa: PLR0915 + async def _bedrock_converse_messages_pt_async( messages: List, model: str, llm_provider: str, @@ -5133,7 +5133,7 @@ class BedrockConverseMessagesProcessor: return assistant_parts -def _bedrock_converse_messages_pt( # noqa: PLR0915 +def _bedrock_converse_messages_pt( messages: List, model: str, llm_provider: str, diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index c8f87d96e2f..c56a70177bf 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1198,7 +1198,7 @@ class RealTimeStreaming: item["content"] = new_content return item - async def client_ack_messages(self): # noqa: PLR0915 + async def client_ack_messages(self): try: while True: message = await self.websocket.receive_text() diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index b495b183ec0..04f6b1241c3 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -209,7 +209,7 @@ class ChunkProcessor: ) return response - def get_combined_tool_content( # noqa: PLR0915 + def get_combined_tool_content( self, tool_call_chunks: List[Dict[str, Any]] ) -> List[ChatCompletionMessageToolCall]: tool_calls_list: List[ChatCompletionMessageToolCall] = [] @@ -604,6 +604,8 @@ class ChunkProcessor: usage_chunk = chunk._hidden_params.get("usage", None) if usage_chunk is not None: + if isinstance(usage_chunk, dict): + usage_chunk = Usage(**usage_chunk) usage_chunk_dict = self._usage_chunk_calculation_helper(usage_chunk) if ( usage_chunk_dict["prompt_tokens"] is not None diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index f3274151e5a..888a9658396 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -92,7 +92,7 @@ def is_async_iterable(obj: Any) -> bool: def print_verbose(print_statement): try: if litellm.set_verbose: - print(print_statement) # noqa + print(print_statement) # noqa: T201 except Exception: pass @@ -295,6 +295,12 @@ class CustomStreamWrapper: if len(self.chunks) < 2: return + # Providers like Vertex Gemini (Flash / Flash Lite with web search) emit + # metadata-only / usage-only chunks with no choices. These get stored in + # self.chunks but carry no comparable content, so skip repetition detection. + if not self.chunks[-1].choices or not self.chunks[-2].choices: + return + last_content = self.chunks[-1].choices[0].delta.content if ( @@ -961,7 +967,7 @@ class CustomStreamWrapper: delta, model_response.choices[0].delta, attribute ) - def return_processed_chunk_logic( # noqa + def return_processed_chunk_logic( # noqa: C901 self, completion_obj: Dict[str, Any], model_response: ModelResponseStream, @@ -1139,7 +1145,7 @@ class CustomStreamWrapper: del model_response.choices[0].delta.reasoning_content return - def chunk_creator(self, chunk: Any): # type: ignore # noqa: PLR0915 + def chunk_creator(self, chunk: Any): # type: ignore if hasattr(chunk, "id"): self.response_id = chunk.id model_response = self.model_response_creator() @@ -1881,7 +1887,7 @@ class CustomStreamWrapper: model_response.choices[0].finish_reason = "tool_calls" return model_response - def __next__(self) -> "ModelResponseStream": # noqa: PLR0915 + def __next__(self) -> "ModelResponseStream": cache_hit = False if ( self.custom_llm_provider is not None @@ -2071,7 +2077,7 @@ class CustomStreamWrapper: return self.completion_stream - async def __anext__(self) -> "ModelResponseStream": # noqa: PLR0915 + async def __anext__(self) -> "ModelResponseStream": cache_hit = False if ( self.custom_llm_provider is not None diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index 2fb29b32a61..5d14f3cc4ae 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -772,7 +772,7 @@ class ModelResponseIterator: ) return results - def chunk_parser(self, chunk: dict) -> ModelResponseStream: # noqa: PLR0915 + def chunk_parser(self, chunk: dict) -> ModelResponseStream: try: type_chunk = chunk.get("type", "") or "" diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 9ecd0df0cb8..cf97c946f1c 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -605,7 +605,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) return _tool_choice - def _map_tool_helper( # noqa: PLR0915 + def _map_tool_helper( self, tool: ChatCompletionToolParam, ) -> Tuple[Optional[AllAnthropicToolsValues], Optional[AnthropicMcpServerTool]]: @@ -1399,7 +1399,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return None - def map_openai_params( # noqa: PLR0915 + def map_openai_params( self, non_default_params: dict, optional_params: dict, @@ -2214,18 +2214,33 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if "inference_geo" in _usage and _usage["inference_geo"] is not None: inference_geo = _usage["inference_geo"] - if ( - "cache_creation_input_tokens" in _usage - and _usage["cache_creation_input_tokens"] is not None - ): - cache_creation_input_tokens = _usage["cache_creation_input_tokens"] - prompt_tokens += cache_creation_input_tokens - if ( - "cache_read_input_tokens" in _usage - and _usage["cache_read_input_tokens"] is not None - ): - cache_read_input_tokens = _usage["cache_read_input_tokens"] - prompt_tokens += cache_read_input_tokens + iterations: Optional[List[Any]] = _usage.get("iterations") + if iterations: + prompt_tokens = sum(it.get("input_tokens", 0) or 0 for it in iterations) + completion_tokens = sum( + it.get("output_tokens", 0) or 0 for it in iterations + ) + cache_creation_input_tokens = sum( + it.get("cache_creation_input_tokens", 0) or 0 for it in iterations + ) + cache_read_input_tokens = sum( + it.get("cache_read_input_tokens", 0) or 0 for it in iterations + ) + prompt_tokens += cache_creation_input_tokens + cache_read_input_tokens + + if not iterations: + if ( + "cache_creation_input_tokens" in _usage + and _usage["cache_creation_input_tokens"] is not None + ): + cache_creation_input_tokens = _usage["cache_creation_input_tokens"] + prompt_tokens += cache_creation_input_tokens + if ( + "cache_read_input_tokens" in _usage + and _usage["cache_read_input_tokens"] is not None + ): + cache_read_input_tokens = _usage["cache_read_input_tokens"] + prompt_tokens += cache_read_input_tokens if "server_tool_use" in _usage and _usage["server_tool_use"] is not None: if ( "web_search_requests" in _usage["server_tool_use"] @@ -2264,7 +2279,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ), ) - raw_input_tokens = usage_object.get("input_tokens", 0) or 0 + raw_input_tokens = ( + prompt_tokens - cache_read_input_tokens - cache_creation_input_tokens + ) prompt_tokens_details = PromptTokensDetailsWrapper( cached_tokens=cache_read_input_tokens, cache_creation_tokens=cache_creation_input_tokens, @@ -2296,6 +2313,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): cache_creation_input_tokens=cache_creation_input_tokens, cache_read_input_tokens=cache_read_input_tokens, completion_tokens_details=completion_token_details, + iterations=iterations, server_tool_use=( ServerToolUse( web_search_requests=web_search_requests, 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 f049abcf47f..a8e2fceb4ee 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -372,7 +372,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): cache_read_input_tokens=0, ) - def __next__(self): # noqa: PLR0915 + def __next__(self): from .transformation import LiteLLMAnthropicMessagesAdapter try: @@ -618,7 +618,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): ) raise StopIteration - async def __anext__(self): # noqa: PLR0915 + async def __anext__(self): from .transformation import LiteLLMAnthropicMessagesAdapter try: diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 150f056dc81..bf425637b56 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -332,7 +332,14 @@ class LiteLLMAnthropicMessagesAdapter: if isinstance(source, dict) else getattr(source, "cache_control", None) ) - if cache_control and model and self.is_anthropic_claude_model(model): + if ( + cache_control + and model + and ( + self.is_anthropic_claude_model(model) + or self.is_bedrock_arn_model(model) + ) + ): # TypedDict objects support dict operations at runtime # Use type ignore consistent with codebase pattern (see anthropic/chat/transformation.py:432) if isinstance(target, dict): @@ -376,7 +383,7 @@ class LiteLLMAnthropicMessagesAdapter: isinstance(tool_type, str) and tool_type.startswith("web_search") ) or tool_name == "web_search" - def translate_anthropic_messages_to_openai( # noqa: PLR0915 + def translate_anthropic_messages_to_openai( self, messages: List[ Union[ @@ -752,6 +759,20 @@ class LiteLLMAnthropicMessagesAdapter: model_lower = model.lower() return "anthropic" in model_lower or "claude" in model_lower + @staticmethod + def is_bedrock_arn_model(model: str) -> bool: + """ + Check if the model string is a Bedrock ARN, such as an Application + Inference Profile (e.g. arn:aws:bedrock:us-east-1:123:application-inference-profile/id). + + These ARNs contain neither "anthropic" nor "claude", so is_anthropic_claude_model + cannot identify them even though, on the /v1/messages endpoint, they point at Claude. + Match ":bedrock:" in the ARN service field so another service's ARN that merely names + bedrock in a resource (arn:aws:sagemaker:.../my-bedrock-endpoint) is not matched. + """ + model_lower = model.lower() + return "arn:" in model_lower and ":bedrock:" in model_lower + @staticmethod def translate_thinking_for_model( thinking: Dict[str, Any], diff --git a/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py b/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py index 4aae85b17fe..6479ee999b0 100644 --- a/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py +++ b/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py @@ -97,7 +97,7 @@ def _read_summary_max_tokens_setting() -> int: return COMPACT_SUMMARY_MAX_TOKENS -async def _check_summary_model_access( # noqa: PLR0915 +async def _check_summary_model_access( user_api_key_auth: Any, summary_model: str, llm_router: Any, @@ -970,7 +970,7 @@ def apply_client_compaction_block_history( ) -async def apply_compact_20260112( # noqa: PLR0915 +async def apply_compact_20260112( *, model: str, messages: List[Dict[str, Any]], diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py index 5f1362e259f..04819a416a2 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py @@ -66,7 +66,7 @@ class AnthropicResponsesStreamWrapper: self._current_block_index += 1 return self._current_block_index - def _process_event(self, event: Any) -> None: # noqa: PLR0915 + def _process_event(self, event: Any) -> None: """Convert one Responses API event into zero or more Anthropic chunks queued for emission.""" event_type = getattr(event, "type", None) if event_type is None and isinstance(event, dict): diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py index 2badc2a3276..4fb1ddf5c46 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py @@ -51,7 +51,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: return source.get("url") return None - def translate_messages_to_responses_input( # noqa: PLR0915 + def translate_messages_to_responses_input( self, messages: List[ Union[ diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 56cf035d0f7..5be3ce22832 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -189,7 +189,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): except Exception as e: raise e - def completion( # noqa: PLR0915 + def completion( self, model: str, messages: list, diff --git a/litellm/llms/azure/completion/handler.py b/litellm/llms/azure/completion/handler.py index 05d5e2f6c68..b8d1ad71d46 100644 --- a/litellm/llms/azure/completion/handler.py +++ b/litellm/llms/azure/completion/handler.py @@ -25,7 +25,7 @@ class AzureTextCompletion(BaseAzureLLM): headers["Authorization"] = f"Bearer {azure_ad_token}" return headers - def completion( # noqa: PLR0915 + def completion( self, model: str, messages: list, diff --git a/litellm/llms/base_llm/base_model_iterator.py b/litellm/llms/base_llm/base_model_iterator.py index bf1bfd06537..422ae947997 100644 --- a/litellm/llms/base_llm/base_model_iterator.py +++ b/litellm/llms/base_llm/base_model_iterator.py @@ -1,8 +1,11 @@ import json from abc import abstractmethod -from typing import List, Optional, Union, cast +from typing import TYPE_CHECKING, List, Optional, Union, cast import litellm + +if TYPE_CHECKING: + import httpx from litellm.types.utils import ( Choices, Delta, @@ -69,6 +72,18 @@ class BaseModelResponseIterator: self.streaming_response = streaming_response self.response_iterator = self.streaming_response self.json_mode = json_mode + self.http_response: Optional["httpx.Response"] = None + + async def aclose(self) -> None: + """Close the upstream HTTP response so the provider connection is + released (and a backend like vLLM aborts generation) when the stream + is abandoned before its natural end. + + ``streaming_response`` is usually a bare ``aiter_lines()`` generator + that holds no reference to the response, so the handler that owns the + response attaches it here after construction.""" + if self.http_response is not None: + await self.http_response.aclose() def chunk_parser( self, chunk: dict diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 7e1020000f4..7b1064ccef9 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -248,7 +248,7 @@ class BedrockConverseLLM(BaseAWSLLM): encoding=encoding, ) - def completion( # noqa: PLR0915 + def completion( self, model: str, messages: list, diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index b5e5e4de6fc..bb261ec85b2 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -2189,7 +2189,7 @@ class AmazonConverseConfig(BaseConfig): real_tools = [t for i, t in enumerate(tools) if i not in json_tool_indices] return real_tools if real_tools else None - def _transform_response( # noqa: PLR0915 + def _transform_response( self, model: str, response: httpx.Response, diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 0a1322a751e..75b560b4d6d 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -473,7 +473,7 @@ class BedrockLLM(BaseAWSLLM): prompt += f"{message['content']}" return prompt, chat_history # type: ignore - def process_response( # noqa: PLR0915 + def process_response( self, model: str, response: httpx.Response, @@ -765,7 +765,7 @@ class BedrockLLM(BaseAWSLLM): return model_response - def completion( # noqa: PLR0915 + def completion( self, model: str, messages: list, diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index 6bb2da1ad44..8fc2375c224 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -257,7 +257,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): return request_data - def transform_response( # noqa: PLR0915 + def transform_response( self, model: str, raw_response: httpx.Response, diff --git a/litellm/llms/bedrock/embed/embedding.py b/litellm/llms/bedrock/embed/embedding.py index 27dc785bf57..b6aa99842d7 100644 --- a/litellm/llms/bedrock/embed/embedding.py +++ b/litellm/llms/bedrock/embed/embedding.py @@ -388,7 +388,7 @@ class BedrockEmbedding(BaseAWSLLM): batch_data=batch_data, ) - def embeddings( # noqa: PLR0915 + def embeddings( self, model: str, input: List[str], diff --git a/litellm/llms/bedrock/image_edit/stability_transformation.py b/litellm/llms/bedrock/image_edit/stability_transformation.py index 2d73e47003d..d00d62a8530 100644 --- a/litellm/llms/bedrock/image_edit/stability_transformation.py +++ b/litellm/llms/bedrock/image_edit/stability_transformation.py @@ -149,7 +149,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): return mapped_params - def transform_image_edit_request( # noqa: PLR0915 + def transform_image_edit_request( self, model: str, prompt: Optional[str], diff --git a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py index 2d6bdb5298a..0522bb249e1 100644 --- a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py @@ -226,7 +226,7 @@ class BedrockPassthroughGuardrailHandler(BaseTranslation): return _is_converse_endpoint(endpoint) @staticmethod - async def de_anonymize_event_stream( # noqa: PLR0915 + async def de_anonymize_event_stream( body_bytes: bytes, proxy_logging_obj: "ProxyLogging", user_api_key_dict: "UserAPIKeyAuth", diff --git a/litellm/llms/bytez/chat/transformation.py b/litellm/llms/bytez/chat/transformation.py index 5b08670f9f2..7d9afe01fa6 100644 --- a/litellm/llms/bytez/chat/transformation.py +++ b/litellm/llms/bytez/chat/transformation.py @@ -191,7 +191,7 @@ class BytezChatConfig(BaseConfig): api_key: Optional[str] = None, json_mode: Optional[bool] = None, ) -> ModelResponse: - json = raw_response.json() # noqa: F811 + json = raw_response.json() error = json.get("error") diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index c3f487997c3..8ac5b47c6e7 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -33,7 +33,10 @@ from litellm.llms.base_llm.anthropic_messages.transformation import ( from litellm.llms.base_llm.audio_transcription.transformation import ( BaseAudioTranscriptionConfig, ) -from litellm.llms.base_llm.base_model_iterator import MockResponseIterator +from litellm.llms.base_llm.base_model_iterator import ( + BaseModelResponseIterator, + MockResponseIterator, +) from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig from litellm.llms.base_llm.chat.transformation import BaseConfig from litellm.llms.base_llm.containers.transformation import BaseContainerConfig @@ -814,6 +817,8 @@ class BaseLLMHTTPHandler: completion_stream = provider_config.get_model_response_iterator( streaming_response=response.aiter_lines(), sync_stream=False ) + if isinstance(completion_stream, BaseModelResponseIterator): + completion_stream.http_response = response # LOGGING logging_obj.post_call( input=messages, @@ -5671,7 +5676,7 @@ class BaseLLMHTTPHandler: ) raise - async def async_responses_websocket( # noqa: PLR0915 + async def async_responses_websocket( self, model: str, websocket: Any, diff --git a/litellm/llms/fastcrw/__init__.py b/litellm/llms/fastcrw/__init__.py new file mode 100644 index 00000000000..d65ed8d3fa1 --- /dev/null +++ b/litellm/llms/fastcrw/__init__.py @@ -0,0 +1,7 @@ +""" +fastCRW API integration module. +""" + +from litellm.llms.fastcrw.search.transformation import FastCRWSearchConfig + +__all__ = ["FastCRWSearchConfig"] diff --git a/litellm/llms/fastcrw/search/__init__.py b/litellm/llms/fastcrw/search/__init__.py new file mode 100644 index 00000000000..4f8023b2db4 --- /dev/null +++ b/litellm/llms/fastcrw/search/__init__.py @@ -0,0 +1,7 @@ +""" +fastCRW Search API module. +""" + +from litellm.llms.fastcrw.search.transformation import FastCRWSearchConfig + +__all__ = ["FastCRWSearchConfig"] diff --git a/litellm/llms/fastcrw/search/transformation.py b/litellm/llms/fastcrw/search/transformation.py new file mode 100644 index 00000000000..ce702266e7b --- /dev/null +++ b/litellm/llms/fastcrw/search/transformation.py @@ -0,0 +1,182 @@ +""" +Calls fastCRW's /v1/search endpoint to search the web. + +fastCRW is a Firecrawl-compatible web data engine (single Rust binary; self-host +or cloud). The search response uses the Firecrawl-compatible envelope +{ "success": true, "data": [ { "title", "url", "description", "markdown"? } ] }. + +fastCRW API Reference: https://fastcrw.com/docs/rest-api +""" + +from typing import Optional, TypedDict, Union + +import httpx + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.search.transformation import ( + BaseSearchConfig, + SearchResponse, + SearchResult, +) +from litellm.secret_managers.main import get_secret_str + + +class _FastCRWSearchRequestRequired(TypedDict): + """Required fields for fastCRW Search API request.""" + + query: str # Required - search query + + +class FastCRWSearchRequest(_FastCRWSearchRequestRequired, total=False): + """ + fastCRW Search API request format. + Based on: https://fastcrw.com/docs/rest-api + """ + + limit: int # Optional - maximum number of results to return + sources: list[ + str + ] # Optional - sources to search ('web', 'images'), default ['web'] + scrapeOptions: dict # Optional - options for scraping search results + + +class FastCRWSearchConfig(BaseSearchConfig): + FASTCRW_API_BASE = "https://fastcrw.com/api/v1" + + @staticmethod + def ui_friendly_name() -> str: + return "fastCRW" + + def validate_environment( + self, + headers: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + **kwargs, + ) -> dict: + """ + Validate environment and return headers. + """ + api_key = api_key or get_secret_str("CRW_API_KEY") + if not api_key: + raise ValueError( + "CRW_API_KEY is not set. Set `CRW_API_KEY` environment variable." + ) + headers["Authorization"] = f"Bearer {api_key}" + headers["Content-Type"] = "application/json" + return headers + + def get_complete_url( + self, + api_base: Optional[str], + optional_params: dict, + data: Optional[Union[dict, list[dict]]] = None, + **kwargs, + ) -> str: + """ + Get complete URL for Search endpoint. + """ + api_base = api_base or get_secret_str("CRW_API_BASE") or self.FASTCRW_API_BASE + + # Append "/search" to the api base if it's not already there + if not api_base.endswith("/search"): + api_base = f"{api_base}/search" + + return api_base + + def transform_search_request( + self, + query: Union[str, list[str]], + optional_params: dict, + **kwargs, + ) -> dict: + """ + Transform Search request to fastCRW API format. + + Transforms Perplexity unified spec parameters: + - query -> query (same) + - max_results -> limit + + All other fastCRW-specific parameters are passed through as-is. + + Args: + query: Search query (string or list of strings). fastCRW only supports single string queries. + optional_params: Optional parameters for the request + + Returns: + Dict with typed request data following FastCRWSearchRequest spec + """ + if isinstance(query, list): + # fastCRW only supports single string queries, join with spaces + query = " ".join(query) + + request_data: FastCRWSearchRequest = { + "query": query, + } + + # Transform Perplexity unified spec parameters to fastCRW format + if "max_results" in optional_params: + request_data["limit"] = optional_params["max_results"] + + # Convert to dict before dynamic key assignments + result_data = dict(request_data) + + # pass through all other parameters as-is + for param, value in optional_params.items(): + if ( + param not in self.get_supported_perplexity_optional_params() + and param not in result_data + ): + result_data[param] = value + + # By default, request markdown content if not explicitly specified + # fastCRW doesn't return content unless explicitly requested via scrapeOptions + if "scrapeOptions" not in result_data: + result_data["scrapeOptions"] = { + "formats": ["markdown"], + "onlyMainContent": True, + } + + return result_data + + def transform_search_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + **kwargs, + ) -> SearchResponse: + """ + Transform fastCRW API response to LiteLLM unified SearchResponse format. + + fastCRW (Firecrawl-compatible) returns: + {"success": true, "data": [{"url": "...", "title": "...", "description": "...", "markdown"?: "..."}, ...]} + + Args: + raw_response: Raw httpx response from fastCRW API + logging_obj: Logging object for tracking + + Returns: + SearchResponse with standardized format + """ + response_json = raw_response.json() + + results = [] + + data = response_json.get("data", []) + + if isinstance(data, list): + for result in data: + snippet = result.get("markdown") or result.get("description", "") + search_result = SearchResult( + title=result.get("title", ""), + url=result.get("url", ""), + snippet=snippet, + date=None, + last_updated=None, + ) + results.append(search_result) + + return SearchResponse( + results=results, + object="search", + ) diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index 51fa395d899..74f6cd4d831 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -1378,7 +1378,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): raise ValueError(f"Unknown openai event: {key}, value: {value}") return openai_event - def transform_realtime_response( # noqa: PLR0915 + def transform_realtime_response( self, message: Union[str, bytes], model: str, diff --git a/litellm/llms/huggingface/embedding/transformation.py b/litellm/llms/huggingface/embedding/transformation.py index 88d42cfcdcc..7cddda617a9 100644 --- a/litellm/llms/huggingface/embedding/transformation.py +++ b/litellm/llms/huggingface/embedding/transformation.py @@ -404,7 +404,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig): ) return completion_response - def convert_to_model_response_object( # noqa: PLR0915 + def convert_to_model_response_object( self, completion_response: Union[List[Dict[str, Any]], Dict[str, Any]], model_response: ModelResponse, diff --git a/litellm/llms/modelscope/chat/transformation.py b/litellm/llms/modelscope/chat/transformation.py new file mode 100644 index 00000000000..162ef1a236c --- /dev/null +++ b/litellm/llms/modelscope/chat/transformation.py @@ -0,0 +1,93 @@ +""" +Translates from OpenAI's `/v1/chat/completions` to ModelScope's `/v1/chat/completions` +""" + +from typing import Any, Coroutine, Literal, Optional, Tuple, Union, cast, overload + +from typing_extensions import override + +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import AllMessageValues + +from ...openai.chat.gpt_transformation import OpenAIGPTConfig + + +def _has_non_text_content(message: AllMessageValues) -> bool: + """Check if a message has non-text content items (e.g. image_url).""" + content = message.get("content") + if not isinstance(content, list): + return False + return any(item.get("type") != "text" for item in content) + + +class ModelScopeChatConfig(OpenAIGPTConfig): + DEFAULT_BASE_URL: str = "https://api-inference.modelscope.cn/v1" + + @overload + def _transform_messages( + self, messages: list[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + + @overload + def _transform_messages( + self, + messages: list[AllMessageValues], + model: str, + is_async: Literal[False] = False, + ) -> list[AllMessageValues]: ... + + def _transform_messages( + self, messages: list[AllMessageValues], model: str, is_async: bool = False + ) -> Union[list[AllMessageValues], Coroutine[Any, Any, list[AllMessageValues]]]: + """ + Flatten text-only content lists to strings for ModelScope. + + Messages with non-text content (e.g. image_url for vision models) + are kept as lists so the parent class can normalize them properly. + """ + messages = [cast(AllMessageValues, {**m}) for m in messages] + for message in messages: + if _has_non_text_content(message): + continue + content = message.get("content") + if isinstance(content, list): + message["content"] = "".join(item.get("text") or "" for item in content) + + if is_async: + return super()._transform_messages( + messages=messages, model=model, is_async=True + ) + else: + return super()._transform_messages( + messages=messages, model=model, is_async=False + ) + + def _get_openai_compatible_provider_info( + self, api_base: Optional[str], api_key: Optional[str] + ) -> Tuple[Optional[str], Optional[str]]: + api_base = ( + api_base or get_secret_str("MODELSCOPE_API_BASE") or self.DEFAULT_BASE_URL + ) # type: ignore + dynamic_api_key = api_key or get_secret_str("MODELSCOPE_API_KEY") + return api_base, dynamic_api_key + + @override + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + """ + If api_base is not provided, use the default ModelScope /chat/completions endpoint. + """ + if not api_base: + api_base = self.DEFAULT_BASE_URL + + if not api_base.endswith("/chat/completions"): + api_base = f"{api_base}/chat/completions" + + return api_base diff --git a/litellm/llms/modelscope/image_generation/__init__.py b/litellm/llms/modelscope/image_generation/__init__.py new file mode 100644 index 00000000000..8b28ea962ce --- /dev/null +++ b/litellm/llms/modelscope/image_generation/__init__.py @@ -0,0 +1,31 @@ +""" +ModelScope Image Generation Module + +Factory function for getting the appropriate config class. +""" + +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) + +from .transformation import ModelScopeImageGenerationConfig + +__all__ = [ + "ModelScopeImageGenerationConfig", + "get_modelscope_image_generation_config", +] + + +def get_modelscope_image_generation_config( + model: str, +) -> BaseImageGenerationConfig: + """ + Get the ModelScope config for image generation. + + Args: + model: The model name (e.g., "modelscope/Qwen/Qwen-Image-Edit") + + Returns: + BaseImageGenerationConfig instance for ModelScope + """ + return ModelScopeImageGenerationConfig() diff --git a/litellm/llms/modelscope/image_generation/transformation.py b/litellm/llms/modelscope/image_generation/transformation.py new file mode 100644 index 00000000000..0d85f7796fb --- /dev/null +++ b/litellm/llms/modelscope/image_generation/transformation.py @@ -0,0 +1,248 @@ +""" +ModelScope Image Generation Config + +Handles transformation between OpenAI-compatible format and ModelScope API format. + +API Reference: https://modelscope.cn/docs/model-service/API-Inference/intro +""" + +from typing import TYPE_CHECKING, Optional, Union + +import httpx +from typing_extensions import override + +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import ( + AllMessageValues, + OpenAIImageGenerationOptionalParams, +) +from litellm.types.utils import ImageObject, ImageResponse + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = object + + +class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): + """ + Configuration for ModelScope image generation. + + Supports text-to-image models like: + - Qwen/Qwen-Image-Edit + - And other ModelScope-hosted image generation models + """ + + DEFAULT_BASE_URL: str = "https://api-inference.modelscope.cn/v1" + + def get_supported_openai_params( + self, model: str + ) -> list[OpenAIImageGenerationOptionalParams]: + """ + Return list of OpenAI params supported by ModelScope. + + ModelScope supports standard OpenAI image generation parameters. + """ + return [ + "n", # Number of images to generate + "size", # Size of the generated images + "response_format", # url or b64_json + "user", # User identifier + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + """ + Map OpenAI parameters to ModelScope parameters. + + ModelScope uses the same parameter names as OpenAI. + """ + supported_params = self.get_supported_openai_params(model) + if drop_params: + non_default_params = { + k: v for k, v in non_default_params.items() if k in supported_params + } + optional_params.update(non_default_params) + return optional_params + + @override + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + """ + Get the complete URL for the ModelScope image generation API request. + """ + base_url: str = ( + api_base or get_secret_str("MODELSCOPE_API_BASE") or self.DEFAULT_BASE_URL + ) + base_url = base_url.rstrip("/") + + # Return the images endpoint + return f"{base_url}/images/generations" + + @override + def validate_environment( + self, + headers: dict, + model: str, + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + """ + Validate environment and set up headers for ModelScope. + """ + final_api_key: Optional[str] = api_key or get_secret_str("MODELSCOPE_API_KEY") + + if not final_api_key: + raise ValueError( + "MODELSCOPE_API_KEY is not set. " + "Please set it via environment variable or pass api_key parameter." + ) + + default_headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {final_api_key}", + } + + headers = {**headers, **default_headers} + return headers + + def transform_image_generation_request( + self, + model: str, + prompt: str, + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + """ + Transform OpenAI-style request to ModelScope request format. + + ModelScope uses the same format as OpenAI for image generation. + """ + # Build the request body (same as OpenAI) + request_data: dict = { + "model": model, + "prompt": prompt, + } + + # Add optional params + for key, value in optional_params.items(): + if key.startswith("_"): + continue + request_data[key] = value + + return request_data + + @override + def transform_image_generation_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ImageResponse, + logging_obj: LiteLLMLoggingObj, + request_data: dict, + optional_params: dict, + litellm_params: dict, + encoding: object, + api_key: Optional[str] = None, + json_mode: Optional[bool] = None, + ) -> ImageResponse: + """ + Transform ModelScope response to OpenAI-compatible ImageResponse. + + ModelScope returns the same format as OpenAI: + {"created": timestamp, "data": [{"url": "..."}]} + """ + try: + response_data = raw_response.json() + except Exception as e: + raise self.get_error_class( + error_message=f"Error parsing ModelScope response: {e}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + # Check for errors in response + if "error" in response_data: + error_msg = response_data["error"].get( + "message", str(response_data["error"]) + ) + raise self.get_error_class( + error_message=f"ModelScope error: {error_msg}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + # Extract images from response + data_list = response_data.get("data", []) + if not model_response.data: + model_response.data = [] + + for item in data_list: + image_obj = ImageObject( + url=item.get("url"), + b64_json=item.get("b64_json"), + revised_prompt=item.get("revised_prompt"), + ) + model_response.data.append(image_obj) + + return model_response + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: Union[dict, httpx.Headers], + ) -> BaseLLMException: + """Return the appropriate error class for ModelScope.""" + from litellm.exceptions import ( + AuthenticationError, + BadRequestError, + InternalServerError, + ) + + if status_code == 400: + return BadRequestError( # type: ignore[return-value] + message=error_message, + model="", + llm_provider="modelscope", + ) + elif status_code == 401: + return AuthenticationError( # type: ignore[return-value] + message=error_message, + model="", + llm_provider="modelscope", + ) + elif status_code >= 500: + return InternalServerError( # type: ignore[return-value] + message=error_message, + model="", + llm_provider="modelscope", + ) + else: + return BadRequestError( # type: ignore[return-value] + message=error_message, + model="", + llm_provider="modelscope", + ) diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 194f29648c4..ea905d8ebca 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -608,7 +608,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): return streaming_response - def completion( # type: ignore # noqa: PLR0915 + def completion( # type: ignore self, model_response: ModelResponse, timeout: Union[float, httpx.Timeout], diff --git a/litellm/llms/openai_like/providers.json b/litellm/llms/openai_like/providers.json index 303e9ba8f9e..0dda047d1ca 100644 --- a/litellm/llms/openai_like/providers.json +++ b/litellm/llms/openai_like/providers.json @@ -143,6 +143,14 @@ "force_store_false": true } }, + "libertai": { + "base_url": "https://api.libertai.io/v1", + "api_key_env": "LIBERTAI_API_KEY", + "api_base_env": "LIBERTAI_API_BASE", + "param_mappings": { + "max_completion_tokens": "max_tokens" + } + }, "empiriolabs": { "base_url": "https://api.empiriolabs.ai/v1", "api_key_env": "EMPIRIOLABS_API_KEY", diff --git a/litellm/llms/predibase/chat/transformation.py b/litellm/llms/predibase/chat/transformation.py index 3d251d24b0d..ce004f60bfc 100644 --- a/litellm/llms/predibase/chat/transformation.py +++ b/litellm/llms/predibase/chat/transformation.py @@ -129,7 +129,7 @@ class PredibaseConfig(BaseConfig): optional_params["response_format"] = value return optional_params - def transform_response( # noqa: PLR0915 + def transform_response( self, model: str, raw_response: Response, diff --git a/litellm/llms/sagemaker/completion/handler.py b/litellm/llms/sagemaker/completion/handler.py index de7be18e8ba..aa4663666c2 100644 --- a/litellm/llms/sagemaker/completion/handler.py +++ b/litellm/llms/sagemaker/completion/handler.py @@ -138,7 +138,7 @@ class SagemakerLLM(BaseAWSLLM): return prepped_request - def completion( # noqa: PLR0915 + def completion( self, model: str, messages: list, diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index c578d6cd28b..f5a2b268263 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -678,7 +678,7 @@ def check_if_part_exists_in_parts( return False -def _gemini_convert_messages_with_history( # noqa: PLR0915 +def _gemini_convert_messages_with_history( messages: List[AllMessageValues], model: Optional[str] = None, litellm_params: Optional[dict] = None, @@ -1176,7 +1176,7 @@ def _rewrite_google_maps_response_format(data: RequestBody) -> None: _rewrite_mime_type_to_response_format(generation_config) -def _transform_request_body( # noqa: PLR0915 +def _transform_request_body( messages: List[AllMessageValues], model: str, optional_params: dict, 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 3ec7b0814dd..dab21e2ce8e 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 @@ -614,9 +614,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return googleSearch, googleSearchRetrieval, enterpriseWebSearch, urlContext - def _map_function( # noqa: PLR0915 - self, value: List[dict], optional_params: dict - ) -> List[Tools]: + def _map_function(self, value: List[dict], optional_params: dict) -> List[Tools]: """ Map OpenAI-style tools/functions to Vertex AI format. @@ -1173,7 +1171,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): optional_params["include_server_side_tool_invocations"] = True return - def map_openai_params( # noqa: PLR0915 + def map_openai_params( self, non_default_params: Dict, optional_params: Dict, @@ -1904,7 +1902,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return False @staticmethod - def _calculate_usage( # noqa: PLR0915 + def _calculate_usage( completion_response: Union[ GenerateContentResponseBody, BidiGenerateContentServerMessage ], @@ -2380,7 +2378,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return annotations @staticmethod - def _process_candidates( # noqa: PLR0915 + def _process_candidates( _candidates: List[Candidates], model_response: Union[ModelResponse, "ModelResponseStream"], standard_optional_params: dict, diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py index 99165c37c93..165dac24903 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py @@ -124,7 +124,7 @@ class GoogleBatchEmbeddings(VertexLLM): return resolved_files - def batch_embeddings( # noqa: PLR0915 + def batch_embeddings( self, model: str, input: GeminiEmbeddingInput, diff --git a/litellm/llms/vertex_ai/vertex_ai_non_gemini.py b/litellm/llms/vertex_ai/vertex_ai_non_gemini.py index cfbab584f6a..c134dee7ad4 100644 --- a/litellm/llms/vertex_ai/vertex_ai_non_gemini.py +++ b/litellm/llms/vertex_ai/vertex_ai_non_gemini.py @@ -77,7 +77,7 @@ def _set_client_in_cache(client_cache_key: str, vertex_llm_model: Any): ) -def completion( # noqa: PLR0915 +def completion( model: str, messages: list, model_response: ModelResponse, @@ -485,7 +485,7 @@ def completion( # noqa: PLR0915 ) -async def async_completion( # noqa: PLR0915 +async def async_completion( llm_model, mode: str, prompt: str, @@ -650,7 +650,7 @@ async def async_completion( # noqa: PLR0915 raise VertexAIError(status_code=500, message=str(e)) -async def async_streaming( # noqa: PLR0915 +async def async_streaming( llm_model, mode: str, prompt: str, diff --git a/litellm/main.py b/litellm/main.py index 18dcdfcd6be..80176cc8b16 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -392,7 +392,7 @@ class AsyncCompletions: @tracer.wrap() @client -async def acompletion( # noqa: PLR0915 +async def acompletion( model: str, # Optional OpenAI params: see https://platform.openai.com/docs/api-reference/chat/create messages: List = [], @@ -1086,7 +1086,7 @@ def _build_custom_pricing_entry( @tracer.wrap() @client -def completion( # type: ignore # noqa: PLR0915 +def completion( # type: ignore model: str, # Optional OpenAI params: see https://platform.openai.com/docs/api-reference/chat/create messages: List = [], @@ -4878,7 +4878,7 @@ def embedding( @client -def embedding( # noqa: PLR0915 +def embedding( model, input=[], # Optional params @@ -6125,7 +6125,7 @@ async def atext_completion( @client -def text_completion( # noqa: PLR0915 +def text_completion( prompt: Union[ str, List[Union[str, List[Union[str, List[int]]]]] ], # Required: The prompt(s) to generate completions for. @@ -6664,7 +6664,7 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse: @client -def transcription( # noqa: PLR0915 +def transcription( model: str, file: FileTypes, ## OPTIONAL OPENAI PARAMS ## @@ -6971,7 +6971,7 @@ async def aspeech(*args, **kwargs) -> HttpxBinaryResponseContent: @client -def speech( # noqa: PLR0915 +def speech( model: str, input: str, voice: Optional[Union[str, dict]] = None, @@ -7572,7 +7572,7 @@ def print_verbose(print_statement): try: verbose_logger.debug(print_statement) if litellm.set_verbose: - print(print_statement) # noqa + print(print_statement) # noqa: T201 except Exception: pass @@ -7662,7 +7662,7 @@ def stream_chunk_builder_text_completion( return TextCompletionResponse(**response) -def stream_chunk_builder( # noqa: PLR0915 +def stream_chunk_builder( chunks: list, messages: Optional[list] = None, start_time=None, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 76a7c0640af..f563ad0c5b5 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -18822,6 +18822,38 @@ "supports_response_schema": true, "supports_vision": true }, + "github_copilot/mai-code-1-flash": { + "cache_read_input_token_cost": 7.5e-08, + "input_cost_per_token": 7.5e-07, + "litellm_provider": "github_copilot", + "max_input_tokens": 128000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 4.5e-06, + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true + }, + "github_copilot/mai-code-1-flash-internal": { + "cache_read_input_token_cost": 7.5e-08, + "input_cost_per_token": 7.5e-07, + "litellm_provider": "github_copilot", + "max_input_tokens": 128000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 4.5e-06, + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true + }, "github_copilot/text-embedding-3-small": { "litellm_provider": "github_copilot", "max_input_tokens": 8191, @@ -40784,6 +40816,174 @@ "litellm_provider": "llamagate", "mode": "embedding" }, + "libertai/hermes-3-8b-tee": { + "max_tokens": 16000, + "max_input_tokens": 16000, + "max_output_tokens": 16000, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": false, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/gemma-4-31b-it": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 4e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/gemma-4-31b-it-thinking": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 4e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_reasoning": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.6-27b": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.6-27b-thinking": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_reasoning": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.6-35b-a3b": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.6-35b-a3b-thinking": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_reasoning": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.5-122b-a10b": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.75e-06, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.5-122b-a10b-thinking": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.75e-06, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_reasoning": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/deepseek-v4-flash": { + "max_tokens": 200000, + "max_input_tokens": 200000, + "max_output_tokens": 200000, + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.75e-06, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": false, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/deepseek-v4-flash-thinking": { + "max_tokens": 200000, + "max_input_tokens": 200000, + "max_output_tokens": 200000, + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.75e-06, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": false, + "supports_reasoning": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/bge-m3": { + "max_tokens": 8192, + "max_input_tokens": 8192, + "input_cost_per_token": 1e-08, + "output_cost_per_token": 0.0, + "litellm_provider": "libertai", + "mode": "embedding", + "source": "https://docs.libertai.io/apis/text/" + }, "sarvam/sarvam-m": { "cache_creation_input_token_cost": 0, "cache_creation_input_token_cost_above_1hr": 0, diff --git a/litellm/mypy.ini b/litellm/mypy.ini index 4702b591124..b65e11bab42 100644 --- a/litellm/mypy.ini +++ b/litellm/mypy.ini @@ -1,13 +1,16 @@ [mypy] -warn_return_any = False +warn_return_any = True ignore_missing_imports = True +disallow_untyped_defs = True mypy_path = litellm/stubs namespace_packages = True disable_error_code = - valid-type, annotation-unchecked, import-untyped +[mypy-litellm.*] +ignore_missing_imports = False + [mypy-google.*] ignore_missing_imports = True diff --git a/litellm/provider_endpoints_support_backup.json b/litellm/provider_endpoints_support_backup.json index e0eeb014c51..db6183edaa0 100644 --- a/litellm/provider_endpoints_support_backup.json +++ b/litellm/provider_endpoints_support_backup.json @@ -1288,6 +1288,23 @@ "interactions": true } }, + "libertai": { + "display_name": "LibertAI (`libertai`)", + "url": "https://docs.litellm.ai/docs/providers/libertai", + "endpoints": { + "chat_completions": true, + "messages": true, + "responses": false, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": false + } + }, "litellm_proxy": { "display_name": "LiteLLM Proxy (`litellm_proxy`)", "url": "https://docs.litellm.ai/docs/providers/litellm_proxy", 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 dcf7660d002..e47fc84b533 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 @@ -67,9 +67,10 @@ def _is_mcp_passthrough_cold_start( spec-compliant WWW-Authenticate challenge instead of surfacing a generic admission error. - Uses "all" semantics (mirrors :meth:`MCPRequestHandler._target_servers_use_oauth2`): - one non-passthrough target in a co-targeted set must not flip the bypass - open for the others. Fails closed when any target cannot be resolved.""" + Uses "all" semantics (mirrors + :meth:`MCPRequestHandler._target_servers_delegate_auth_to_upstream`): one + non-passthrough target in a co-targeted set must not flip the bypass open + for the others. Fails closed when any target cannot be resolved.""" if not mcp_servers: return False from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -124,7 +125,7 @@ class MCPRequestHandler: LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value @staticmethod - async def process_mcp_request( # noqa: PLR0915 + async def process_mcp_request( scope: Scope, ) -> Tuple[ UserAPIKeyAuth, @@ -214,101 +215,64 @@ class MCPRequestHandler: # Only OAuth metadata routes registered under /.well-known/ are public. if request_route.startswith("/.well-known/"): validated_user_api_key_auth = UserAPIKeyAuth() - elif ( - not litellm_api_key - and MCPRequestHandler._target_servers_delegate_auth_to_upstream( # noqa: E501 - path=request_route, - mcp_servers=mcp_servers, - client_ip=IPAddressUtils.get_mcp_client_ip(request), - ) - ): - # Operator opted this oauth2 server into upstream-delegated auth - # (PKCE passthrough): skip LiteLLM API-key/SSO entirely so the - # client authenticates directly with the upstream MCP server. - # Fires ONLY when neither x-litellm-api-key nor Authorization is - # present. If any LiteLLM key is supplied (primary or secondary - # header), we fall through so user_id is resolved, spend/rate - # limiting apply, and any stored OAuth token can be retrieved - # and forwarded upstream. Gated by - # _target_servers_delegate_auth_to_upstream, which only returns - # True when EVERY target is auth_type=oauth2 AND has the - # delegate_auth_to_upstream flag set — fails closed otherwise. - validated_user_api_key_auth = UserAPIKeyAuth() elif has_explicit_litellm_key: - # Explicit x-litellm-api-key provided - always validate normally + # An explicit x-litellm-api-key is always a LiteLLM credential, even + # for a delegated server, so validate it: identity / spend / rate + # limits resolve and any stored upstream token can be forwarded. validated_user_api_key_auth = await user_api_key_auth( api_key=litellm_api_key, request=request ) + elif MCPRequestHandler._target_servers_delegate_auth_to_upstream( + path=request_route, + mcp_servers=mcp_servers, + client_ip=IPAddressUtils.get_mcp_client_ip(request), + ): + # Operator opted this oauth2 server into upstream-delegated auth: the + # client authenticates directly with the upstream MCP server, so any + # Authorization bearer is an upstream token, never a LiteLLM key. Skip + # LiteLLM validation entirely — covering both the no-credential + # discovery request and the authenticated call carrying the upstream + # bearer — so a tool call that succeeds never carries a phantom 401 + # auth span; the bearer is forwarded upstream unchanged. Gated by + # _target_servers_delegate_auth_to_upstream, which returns True only + # when EVERY target is auth_type=oauth2 with delegate_auth_to_upstream + # set; fails closed otherwise. + validated_user_api_key_auth = UserAPIKeyAuth() elif oauth2_headers: - # No x-litellm-api-key, but Authorization header present. - # Could be a LiteLLM key (backward compat) OR an opaque OAuth2 token - # the operator wants forwarded to an upstream OAuth2-mode MCP server. - # Try LiteLLM auth first; on auth failure, only fall back to anonymous - # passthrough when the request actually targets a server whose operator - # configured ``auth_type=oauth2``. For any other server (api_key, - # bearer_token, basic, etc.), a failed LiteLLM auth is a real failure - # and must propagate — otherwise an attacker can exchange any garbage - # bearer for an anonymous session. + # Authorization on a non-delegated server: the bearer must be a real + # LiteLLM credential, so a failed validation is a genuine 401/403 and + # propagates. The sole anonymous fallback is the auth_type=none + # pass-through cold-start (RFC 9728 discovery return), gated on a 401 + # so a recognized-but-forbidden key still fails closed. + client_ip = IPAddressUtils.get_mcp_client_ip(request) try: validated_user_api_key_auth = await user_api_key_auth( api_key=litellm_api_key, request=request ) except (HTTPException, ProxyException) as e: - # HTTPException.status_code is int; ProxyException.code is - # normalized to str in its __init__ but can be ``"None"`` or any - # non-numeric string when the caller didn't supply a numeric - # code, so we compare against both int and str forms rather - # than coercing (``int("None")`` would raise ValueError and - # rewrite the auth error as a 500). + # ProxyException.code is normalized to str (possibly "None"), so + # compare both int and str forms rather than coercing. status = e.status_code if isinstance(e, HTTPException) else e.code - is_auth_error = status in (401, 403, "401", "403") is_unauthenticated = status in (401, "401") - client_ip = IPAddressUtils.get_mcp_client_ip(request) - if is_auth_error and MCPRequestHandler._target_servers_use_oauth2( - path=request_route, - mcp_servers=mcp_servers, - client_ip=client_ip, + mcp_servers_from_path = _parse_mcp_server_names_from_path( + request_route, mcp_servers + ) + if ( + is_unauthenticated + and mcp_servers_from_path is not None + and not _has_client_supplied_mcp_auth( + mcp_auth_header, + mcp_server_auth_headers, + ) + and _is_mcp_passthrough_cold_start( + mcp_servers_from_path, client_ip=client_ip + ) ): verbose_logger.debug( - "MCP OAuth2: target server is OAuth2-mode, treating " - "Authorization as upstream OAuth2 token passthrough" + "MCP pass-through return: forwarding Authorization as " + "upstream OAuth token for delegated auth" ) validated_user_api_key_auth = UserAPIKeyAuth() - elif is_unauthenticated: - # Pass-through cold-start return: per RFC 9728 / MCP - # Authorization spec the client completes upstream OAuth - # discovery and returns with ``Authorization: Bearer - # ``. For ``auth_type=none`` passthrough - # servers that bearer is not a LiteLLM key (auth above - # failed) but is meant to be forwarded upstream - # unchanged. Fall back to anonymous admission so the - # caller is not rejected for following the discovery - # flow without also setting ``x-litellm-api-key``. - # Only trigger on 401 (token unrecognized); a 403 means - # the key WAS recognized but is forbidden (e.g. over - # budget / rate limited) and must propagate so those - # controls are not bypassed via anonymous admission. - mcp_servers_from_path = _parse_mcp_server_names_from_path( - request_route, mcp_servers - ) - if ( - mcp_servers_from_path is not None - and not _has_client_supplied_mcp_auth( - mcp_auth_header, - mcp_server_auth_headers, - ) - and _is_mcp_passthrough_cold_start( - mcp_servers_from_path, client_ip=client_ip - ) - ): - verbose_logger.debug( - "MCP pass-through return: target server is " - "passthrough, treating Authorization as " - "upstream OAuth token for delegated auth" - ) - validated_user_api_key_auth = UserAPIKeyAuth() - else: - raise else: raise else: @@ -412,45 +376,6 @@ class MCPRequestHandler: return [single_server_match.group(1)] return [servers_and_path] - @staticmethod - def _target_servers_use_oauth2( - path: str, mcp_servers: Optional[List[str]], client_ip: Optional[str] - ) -> bool: - """ - True only when EVERY MCP server the request targets is configured for - ``auth_type == oauth2``. If any target is non-OAuth2 — or if the target - cannot be resolved at all — return False so the caller fails closed. - - Used to gate the "treat Authorization as opaque OAuth2 token" fallback - in :meth:`process_mcp_request` so a failed LiteLLM-auth cannot be - exchanged for an anonymous session against a non-OAuth2 server. - """ - # Inline imports avoid a circular dependency: mcp_server_manager imports - # from this module. - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - from litellm.types.mcp import MCPAuth - - # Resolve the same target list downstream routing will use. For - # ``/mcp/...`` routes, ``extract_mcp_auth_context`` overrides the - # ``x-mcp-servers`` header with path-derived names, so we must mirror - # that here — otherwise a caller could set the header to a permissive - # server while the path targets a stricter one (header/path TOCTOU). - target_names = MCPRequestHandler._resolve_target_server_names( - path=path, mcp_servers_header=mcp_servers - ) - if not target_names: - return False - - for name in target_names: - server = global_mcp_server_manager.get_mcp_server_by_name( - name, client_ip=client_ip - ) - if server is None or server.auth_type != MCPAuth.oauth2: - return False - return True - @staticmethod def _target_servers_delegate_auth_to_upstream( path: str, mcp_servers: Optional[List[str]], client_ip: Optional[str] @@ -472,8 +397,8 @@ class MCPRequestHandler: ) from litellm.types.mcp import MCPAuth - # See _target_servers_use_oauth2: must mirror the downstream - # header-vs-path override or an attacker could set + # Must mirror the downstream header-vs-path override + # (``extract_mcp_auth_context``) or an attacker could set # ``x-mcp-servers`` to a delegate-enabled server while the URL path # targets a non-delegate server, skipping LiteLLM auth for it. target_names = MCPRequestHandler._resolve_target_server_names( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 5e419b5c0a3..3c3f2afad6d 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3355,7 +3355,7 @@ class MCPServerManager: ) ) - async def _call_regular_mcp_tool( # noqa: PLR0915 + async def _call_regular_mcp_tool( self, mcp_server: MCPServer, original_tool_name: str, diff --git a/litellm/proxy/_experimental/mcp_server/sampling_handler.py b/litellm/proxy/_experimental/mcp_server/sampling_handler.py index 1637c9eb0b9..b659ba6f813 100644 --- a/litellm/proxy/_experimental/mcp_server/sampling_handler.py +++ b/litellm/proxy/_experimental/mcp_server/sampling_handler.py @@ -661,7 +661,7 @@ def _convert_openai_response_to_mcp_result( ) -async def _check_model_access( # noqa: PLR0915 +async def _check_model_access( model: str, user_api_key_auth: Any ) -> Optional["ErrorData"]: """Enforce model-permission checks for MCP sampling requests. diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 746fc4e7d3f..1d9a4479f05 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -617,7 +617,7 @@ if MCP_AVAILABLE: active_mcp_session_var.reset(_session_reset_token) @server.call_tool() - async def mcp_server_tool_call( # noqa: PLR0915 + async def mcp_server_tool_call( name: str, arguments: Dict[str, Any] | None ) -> CallToolResult: """ @@ -1591,7 +1591,7 @@ if MCP_AVAILABLE: _mcp_gateway_initialize_instructions.reset(instructions_token) _mcp_gateway_server_name.reset(server_name_token) - async def _get_tools_from_mcp_servers( # noqa: PLR0915 + async def _get_tools_from_mcp_servers( user_api_key_auth: Optional[UserAPIKeyAuth], mcp_auth_header: Optional[str], mcp_servers: Optional[List[str]], @@ -2435,7 +2435,7 @@ if MCP_AVAILABLE: }, ) - async def execute_mcp_tool( # noqa: PLR0915 + async def execute_mcp_tool( name: str, arguments: Dict[str, Any], allowed_mcp_servers: List[MCPServer], @@ -3642,7 +3642,7 @@ if MCP_AVAILABLE: detail="Forbidden", ) - async def handle_streamable_http_mcp( # noqa: PLR0915 + async def handle_streamable_http_mcp( scope: Scope, receive: Receive, send: Send ) -> None: """Handle MCP requests through StreamableHTTP.""" diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 493e09e3af1..765e90bc896 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -373,6 +373,8 @@ class LiteLLMRoutes(enum.Enum): # vector stores "/vector_stores", "/v1/vector_stores", + "/vector_stores/{vector_store_id}", + "/v1/vector_stores/{vector_store_id}", "/vector_stores/{vector_store_id}/search", "/v1/vector_stores/{vector_store_id}/search", "/vector_stores/{vector_store_id}/files", @@ -2150,6 +2152,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): master_key: Optional[str] = Field( None, description="require a key for all calls to proxy" ) + allow_cli_sso_verification_uri_complete: bool | None = Field( + None, + description="opt-in to RFC 8628 verification_uri_complete for the CLI SSO device flow, pre-filling the user_code in the browser. Off by default; intended for same-host clients where the device that starts the flow and the browser run on the same machine", + ) database_url: Optional[str] = Field( None, description="connect to a postgres db - needed for generating temporary keys + tracking spend / key", diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 7b2f75e1cff..7446f61ad1c 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -509,7 +509,7 @@ async def get_agent_card( tags=["[beta] A2A Agents"], dependencies=[Depends(user_api_key_auth)], ) -async def invoke_agent_a2a( # noqa: PLR0915 +async def invoke_agent_a2a( agent_id: str, request: Request, fastapi_response: Response, diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 900386f3d7b..1995ff275c9 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -28,7 +28,7 @@ router = APIRouter() tags=["[beta] Anthropic `/v1/messages`"], dependencies=[Depends(user_api_key_auth)], ) -async def anthropic_response( # noqa: PLR0915 +async def anthropic_response( fastapi_response: Response, request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index aa967732a90..814346eddf8 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -519,7 +519,7 @@ MODEL_DISCOVERY_ROUTES = frozenset( ) -async def common_checks( # noqa: PLR0915 +async def common_checks( request_body: dict, team_object: Optional[LiteLLM_TeamTable], user_object: Optional[LiteLLM_UserTable], diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index fd6ff2ada7f..90845dfd824 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -1954,7 +1954,7 @@ class JWTAuthManager: return None, None, None @staticmethod - async def auth_builder( # noqa: PLR0915 + async def auth_builder( api_key: str, jwt_handler: JWTHandler, request_data: dict, diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index d0818b95363..bd2e7560430 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -103,7 +103,7 @@ class LoginResult: self.login_method = login_method -async def authenticate_user( # noqa: PLR0915 +async def authenticate_user( username: str, password: str, master_key: Optional[str], diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 666c01562b5..6f359e52eeb 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -979,7 +979,7 @@ def _ensure_parent_otel_span_on_request_state(request: Request) -> None: request.state.parent_otel_span = parent_otel_span -async def _user_api_key_auth_builder( # noqa: PLR0915 +async def _user_api_key_auth_builder( request: Request, api_key: str, azure_api_key_header: str, @@ -2126,7 +2126,7 @@ def _team_obj_from_token(valid_token: UserAPIKeyAuth) -> LiteLLM_TeamTableCached @tracer.wrap() -async def _run_centralized_common_checks( # noqa: PLR0915 +async def _run_centralized_common_checks( user_api_key_auth_obj: UserAPIKeyAuth, request: Request, request_data: dict, diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index ea479a5721b..344f90aa144 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -58,7 +58,7 @@ router = APIRouter() dependencies=[Depends(user_api_key_auth)], tags=["batch"], ) -async def create_batch( # noqa: PLR0915 +async def create_batch( request: Request, fastapi_response: Response, provider: Optional[str] = None, @@ -343,7 +343,7 @@ async def create_batch( # noqa: PLR0915 dependencies=[Depends(user_api_key_auth)], tags=["batch"], ) -async def retrieve_batch( # noqa: PLR0915 +async def retrieve_batch( request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 90ad0f28808..41cadd5bbc3 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -17,10 +17,12 @@ from typing import ( Union, ) +import anyio import httpx import orjson from fastapi import HTTPException, Request, status from fastapi.responses import JSONResponse, Response, StreamingResponse +from starlette.types import Receive, Scope, Send import litellm from litellm._logging import _redact_string, verbose_proxy_logger @@ -240,7 +242,65 @@ def _extract_error_from_sse_chunk(event_line: Union[str, bytes]) -> dict: return default_error -async def create_response( # noqa: PLR0915 +async def _aclose_upstream_response(response: Any) -> None: + """Release the upstream HTTP connection when a stream ends for any + reason, including client disconnect. Mirrors the finally block of + async_data_generator in proxy_server.py.""" + with anyio.CancelScope(shield=True): + if hasattr(response, "aclose"): + try: + await response.aclose() + except BaseException as e: + verbose_proxy_logger.debug( + "error closing upstream response stream: %s", e + ) + + +class _UpstreamClosingStreamingResponse(StreamingResponse): + """StreamingResponse that always closes its body iterator and the wrapped + upstream generator. + + When the client disconnects mid-stream, Starlette abandons the body + iterator without calling aclose(), leaving the upstream LLM connection + open until garbage collection; the backend (e.g. vLLM) keeps generating + into a dead pipe. The upstream generator is closed directly (not via the + body iterator) because aclose() on a never-started generator skips its + body, so a cascade through it would be a no-op if the client disconnects + before the first chunk is sent. + """ + + def __init__( + self, + content: AsyncGenerator[str, None], + *, + media_type: Optional[str] = None, + headers: Optional[dict] = None, + status_code: int = status.HTTP_200_OK, + upstream_generator: Optional[AsyncGenerator[str, None]] = None, + ) -> None: + super().__init__( + content, status_code=status_code, headers=headers, media_type=media_type + ) + self._upstream_generator = upstream_generator + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + try: + await super().__call__(scope, receive, send) + finally: + with anyio.CancelScope(shield=True): + for target in (self.body_iterator, self._upstream_generator): + aclose = getattr(target, "aclose", None) + if aclose is None: + continue + try: + await aclose() + except BaseException as e: + verbose_proxy_logger.debug( + "error closing streaming generator: %s", e + ) + + +async def create_response( generator: AsyncGenerator[str, None], media_type: str, headers: dict, @@ -366,11 +426,12 @@ async def create_response( # noqa: PLR0915 with tracer.trace(DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE): yield chunk - return StreamingResponse( + return _UpstreamClosingStreamingResponse( combined_generator(), media_type=media_type, headers=streaming_headers, status_code=final_status_code, + upstream_generator=generator, ) @@ -1087,7 +1148,7 @@ class ProxyBaseLLMRequestProcessing: _payload_str, ) - async def base_process_llm_request( # noqa: PLR0915 + async def base_process_llm_request( self, request: Request, fastapi_response: Response, @@ -1702,6 +1763,23 @@ class ProxyBaseLLMRequestProcessing: response=completed_obj, user_api_key_dict=user_api_key_dict, ) + else: + # Silent skip caused #30210: the proxy's Router wrapper + # of the responses streaming iterator wasn't propagating + # ``completed_response``, so this hook recorded nothing + # and follow-up /v1/containers//files calls 403'd + # for non-admin keys with no proxy-side hint. Log a + # warning so future regressions of the same shape + # surface in operator logs. + verbose_proxy_logger.warning( + "Container ownership recording skipped on streaming " + "/v1/responses: no completed_response on stream " + "iterator %s. If this stream created any tool " + "container (e.g. code_interpreter), follow-up " + "/v1/containers//files calls will 403 for " + "non-admin keys.", + type(original_stream_response).__name__, + ) except Exception as e: verbose_proxy_logger.exception( "Container ownership recording failed after streaming responses call: %s", @@ -2424,6 +2502,8 @@ class ProxyBaseLLMRequestProcessing: code=getattr(e, "status_code", 500), ) yield serialize_error(proxy_exception) + finally: + await _aclose_upstream_response(response) @staticmethod def async_sse_data_generator( diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index c630294c1ec..71dce163b78 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -4,6 +4,7 @@ from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Literal, import litellm from litellm import get_secret from litellm._logging import verbose_proxy_logger +from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.proxy._types import CommonProxyErrors, LiteLLMPromptInjectionParams @@ -35,7 +36,7 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging -def initialize_callbacks_on_proxy( # noqa: PLR0915 +def initialize_callbacks_on_proxy( value: Any, premium_user: bool, config_file_path: str, @@ -497,6 +498,7 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset( "guardrail_config", "_guardrail_pipelines", "_pipeline_managed_guardrails", + PRE_CALL_EXECUTED_GUARDRAILS_KEY, "disable_global_guardrails", "disable_global_guardrail", "opted_out_global_guardrails", diff --git a/litellm/proxy/common_utils/debug_utils.py b/litellm/proxy/common_utils/debug_utils.py index 9b2c3ddce46..4cc62e1adbd 100644 --- a/litellm/proxy/common_utils/debug_utils.py +++ b/litellm/proxy/common_utils/debug_utils.py @@ -92,12 +92,12 @@ if os.environ.get("LITELLM_PROFILE", "false").lower() == "true": try: import objgraph # type: ignore - print("growth of objects") # noqa + print("growth of objects") # noqa: T201 objgraph.show_growth() - print("\n\nMost common types") # noqa + print("\n\nMost common types") # noqa: T201 objgraph.show_most_common_types() roots = objgraph.get_leaking_objects() - print("\n\nLeaking objects") # noqa + print("\n\nLeaking objects") # noqa: T201 objgraph.show_most_common_types(objects=roots) except ImportError: raise ImportError( @@ -739,7 +739,7 @@ async def get_otel_spans(): else: recorded_spans = [] - print("Spans: ", recorded_spans) # noqa + print("Spans: ", recorded_spans) # noqa: T201 most_recent_parent = None most_recent_start_time = 1000000 diff --git a/litellm/proxy/common_utils/timezone_utils.py b/litellm/proxy/common_utils/timezone_utils.py index 700a9197f6f..32f9f47d519 100644 --- a/litellm/proxy/common_utils/timezone_utils.py +++ b/litellm/proxy/common_utils/timezone_utils.py @@ -15,7 +15,7 @@ def get_budget_reset_timezone(): return getattr(litellm, "timezone", None) or "UTC" -def get_budget_reset_time(budget_duration: str): +def get_budget_reset_time(budget_duration: str) -> datetime: """ Get the budget reset time based on the configured timezone. Falls back to UTC if not specified. diff --git a/litellm/proxy/db/check_migration.py b/litellm/proxy/db/check_migration.py index bf180c1132d..2aacaed8aff 100644 --- a/litellm/proxy/db/check_migration.py +++ b/litellm/proxy/db/check_migration.py @@ -54,7 +54,7 @@ def check_prisma_schema_diff_helper(db_url: str) -> Tuple[bool, List[str]]: subprocess.CalledProcessError: If the Prisma command fails. Exception: For any other errors during execution. """ - verbose_logger.debug("Checking for Prisma schema diff...") # noqa: T201 + verbose_logger.debug("Checking for Prisma schema diff...") try: result = subprocess.run( [ diff --git a/litellm/proxy/db/create_views.py b/litellm/proxy/db/create_views.py index 97525a528d0..d9e21fc5d2a 100644 --- a/litellm/proxy/db/create_views.py +++ b/litellm/proxy/db/create_views.py @@ -11,7 +11,7 @@ _db = Any _VIEW_NOT_FOUND_MARKERS = ("does not exist", "no such table", "undefined table") -async def create_missing_views(db: _db): # noqa: PLR0915 +async def create_missing_views(db: _db): """ -------------------------------------------------- NOTE: Copy of `litellm/db_scripts/create_views.py`. diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index e7f14df5294..4b7b20d75d0 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1128,7 +1128,7 @@ class DBSpendUpdateWriter: "_flush_tool_discovery_queue error (non-blocking): %s", e ) - async def _commit_spend_updates_to_db( # noqa: PLR0915 + async def _commit_spend_updates_to_db( self, prisma_client: PrismaClient, n_retry_times: int, diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py index ff802223f21..72b9b7dc3c1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py @@ -121,7 +121,7 @@ class lakeraAI_Moderation(CustomGuardrail): return None - async def _check( # noqa: PLR0915 + async def _check( self, data: dict, user_api_key_dict: UserAPIKeyAuth, diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index e5200394b55..e8887fa712a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -261,7 +261,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): ) return "" - async def _call_panw_api( # noqa: PLR0915 + async def _call_panw_api( self, content: str = "", is_response: bool = False, @@ -1762,7 +1762,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): return rd.get("name") if ("arguments" in rd or "mcp_arguments" in rd) else None @log_guardrail_information - async def apply_guardrail( # noqa: PLR0915 + async def apply_guardrail( self, inputs: GenericGuardrailAPIInputs, request_data: dict, diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index e723c07e3c4..7d6d1adb05e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -28,7 +28,7 @@ from typing import ( import aiohttp -import litellm # noqa: E401 +import litellm from litellm import get_secret from litellm._logging import verbose_proxy_logger from litellm.types.utils import GenericGuardrailAPIInputs @@ -1432,7 +1432,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): try: verbose_proxy_logger.debug(print_statement) if litellm.set_verbose: - print(print_statement) # noqa + print(print_statement) # noqa: T201 except Exception: pass diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 2a2c758fa8a..09fff71062b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -288,7 +288,7 @@ class UnifiedLLMGuardrails(CustomLogger): return response - async def async_post_call_streaming_iterator_hook( # noqa: PLR0915 + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, response: Any, diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index a80bb817890..b99ea8f14a0 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -5,6 +5,8 @@ import os from datetime import datetime, timezone from typing import Any, Dict, List, Literal, Optional, Set, Type, cast +from pydantic import ValidationError + import litellm from litellm import Router from litellm._logging import verbose_proxy_logger @@ -601,21 +603,25 @@ class InMemoryGuardrailHandler: def delete_in_memory_guardrail(self, guardrail_id: str) -> None: """ Delete a guardrail in memory and remove from litellm callbacks. + + The callback is purged from every callback list, not just + litellm.callbacks: request handling promotes guardrail callbacks into the + success/failure/async lists, so removing it from only litellm.callbacks + leaves the old instance stranded in those lists on every re-initialization. """ # Remove from in-memory storage self.IN_MEMORY_GUARDRAILS.pop(guardrail_id, None) self._sources.pop(guardrail_id, None) - # Remove the callback from litellm.callbacks custom_guardrail_callback = self.guardrail_id_to_custom_guardrail.pop( guardrail_id, None ) - if custom_guardrail_callback: - litellm.logging_callback_manager.remove_callback_from_list_by_object( - callback_list=litellm.callbacks, - obj=custom_guardrail_callback, - require_self=False, - ) + if custom_guardrail_callback is None: + return + + litellm.logging_callback_manager.remove_callback_from_all_lists( + custom_guardrail_callback + ) def list_in_memory_guardrails(self) -> List[Guardrail]: """ @@ -657,6 +663,34 @@ class InMemoryGuardrailHandler: self.delete_in_memory_guardrail(guardrail_id) return stale_ids + @staticmethod + def _normalize_litellm_params_for_comparison( + params: Optional[Any], + ) -> Optional[Dict[str, Any]]: + """ + Render litellm_params to a canonical dict so an in-memory LitellmParams and + the raw dict loaded from the DB compare equal when they describe the same + config. The in-memory side is a LitellmParams whose model_dump() carries + every field default and coerces enums, while the DB side is the raw stored + dict holding only the keys originally provided. Comparing those two shapes + directly never matches, so each DB poll would re-initialize the guardrail + forever; normalizing both through LitellmParams keeps the diff meaningful. + """ + if params is None: + return None + if isinstance(params, LitellmParams): + return params.model_dump() + if isinstance(params, dict): + try: + return LitellmParams(**params).model_dump() + except ValidationError as e: + verbose_proxy_logger.warning( + f"Could not normalize guardrail litellm_params for comparison; " + f"treating the guardrail as changed. Error: {e}" + ) + return params + return params + def _has_guardrail_params_changed( self, guardrail_id: str, new_guardrail: Guardrail ) -> bool: @@ -673,19 +707,11 @@ class InMemoryGuardrailHandler: return True # Compare litellm_params - existing_params = existing.get("litellm_params") - new_params = new_guardrail.get("litellm_params") - - # Convert to dicts for comparison - existing_dict = ( - existing_params.model_dump() - if isinstance(existing_params, LitellmParams) - else existing_params + existing_dict = self._normalize_litellm_params_for_comparison( + existing.get("litellm_params") ) - new_dict = ( - new_params.model_dump() - if isinstance(new_params, LitellmParams) - else new_params + new_dict = self._normalize_litellm_params_for_comparison( + new_guardrail.get("litellm_params") ) # Compare and identify specific differences diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index e0d018d4344..8a432eb2f42 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -167,7 +167,7 @@ async def test_endpoint(request: Request): tags=["health"], dependencies=[Depends(user_api_key_auth)], ) -async def health_services_endpoint( # noqa: PLR0915 +async def health_services_endpoint( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), service: services = fastapi.Query(description="Specify the service being hit."), ): diff --git a/litellm/proxy/hooks/batch_redis_get.py b/litellm/proxy/hooks/batch_redis_get.py index c608317f4eb..f734b19681d 100644 --- a/litellm/proxy/hooks/batch_redis_get.py +++ b/litellm/proxy/hooks/batch_redis_get.py @@ -33,7 +33,7 @@ class _PROXY_BatchRedisRequests(CustomLogger): elif debug_level == "INFO": verbose_proxy_logger.debug(print_statement) if litellm.set_verbose is True: - print(print_statement) # noqa + print(print_statement) # noqa: T201 async def async_pre_call_hook( self, diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 874e5aa1939..d36e9858b5a 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -50,7 +50,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): try: verbose_proxy_logger.debug(print_statement) if litellm.set_verbose: - print(print_statement) # noqa + print(print_statement) # noqa: T201 except Exception: pass @@ -239,7 +239,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): request_count_end_user_id=results[5], ) - async def async_pre_call_hook( # noqa: PLR0915 + async def async_pre_call_hook( self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, @@ -506,9 +506,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): return - async def async_log_success_event( # noqa: PLR0915 - self, kwargs, response_obj, start_time, end_time - ): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, ) @@ -769,7 +767,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): litellm_parent_otel_span=litellm_parent_otel_span, ) except Exception as e: - self.print_verbose(e) # noqa + self.print_verbose(e) async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: diff --git a/litellm/proxy/hooks/prompt_injection_detection.py b/litellm/proxy/hooks/prompt_injection_detection.py index f1b688948f2..6678ccd7e0b 100644 --- a/litellm/proxy/hooks/prompt_injection_detection.py +++ b/litellm/proxy/hooks/prompt_injection_detection.py @@ -72,7 +72,7 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): verbose_proxy_logger.debug(print_statement) if litellm.set_verbose is True: - print(print_statement) # noqa + print(print_statement) # noqa: T201 def update_environment(self, router: Optional[Router] = None): self.llm_router = router diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index fca395f889c..0e21fd8e1f0 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -13,6 +13,7 @@ from starlette.datastructures import Headers import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging +from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host @@ -161,6 +162,7 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS = ( "secret_fields", "_guardrail_pipelines", "_pipeline_managed_guardrails", + PRE_CALL_EXECUTED_GUARDRAILS_KEY, ) _UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS = frozenset( @@ -1315,7 +1317,7 @@ class LiteLLMProxyRequestSetup: ) -async def add_litellm_data_to_request( # noqa: PLR0915 +async def add_litellm_data_to_request( data: dict, request: Request, user_api_key_dict: UserAPIKeyAuth, diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index 698155a5c26..e35ec2933d0 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -183,10 +183,23 @@ async def update_budget( except ValueError as e: raise HTTPException(status_code=400, detail={"error": str(e)}) + # recompute budget_reset_at when the duration changes, unless the caller pinned a reset time explicitly + recomputed_reset_at = ( + { + "budget_reset_at": get_budget_reset_time( + budget_duration=budget_obj.budget_duration + ) + } + if budget_obj.budget_duration is not None + and "budget_reset_at" not in budget_obj.model_fields_set + else {} + ) + response = await BudgetRepository(prisma_client).table.update( where={"budget_id": budget_obj.budget_id}, data={ **budget_obj.model_dump(exclude_unset=True), # type: ignore + **recomputed_reset_at, "updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name, }, # type: ignore ) diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 13107b68864..e258ddc0410 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -726,7 +726,7 @@ def _key_metadata( return KeyMetadata(key_alias=meta.get("key_alias"), team_id=meta.get("team_id")) -def _aggregate_grouping_sets_records_sync( # noqa: PLR0915 +def _aggregate_grouping_sets_records_sync( *, records: List[Any], api_key_metadata: Dict[str, Dict[str, Any]], diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index f1a34bb0ed4..50a1bc23a6d 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -885,6 +885,19 @@ async def get_customer_daily_activity( """ Get daily activity for specific organizations or all accessible organizations. """ + if ( + user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN + and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY + ): + raise HTTPException( + status_code=401, + detail={ + "error": "Admin-only endpoint. Your user role={}".format( + user_api_key_dict.user_role + ) + }, + ) + from litellm.proxy.proxy_server import prisma_client if prisma_client is None: diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index b3a5c66e9e1..ba7013570fe 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -755,7 +755,7 @@ def _build_user_info_response( response_model=UserInfoResponse, ) @management_endpoint_wrapper -async def user_info( # noqa: PLR0915 +async def user_info( request: Request, user_id: Optional[str] = fastapi.Query( default=None, description="User ID in the request parameters" @@ -1082,7 +1082,7 @@ def _process_keys_for_user_info( continue try: - _key: dict = key.model_dump() # noqa + _key: dict = key.model_dump() except Exception: # if using pydantic v1 _key = key.dict() diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index c980f6f5260..132060be76b 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -675,7 +675,7 @@ def _enforce_upperbound_key_params( ) -async def _common_key_generation_helper( # noqa: PLR0915 +async def _common_key_generation_helper( data: GenerateKeyRequest, user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str], @@ -2428,7 +2428,7 @@ async def _validate_update_key_data( "/key/update", tags=["key management"], dependencies=[Depends(user_api_key_auth)] ) @management_endpoint_wrapper -async def update_key_fn( # noqa: PLR0915 +async def update_key_fn( request: Request, data: UpdateKeyRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -3377,7 +3377,7 @@ async def info_key_fn( ) ## REMOVE HASHED TOKEN INFO BEFORE RETURNING ## try: - key_info = key_info.model_dump() # noqa + key_info = key_info.model_dump() except Exception: # if using pydantic v1 key_info = key_info.dict() @@ -3419,7 +3419,7 @@ def _check_model_access_group( return True -async def generate_key_helper_fn( # noqa: PLR0915 +async def generate_key_helper_fn( request_type: Literal[ "user", "key" ], # identifies if this request is from /user/new or /key/generate @@ -4070,7 +4070,7 @@ async def delete_key_aliases( ) -async def _rotate_master_key( # noqa: PLR0915 +async def _rotate_master_key( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, current_master_key: str, @@ -4412,7 +4412,7 @@ async def _execute_virtual_key_regeneration( dependencies=[Depends(user_api_key_auth)], ) @management_endpoint_wrapper -async def regenerate_key_fn( # noqa: PLR0915 +async def regenerate_key_fn( key: Optional[str] = None, data: Optional[RegenerateKeyRequest] = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), diff --git a/litellm/proxy/management_endpoints/policy_endpoints/__init__.py b/litellm/proxy/management_endpoints/policy_endpoints/__init__.py index 157d15c2710..862c92bace9 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints/__init__.py +++ b/litellm/proxy/management_endpoints/policy_endpoints/__init__.py @@ -7,7 +7,7 @@ continue to work. Patch targets also resolve correctly since names are imported directly into this namespace. """ -from litellm.proxy.management_endpoints.policy_endpoints.endpoints import * # noqa: F401, F403 +from litellm.proxy.management_endpoints.policy_endpoints.endpoints import * # noqa: F403 from litellm.proxy.management_endpoints.policy_endpoints.endpoints import ( # noqa: F401 _build_all_names_per_competitor, _build_comparison_blocked_words, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 1a0a57c71fd..4d4d1ef2774 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -933,7 +933,7 @@ def _check_team_budget_update_authority( response_model=LiteLLM_TeamTable, ) @management_endpoint_wrapper -async def new_team( # noqa: PLR0915 +async def new_team( data: NewTeamRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -1637,7 +1637,7 @@ def validate_team_org_change( "/team/update", tags=["team management"], dependencies=[Depends(user_api_key_auth)] ) @management_endpoint_wrapper -async def update_team( # noqa: PLR0915 +async def update_team( data: UpdateTeamRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -3647,7 +3647,7 @@ async def team_info( ## REMOVE HASHED TOKEN INFO before returning ## for key in keys: try: - key = key.model_dump() # noqa + key = key.model_dump() except Exception: # if using pydantic v1 key = key.dict() @@ -4244,6 +4244,14 @@ async def _enforce_list_team_v2_access( status_code=403, detail={"error": "You can only view teams within your organizations."}, ) + # When the caller is an org admin querying their own teams (or no + # specific user), null out user_id so that + # _build_team_list_where_conditions scopes only by organization_id + # — org admins should see all teams in their orgs, not just teams + # they are a direct member of. Keep user_id when the org admin + # explicitly queries a *different* user's teams. + if user_id is None or user_id == user_api_key_dict.user_id: + user_id = None verbose_proxy_logger.debug( "list_team_v2: org admin access for user=%s, org_ids=%s, user_id_filter=%s", user_api_key_dict.user_id, diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 4812bed2f21..91a5c109acf 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -114,7 +114,7 @@ from litellm.repositories.table_repositories import SSOConfigRepository from litellm.repositories.team_repository import TeamRepository from litellm.repositories.user_repository import UserRepository from litellm.secret_managers.main import get_secret_bool, str_to_bool -from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403, F401 +from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403 from litellm.types.proxy.management_endpoints.ui_sso import ( DefaultTeamSSOParams, MicrosoftGraphAPIUserGroupDirectoryObject, @@ -145,6 +145,9 @@ _CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS = 60 _CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS = 30 _CLI_SSO_USER_CODE_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" _CLI_SSO_LOGIN_ID_RE = re.compile(r"^cli-[A-Za-z0-9_-]{12,124}$") +_CLI_SSO_USER_CODE_RE = re.compile( + rf"^[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}-[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}$" +) _CLI_SSO_SCALAR_TYPES = (str, int, float, bool) _CLI_SSO_DEST_KEY_RE = re.compile(r"^[A-Za-z0-9_.-]+$") _CLI_SSO_SECRET_KEY_FRAGMENTS = frozenset( @@ -182,6 +185,45 @@ def _is_valid_cli_sso_login_id(login_id: Optional[str]) -> bool: return isinstance(login_id, str) and bool(_CLI_SSO_LOGIN_ID_RE.fullmatch(login_id)) +def _is_valid_cli_sso_user_code(user_code: str | None) -> bool: + return isinstance(user_code, str) and bool( + _CLI_SSO_USER_CODE_RE.fullmatch(user_code) + ) + + +def _cli_sso_verification_uri_complete_enabled() -> bool: + from litellm.proxy.proxy_server import general_settings + + return bool( + general_settings.get( # any-ok: operator opt-in read from the untyped general_settings dict + "allow_cli_sso_verification_uri_complete", False + ) + ) + + +def _cli_sso_start_response_body( + *, + login_id: str, + poll_secret: str, + user_code: str, + verification_uri_complete: str | None, +) -> dict[str, str | int]: + if verification_uri_complete is None: + return { + "login_id": login_id, + "poll_secret": poll_secret, + "user_code": user_code, + "expires_in": CLI_SSO_SESSION_TTL_SECONDS, + } + return { + "login_id": login_id, + "poll_secret": poll_secret, + "user_code": user_code, + "verification_uri_complete": verification_uri_complete, + "expires_in": CLI_SSO_SESSION_TTL_SECONDS, + } + + def _get_cli_sso_start_rate_limit_cache_key( request: Request, use_x_forwarded_for: Optional[bool] = False ) -> str: @@ -478,10 +520,20 @@ def _cli_poll_attribution_metadata_from_session( def _render_cli_sso_verification_page( - verify_url: str, browser_complete_token: str + verify_url: str, + browser_complete_token: str, + prefill_user_code: str | None = None, ) -> str: escaped_verify_url = escape(verify_url, quote=True) escaped_browser_complete_token = escape(browser_complete_token, quote=True) + user_code_value_attr = ( + f' value="{escape(prefill_user_code, quote=True)}"' if prefill_user_code else "" + ) + instructions = ( + "Confirm the verification code below to finish this login." + if prefill_user_code + else "Enter the verification code shown in your terminal to finish this login." + ) return f""" @@ -535,11 +587,11 @@ def _render_cli_sso_verification_page(

Complete CLI Login

-

Enter the verification code shown in your terminal to finish this login.

+

{instructions}

- +
@@ -573,12 +625,29 @@ async def cli_sso_start(request: Request): } _set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow) - return { - "login_id": login_id, - "poll_secret": poll_secret, - "user_code": user_code, - "expires_in": CLI_SSO_SESSION_TTL_SECONDS, - } + verification_uri_complete: str | None = ( + ( + get_custom_url( + request_base_url=str(request.base_url), route="sso/key/generate" + ) + + "?" + + urlencode( + { + "source": LITELLM_CLI_SOURCE_IDENTIFIER, + "key": login_id, + "user_code": user_code, + } + ) + ) + if _cli_sso_verification_uri_complete_enabled() + else None + ) + return _cli_sso_start_response_body( + login_id=login_id, + poll_secret=poll_secret, + user_code=user_code, + verification_uri_complete=verification_uri_complete, + ) @router.post( @@ -829,7 +898,8 @@ async def google_login( key: Optional[str] = None, existing_key: Optional[str] = None, return_to: Optional[str] = None, -): # noqa: PLR0915 + user_code: str | None = None, +): """ Create Proxy API Keys using Google Workspace SSO. Requires setting PROXY_BASE_URL in .env PROXY_BASE_URL should be the your deployed proxy endpoint, e.g. PROXY_BASE_URL="https://litellm-production-7002.up.railway.app/" @@ -897,6 +967,7 @@ async def google_login( cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state( source=source, key=key, + user_code=(user_code if _cli_sso_verification_uri_complete_enabled() else None), ) # check if user defined a custom auth sso sign in handler, if yes, use it @@ -1833,7 +1904,7 @@ async def check_and_update_if_proxy_admin_id( @router.get("/sso/callback", tags=["experimental"], include_in_schema=False) -async def auth_callback(request: Request, state: Optional[str] = None): # noqa: PLR0915 +async def auth_callback(request: Request, state: Optional[str] = None): """Verify login""" verbose_proxy_logger.info(f"Starting SSO callback with state: {state}") @@ -1921,14 +1992,16 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa: ) if state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"): - # State format: {PREFIX}:{login_id} - state_parts = state.split(":", 1) + # State format: {PREFIX}:{login_id}[:{user_code}] + state_parts = state.split(":", 2) key_id = state_parts[1] if len(state_parts) > 1 else None + prefill_user_code = state_parts[2] if len(state_parts) > 2 else None verbose_proxy_logger.info("CLI SSO callback detected") return await cli_sso_callback( request=request, key=key_id, + prefill_user_code=prefill_user_code, result=result, received_response=received_response, ) @@ -2008,6 +2081,7 @@ async def _complete_cli_sso_callback_session( prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, + prefill_user_code: str | None = None, ): from fastapi.responses import HTMLResponse @@ -2071,6 +2145,7 @@ async def _complete_cli_sso_callback_session( content=_render_cli_sso_verification_page( verify_url=verify_url, browser_complete_token=browser_complete_token, + prefill_user_code=prefill_user_code, ), status_code=200, ) @@ -2081,6 +2156,7 @@ async def cli_sso_callback( key: Optional[str] = None, result: Optional[Union[OpenID, dict]] = None, received_response: Optional[dict] = None, + prefill_user_code: str | None = None, ): """CLI SSO callback - stores session info for JWT generation on polling""" verbose_proxy_logger.info("CLI SSO callback") @@ -2137,6 +2213,7 @@ async def cli_sso_callback( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + prefill_user_code=prefill_user_code, ) except ProxyException: raise @@ -3053,21 +3130,27 @@ class SSOAuthenticationHandler: @staticmethod def _get_cli_state( - source: Optional[str], key: Optional[str], existing_key: Optional[str] = None + source: str | None, + key: str | None, + existing_key: str | None = None, + user_code: str | None = None, ) -> Optional[str]: """ Checks the request 'source' if a cli state token was passed in This is used to authenticate through the CLI login flow. - The state parameter format is: {PREFIX}:{login_id} + The state parameter format is: {PREFIX}:{login_id}[:{user_code}] - The state parameter is used to pass data through the OAuth flow without changing the callback URL + - user_code is appended only for the opt-in verification_uri_complete flow so the verify page can pre-fill it """ from litellm.constants import ( LITELLM_CLI_SESSION_TOKEN_PREFIX, ) if source == LITELLM_CLI_SOURCE_IDENTIFIER and key: + if _is_valid_cli_sso_user_code(user_code): + return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}:{user_code}" return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}" else: return None @@ -3145,7 +3228,7 @@ class SSOAuthenticationHandler: ) @staticmethod - async def get_redirect_response_from_openid( # noqa: PLR0915 + async def get_redirect_response_from_openid( result: Union[OpenID, dict, CustomOpenID], request: Request, received_response: Optional[dict] = None, diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 3e5873c2655..f43e876d111 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -284,7 +284,7 @@ async def route_create_file( dependencies=[Depends(user_api_key_auth)], tags=["files"], ) -async def create_file( # noqa: PLR0915 +async def create_file( request: Request, fastapi_response: Response, purpose: str = Form(...), @@ -589,7 +589,7 @@ async def create_file( # noqa: PLR0915 dependencies=[Depends(user_api_key_auth)], tags=["files"], ) -async def get_file_content( # noqa: PLR0915 +async def get_file_content( request: Request, fastapi_response: Response, file_id: str, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py index a912a88a993..6feb4e36bf9 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py @@ -366,7 +366,7 @@ class AnthropicPassthroughLoggingHandler: ) @staticmethod - def _collapse_pure_text_chunks( # noqa: PLR0915 + def _collapse_pure_text_chunks( all_chunks: Sequence[Union[str, bytes]], ) -> Optional[List[str]]: """ @@ -551,7 +551,7 @@ class AnthropicPassthroughLoggingHandler: return complete_streaming_response @staticmethod - def batch_creation_handler( # noqa: PLR0915 + def batch_creation_handler( httpx_response: httpx.Response, logging_obj: LiteLLMLoggingObj, url_route: str, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py index adb1278fee5..0875f1d5508 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py @@ -71,7 +71,7 @@ class CoherePassthroughLoggingHandler(BasePassthroughLoggingHandler): complete_streaming_response = stream_chunk_builder(chunks=all_openai_chunks) return complete_streaming_response - def cohere_passthrough_handler( # noqa: PLR0915 + def cohere_passthrough_handler( self, httpx_response: httpx.Response, response_body: dict, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py index 9f353226dd0..b77c6e2f655 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py @@ -275,7 +275,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): return litellm_model_response, response_cost @staticmethod - def openai_passthrough_handler( # noqa: PLR0915 + def openai_passthrough_handler( httpx_response: httpx.Response, response_body: dict, logging_obj: LiteLLMLoggingObj, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 6a138532617..73d4245670a 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -645,7 +645,7 @@ class VertexPassthroughLoggingHandler: return kwargs @staticmethod - def batch_prediction_jobs_handler( # noqa: PLR0915 + def batch_prediction_jobs_handler( httpx_response: httpx.Response, logging_obj: LiteLLMLoggingObj, url_route: str, diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e3cb9dec884..b84746758fb 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -143,7 +143,7 @@ async def set_env_variables_in_header(custom_headers: Optional[dict]) -> Optiona return headers -async def chat_completion_pass_through_endpoint( # noqa: PLR0915 +async def chat_completion_pass_through_endpoint( fastapi_response: Response, request: Request, adapter_id: str, @@ -701,7 +701,7 @@ from litellm.passthrough.timeout_utils import ( ) -async def pass_through_request( # noqa: PLR0915 +async def pass_through_request( request: Request, target: str, custom_headers: dict, @@ -1540,7 +1540,7 @@ async def _parse_request_data_by_content_type( return query_params_data, custom_body_data, file_data, stream -def create_pass_through_route( # noqa: PLR0915 +def create_pass_through_route( endpoint, target: str, custom_headers: Optional[Mapping[str, Any]] = None, @@ -1776,7 +1776,7 @@ def create_websocket_passthrough_route( return websocket_endpoint_func -async def websocket_passthrough_request( # noqa: PLR0915 +async def websocket_passthrough_request( websocket: WebSocket, target: str, custom_headers: dict, diff --git a/litellm/proxy/policy_engine/pipeline_executor.py b/litellm/proxy/policy_engine/pipeline_executor.py index 3c5a1d67be4..e46e3e1dc9f 100644 --- a/litellm/proxy/policy_engine/pipeline_executor.py +++ b/litellm/proxy/policy_engine/pipeline_executor.py @@ -171,6 +171,10 @@ class PipelineExecutor: data=data, call_type=call_type, # type: ignore ) + if isinstance(callback, CustomGuardrail): + callback.mark_pre_call_hook_ran(data) + if isinstance(response, dict): + callback.mark_pre_call_hook_ran(response) elif mode == "post_call": response = await target.async_post_call_success_hook( user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/post_call_rules.py b/litellm/proxy/post_call_rules.py index 23ec93f5b30..6200bee4d7b 100644 --- a/litellm/proxy/post_call_rules.py +++ b/litellm/proxy/post_call_rules.py @@ -1,5 +1,5 @@ def post_response_rule(input): # receives the model response - print(f"post_response_rule:input={input}") # noqa + print(f"post_response_rule:input={input}") # noqa: T201 if len(input) < 200: return { "decision": False, diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 8c3fa952903..e1fb65074cd 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -102,9 +102,9 @@ class ProxyInitializationHelpers: @staticmethod def _run_health_check(host, port): - print("\nLiteLLM: Health Testing models in config") # noqa + print("\nLiteLLM: Health Testing models in config") response = httpx.get(url=f"http://{host}:{port}/health") - print(json.dumps(response.json(), indent=4)) # noqa + print(json.dumps(response.json(), indent=4)) @staticmethod def _run_test_chat_completion( @@ -138,7 +138,7 @@ class ProxyInitializationHelpers: ) click.echo(f"\nLiteLLM: response from proxy {response}") - print( # noqa + print( f"\n LiteLLM: Making a test ChatCompletions + streaming r equest to proxy. Model={request_model}" ) @@ -154,11 +154,11 @@ class ProxyInitializationHelpers: ) for chunk in stream_response: click.echo(f"LiteLLM: streaming response from proxy {chunk}") - print("\n making completion request to proxy") # noqa + print("\n making completion request to proxy") completion_response = client.completions.create( model=request_model, prompt="this is a test request, write a short poem" ) - print(completion_response) # noqa + print(completion_response) @staticmethod def _get_default_unvicorn_init_args( @@ -184,7 +184,7 @@ class ProxyInitializationHelpers: "port": port, } if log_config is not None: - print(f"Using log_config: {log_config}") # noqa + print(f"Using log_config: {log_config}") uvicorn_args["log_config"] = log_config elif litellm.json_logs: # Use JSON log config for uvicorn to ensure all logs (including exceptions) are JSON @@ -198,7 +198,7 @@ class ProxyInitializationHelpers: ): uvicorn_args["timeout_worker_healthcheck"] = timeout_worker_healthcheck else: - print( # noqa + print( f"\033[1;33mLiteLLM Proxy: --timeout_worker_healthcheck " f"requires uvicorn>=0.37.0, but installed uvicorn=={uvicorn.__version__}. " f"Ignoring the flag.\033[0m" @@ -304,15 +304,15 @@ class ProxyInitializationHelpers: from hypercorn.asyncio import serve from hypercorn.config import Config - print( # noqa - f"\033[1;32mLiteLLM Proxy: Starting server on {host}:{port} using Hypercorn\033[0m\n" # noqa - ) # noqa + print( + f"\033[1;32mLiteLLM Proxy: Starting server on {host}:{port} using Hypercorn\033[0m\n" + ) config = Config() config.bind = [f"{host}:{port}"] if ssl_certfile_path is not None and ssl_keyfile_path is not None: - print( # noqa - f"\033[1;32mLiteLLM Proxy: Using SSL with certfile: {ssl_certfile_path} and keyfile: {ssl_keyfile_path}\033[0m\n" # noqa + print( + f"\033[1;32mLiteLLM Proxy: Using SSL with certfile: {ssl_certfile_path} and keyfile: {ssl_keyfile_path}\033[0m\n" ) config.certfile = ssl_certfile_path config.keyfile = ssl_keyfile_path @@ -342,16 +342,16 @@ class ProxyInitializationHelpers: from granian import Granian from granian.constants import Interfaces - print( # noqa + print( f"\033[1;32mLiteLLM Proxy: Starting server on {host}:{port} using Granian\033[0m\n" ) if max_requests_before_restart is not None: - print( # noqa + print( "\033[1;33mLiteLLM: --max_requests_before_restart is not supported by Granian " "(Granian uses workers_lifetime in seconds, not a per-request limit).\033[0m\n" ) if ciphers is not None: - print( # noqa + print( "\033[1;33mLiteLLM: --ciphers is not applied when using --run_granian.\033[0m\n" ) @@ -366,7 +366,7 @@ class ProxyInitializationHelpers: if granian_runtime_threads is not None: kwargs["runtime_threads"] = granian_runtime_threads if ssl_certfile_path is not None and ssl_keyfile_path is not None: - print( # noqa + print( f"\033[1;32mLiteLLM Proxy: Using SSL with certfile: {ssl_certfile_path} and keyfile: {ssl_keyfile_path}\033[0m\n" ) kwargs["ssl_cert"] = Path(ssl_certfile_path) @@ -419,19 +419,19 @@ class ProxyInitializationHelpers: }' \n """ - print() # noqa - print( # noqa + print() + print( '\033[1;34mLiteLLM: Test your local proxy with: "litellm --test" This runs an openai.ChatCompletion request to your proxy [In a new terminal tab]\033[0m\n' ) - print( # noqa + print( f"\033[1;34mLiteLLM: Curl Command Test for your local proxy\n {curl_command} \033[0m\n" ) - print( # noqa + print( "\033[1;34mDocs: https://docs.litellm.ai/docs/simple_proxy\033[0m\n" - ) # noqa - print( # noqa + ) + print( f"\033[1;34mSee all Router/Swagger docs on http://0.0.0.0:{port} \033[0m\n" - ) # noqa + ) def load_config(self): # note: This Loads the gunicorn config - has nothing to do with LiteLLM Proxy config @@ -451,8 +451,8 @@ class ProxyInitializationHelpers: # gunicorn app function return self.application - print( # noqa - f"\033[1;32mLiteLLM Proxy: Starting server on {host}:{port} with {num_workers} workers\033[0m\n" # noqa + print( + f"\033[1;32mLiteLLM Proxy: Starting server on {host}:{port} with {num_workers} workers\033[0m\n" ) gunicorn_options = { "bind": f"{host}:{port}", @@ -478,8 +478,8 @@ class ProxyInitializationHelpers: gunicorn_options["child_exit"] = child_exit if ssl_certfile_path is not None and ssl_keyfile_path is not None: - print( # noqa - f"\033[1;32mLiteLLM Proxy: Using SSL with certfile: {ssl_certfile_path} and keyfile: {ssl_keyfile_path}\033[0m\n" # noqa + print( + f"\033[1;32mLiteLLM Proxy: Using SSL with certfile: {ssl_certfile_path} and keyfile: {ssl_keyfile_path}\033[0m\n" ) gunicorn_options["certfile"] = ssl_certfile_path gunicorn_options["keyfile"] = ssl_keyfile_path @@ -496,7 +496,7 @@ class ProxyInitializationHelpers: except Exception as e: print(f""" LiteLLM Warning: proxy started with `ollama` model\n`ollama serve` failed with Exception{e}. \nEnsure you run `ollama serve` - """) # noqa # noqa + """) @staticmethod def _is_port_in_use(port): @@ -557,7 +557,7 @@ class ProxyInitializationHelpers: os.makedirs(multiproc_dir, exist_ok=True) wipe_directory(multiproc_dir) action = "Auto-created" if auto_created else "Using existing" - print(f"LiteLLM: {action} PROMETHEUS_MULTIPROC_DIR={multiproc_dir}") # noqa + print(f"LiteLLM: {action} PROMETHEUS_MULTIPROC_DIR={multiproc_dir}") @click.command() @@ -814,7 +814,7 @@ class ProxyInitializationHelpers: default=False, help="Enable uvicorn hot reload (dev only). Also reloads when the --config YAML file changes. Incompatible with --num_workers>1, --run_gunicorn, and --run_hypercorn.", ) -def run_server( # noqa: PLR0915 +def run_server( cli_args, host, port, @@ -1185,7 +1185,7 @@ def run_server( # noqa: PLR0915 check_prisma_schema_diff(db_url=None) else: if not use_v2_migration_resolver: - print( # noqa + print( "\033[1;33mLiteLLM Proxy: Using default (v1) migration resolver. " "If your deployment has seen schema thrashing during rolling " "deploys, try --use_v2_migration_resolver (safer: avoids the " @@ -1201,7 +1201,7 @@ def run_server( # noqa: PLR0915 # (e.g. non-idempotent failures, permission issues). # v1 never raises here, so this only fires when the # operator opted into v2. - print( # noqa + print( "\033[1;31mLiteLLM Proxy: Database migration cannot proceed. " f"{e}\033[0m", file=sys.stderr, @@ -1210,19 +1210,19 @@ def run_server( # noqa: PLR0915 sys.exit(2) if not setup_ok: if enforce_prisma_migration_check: - print( # noqa + print( "\033[1;31mLiteLLM Proxy: Database setup failed after multiple retries. " "The proxy cannot start safely. Please check your database connection and migration status.\033[0m" ) sys.exit(1) else: - print( # noqa + print( "\033[1;33mLiteLLM Proxy: Database migration failed but continuing startup. " "Set --enforce_prisma_migration_check or ENFORCE_PRISMA_MIGRATION_CHECK=true to exit on failure.\033[0m" ) else: - print( # noqa - f"Unable to connect to DB. DATABASE_URL found in environment, but prisma package not found." # noqa + print( + f"Unable to connect to DB. DATABASE_URL found in environment, but prisma package not found." # noqa: F541 ) if port == 4000 and ProxyInitializationHelpers._is_port_in_use(port): port = random.randint(1024, 49152) @@ -1233,7 +1233,7 @@ def run_server( # noqa: PLR0915 litellm._turn_on_debug() # DO NOT DELETE - enables global variables to work across files - from litellm.proxy.proxy_server import app # noqa + from litellm.proxy.proxy_server import app # Auto-create PROMETHEUS_MULTIPROC_DIR for multi-worker setups ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir( @@ -1243,9 +1243,7 @@ def run_server( # noqa: PLR0915 # Skip server startup if requested (after all setup is done) if skip_server_startup: - print( # noqa - "LiteLLM: Setup complete. Skipping server startup as requested." - ) + print("LiteLLM: Setup complete. Skipping server startup as requested.") return running_uvicorn = run_gunicorn is False and run_hypercorn is False @@ -1263,8 +1261,8 @@ def run_server( # noqa: PLR0915 uvicorn_args["limit_max_requests"] = max_requests_before_restart if run_gunicorn is False and run_hypercorn is False and run_granian is False: if ssl_certfile_path is not None and ssl_keyfile_path is not None: - print( # noqa - f"\033[1;32mLiteLLM Proxy: Using SSL with certfile: {ssl_certfile_path} and keyfile: {ssl_keyfile_path}\033[0m\n" # noqa + print( + f"\033[1;32mLiteLLM Proxy: Using SSL with certfile: {ssl_certfile_path} and keyfile: {ssl_keyfile_path}\033[0m\n" ) uvicorn_args["ssl_keyfile"] = ssl_keyfile_path uvicorn_args["ssl_certfile"] = ssl_certfile_path diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b7ad86319ae..cae64cdf316 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -180,26 +180,26 @@ def generate_feedback_box(): # Select a random message message = random.choice(list_of_messages) - print() # noqa - print("\033[1;37m" + "#" + "-" * box_width + "#\033[0m") # noqa - print("\033[1;37m" + "#" + " " * box_width + "#\033[0m") # noqa - print("\033[1;37m" + "# {:^59} #\033[0m".format(message)) # noqa - print( # noqa + print() # noqa: T201 + print("\033[1;37m" + "#" + "-" * box_width + "#\033[0m") # noqa: T201 + print("\033[1;37m" + "#" + " " * box_width + "#\033[0m") # noqa: T201 + print("\033[1;37m" + "# {:^59} #\033[0m".format(message)) # noqa: T201 + print( # noqa: T201 "\033[1;37m" + "# {:^59} #\033[0m".format("https://github.com/BerriAI/litellm/issues/new") - ) # noqa - print("\033[1;37m" + "#" + " " * box_width + "#\033[0m") # noqa - print("\033[1;37m" + "#" + "-" * box_width + "#\033[0m") # noqa - print() # noqa - print(" Thank you for using LiteLLM! - Krrish & Ishaan") # noqa - print() # noqa - print() # noqa - print() # noqa - print( # noqa + ) + print("\033[1;37m" + "#" + " " * box_width + "#\033[0m") # noqa: T201 + print("\033[1;37m" + "#" + "-" * box_width + "#\033[0m") # noqa: T201 + print() # noqa: T201 + print(" Thank you for using LiteLLM! - Krrish & Ishaan") # noqa: T201 + print() # noqa: T201 + print() # noqa: T201 + print() # noqa: T201 + print( # noqa: T201 "\033[1;31mGive Feedback / Get Help: https://github.com/BerriAI/litellm/issues/new\033[0m" - ) # noqa - print() # noqa - print() # noqa + ) + print() # noqa: T201 + print() # noqa: T201 import contextlib @@ -745,7 +745,7 @@ async def _initialize_shared_aiohttp_session(): @asynccontextmanager -async def proxy_startup_event(app: FastAPI): # noqa: PLR0915 +async def proxy_startup_event(app: FastAPI): global prisma_client, master_key, use_background_health_checks, llm_router, llm_model_list, general_settings, proxy_budget_rescheduler_min_time, proxy_budget_rescheduler_max_time, litellm_proxy_admin_name, db_writer_client, store_model_in_db, premium_user, _license_check, proxy_batch_polling_interval, shared_aiohttp_session import json @@ -2496,7 +2496,7 @@ async def _invalidate_spend_counter(counter_key: str): ) -async def update_cache( # noqa: PLR0915 +async def update_cache( token: Optional[str], user_id: Optional[str], end_user_id: Optional[str], @@ -3805,9 +3805,9 @@ class ProxyConfig: search_tools_parsed: List[SearchToolTypedDict] = [] - print( # noqa + print( # noqa: T201 "\033[32mLiteLLM: Proxy initialized with Search Tools:\033[0m" - ) # noqa + ) for search_tool in search_tools_raw: # Display loaded search tool @@ -3815,7 +3815,9 @@ class ProxyConfig: search_provider = search_tool.get("litellm_params", {}).get( "search_provider", "" ) - print(f"\033[32m {search_tool_name} ({search_provider})\033[0m") # noqa + print( # noqa: T201 + f"\033[32m {search_tool_name} ({search_provider})\033[0m" + ) # Handle os.environ/ variables in litellm_params litellm_params = search_tool.get("litellm_params", {}) @@ -3898,7 +3900,7 @@ class ProxyConfig: premium_user = _license_check.is_premium() return - async def load_config( # noqa: PLR0915 + async def load_config( self, router: Optional[litellm.Router], config_file_path: str ): """ @@ -3925,7 +3927,7 @@ class ProxyConfig: reset_color_code = "\033[0m" for key, value in litellm_settings.items(): if key == "cache" and value is True: - print(f"{blue_color_code}\nSetting Cache on Proxy") # noqa + print(f"{blue_color_code}\nSetting Cache on Proxy") # noqa: T201 from litellm.caching.caching import Cache cache_params = {} @@ -4120,9 +4122,9 @@ class ProxyConfig: "mounting metrics endpoint" ) PrometheusLogger._mount_metrics_endpoint() - print( # noqa + print( # noqa: T201 f"{blue_color_code} Initialized Success Callbacks - {litellm.success_callback} {reset_color_code}" - ) # noqa + ) elif key == "failure_callback": litellm.failure_callback = [] @@ -4141,9 +4143,9 @@ class ProxyConfig: litellm.logging_callback_manager.add_litellm_failure_callback( callback ) - print( # noqa + print( # noqa: T201 f"{blue_color_code} Initialized Failure Callbacks - {litellm.failure_callback} {reset_color_code}" - ) # noqa + ) elif key == "audit_log_callbacks": from litellm.proxy.management_helpers.audit_logs import ( reset_audit_log_callback_cache, @@ -4167,9 +4169,9 @@ class ProxyConfig: "store_audit_logs", litellm.store_audit_logs ) if _store_audit_logs: - print( # noqa + print( # noqa: T201 f"{blue_color_code} Initialized Audit Log Callbacks - {litellm.audit_log_callbacks} {reset_color_code}" - ) # noqa + ) else: verbose_proxy_logger.warning( "'audit_log_callbacks' is configured but 'store_audit_logs' is not enabled. " @@ -4516,15 +4518,15 @@ class ProxyConfig: model_list = config.get("model_list", None) if model_list: router_params["model_list"] = model_list - print( # noqa + print( # noqa: T201 "\033[32mLiteLLM: Proxy initialized with Config, Set models:\033[0m" - ) # noqa + ) for model in model_list: ### LOAD FROM os.environ/ ### for k, v in model["litellm_params"].items(): if isinstance(v, str) and v.startswith("os.environ/"): model["litellm_params"][k] = get_secret(v) - print(f"\033[32m {model.get('model_name', '')}\033[0m") # noqa + print(f"\033[32m {model.get('model_name', '')}\033[0m") # noqa: T201 litellm_model_name = model["litellm_params"]["model"] litellm_model_api_base = model["litellm_params"].get("api_base", None) if "ollama" in litellm_model_name and litellm_model_api_base is None: @@ -6629,7 +6631,7 @@ def save_worker_config(**data): os.environ["WORKER_CONFIG"] = json.dumps(data) -async def initialize( # noqa: PLR0915 +async def initialize( model=None, alias=None, api_base=None, @@ -7020,7 +7022,7 @@ def _format_streaming_sse_chunk(chunk: Union[str, bytes]) -> Union[str, bytes]: return f"data: {chunk}\n\n" -async def async_data_generator( # noqa: PLR0915 +async def async_data_generator( response, user_api_key_dict: UserAPIKeyAuth, request_data: dict ): verbose_proxy_logger.debug("inside generator") @@ -7468,7 +7470,7 @@ class ProxyStartupEvent: ) @classmethod - async def initialize_scheduled_background_jobs( # noqa: PLR0915 + async def initialize_scheduled_background_jobs( cls, general_settings: dict, prisma_client: PrismaClient, @@ -7996,7 +7998,7 @@ class ProxyStartupEvent: and proxy_logging_obj.slack_alerting_instance.alerting is not None and prisma_client is not None ): - print("Alerting: Initializing Weekly/Monthly Spend Reports") # noqa + print("Alerting: Initializing Weekly/Monthly Spend Reports") # noqa: T201 spend_report_frequency: str = ( general_settings.get("spend_report_frequency", "7d") or "7d" ) @@ -8490,7 +8492,7 @@ async def model_info( tags=["chat/completions"], responses={200: {"description": "Successful response"}, **ERROR_RESPONSES}, ) # azure compatible endpoint -async def chat_completion( # noqa: PLR0915 +async def chat_completion( request: Request, fastapi_response: Response, model: Optional[str] = None, @@ -8679,7 +8681,7 @@ async def chat_completion( # noqa: PLR0915 dependencies=[Depends(user_api_key_auth)], tags=["completions"], ) -async def completion( # noqa: PLR0915 +async def completion( request: Request, fastapi_response: Response, model: Optional[str] = None, @@ -8894,7 +8896,7 @@ async def completion( # noqa: PLR0915 response_class=ORJSONResponse, tags=["embeddings"], ) # azure compatible endpoint -async def embeddings( # noqa: PLR0915 +async def embeddings( request: Request, fastapi_response: Response, model: Optional[str] = None, @@ -12640,7 +12642,7 @@ def _get_proxy_model_info(model: dict) -> dict: tags=["model management"], dependencies=[Depends(user_api_key_auth)], ) -async def model_info_v1( # noqa: PLR0915 +async def model_info_v1( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), litellm_model_id: Optional[str] = None, include_team_models: Optional[bool] = fastapi.Query( @@ -13370,7 +13372,7 @@ async def fallback_login(request: Request): @router.post( "/login", include_in_schema=False ) # hidden since this is a helper for UI sso login -async def login(request: Request): # noqa: PLR0915 +async def login(request: Request): global premium_user, general_settings, master_key from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object from litellm.proxy.utils import get_custom_url @@ -13420,7 +13422,7 @@ async def login(request: Request): # noqa: PLR0915 @router.post( "/v2/login", include_in_schema=False ) # hidden helper for UI logins via API -async def login_v2(request: Request): # noqa: PLR0915 +async def login_v2(request: Request): global premium_user, general_settings, master_key from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object from litellm.proxy.utils import get_custom_url @@ -13495,7 +13497,7 @@ async def login_v2(request: Request): # noqa: PLR0915 @router.post( "/v3/login", include_in_schema=False ) # control-plane login — always returns token in body for cross-origin use -async def login_v3(request: Request): # noqa: PLR0915 +async def login_v3(request: Request): global premium_user, general_settings, master_key from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object from litellm.proxy.utils import get_custom_url @@ -14409,7 +14411,7 @@ async def invitation_delete( dependencies=[Depends(user_api_key_auth)], include_in_schema=False, ) -async def update_config( # noqa: PLR0915 +async def update_config( config_info: ConfigYAML, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): @@ -15026,7 +15028,7 @@ async def delete_callback( include_in_schema=False, dependencies=[Depends(user_api_key_auth)], ) -async def get_config(): # noqa: PLR0915 +async def get_config(): """ For Admin UI - allows admin to view config via UI # return the callbacks and the env variables for the callback diff --git a/litellm/proxy/response_polling/background_streaming.py b/litellm/proxy/response_polling/background_streaming.py index 03039d4f441..a69e6734d71 100644 --- a/litellm/proxy/response_polling/background_streaming.py +++ b/litellm/proxy/response_polling/background_streaming.py @@ -21,7 +21,7 @@ from litellm.proxy.response_polling.polling_handler import ResponsePollingHandle from litellm.types.llms.openai import ResponsesAPIStatus -async def background_streaming_task( # noqa: PLR0915 +async def background_streaming_task( polling_id: str, data: dict, polling_handler: ResponsePollingHandler, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 3626a21516d..bbd8b75fdd6 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -264,7 +264,7 @@ async def add_shared_session_to_data(data: dict) -> None: pass -async def route_request( # noqa: PLR0915 - Complex routing function, refactoring tracked separately +async def route_request( data: dict, llm_router: Optional[LitellmRouter], user_model: Optional[str], diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index ef06adb27fc..0ba77dcd2f0 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1734,7 +1734,7 @@ async def calculate_spend(request: SpendCalculateRequest): 200: {"model": List[LiteLLM_SpendLogs]}, }, ) -async def ui_view_spend_logs( # noqa: PLR0915 +async def ui_view_spend_logs( request: Request, api_key: Optional[str] = fastapi.Query( default=None, @@ -2273,7 +2273,7 @@ async def ui_view_request_response_for_request_id( 200: {"model": List[LiteLLM_SpendLogs]}, }, ) -async def view_spend_logs( # noqa: PLR0915 +async def view_spend_logs( api_key: Optional[str] = fastapi.Query( default=None, description="Get spend logs based on api key", diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index d215294fd04..aef06a3c668 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -228,9 +228,7 @@ def _extract_usage_for_ocr_call(response_obj: Any, response_obj_dict: dict) -> d return {} -def get_logging_payload( # noqa: PLR0915 - kwargs, response_obj, start_time, end_time -) -> SpendLogsPayload: +def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogsPayload: if kwargs is None: kwargs = {} diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index e81bf85604d..74f12a1eeb1 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -191,7 +191,7 @@ def print_verbose(print_statement): verbose_proxy_logger.debug("{}\n{}".format(print_statement, traceback.format_exc())) if litellm.set_verbose: - print(f"LiteLLM Proxy: {_redact_string(str(print_statement))}") # noqa + print(f"LiteLLM Proxy: {_redact_string(str(print_statement))}") # noqa: T201 def _get_email_logger_class(): @@ -1171,6 +1171,8 @@ class ProxyLogging: response=response, data=data, call_type=call_type ) + callback.mark_pre_call_hook_ran(data) + except SensitiveDataRouteException: status = "intervened" raise @@ -3345,7 +3347,7 @@ class PrismaClient: on_backoff=on_backoff, # specifying the function to call on backoff ) @log_db_metrics - async def get_data( # noqa: PLR0915 + async def get_data( self, token: Optional[Union[str, list]] = None, user_id: Optional[str] = None, @@ -3788,7 +3790,7 @@ class PrismaClient: max_time=10, # maximum total time to retry for on_backoff=on_backoff, # specifying the function to call on backoff ) - async def insert_data( # noqa: PLR0915 + async def insert_data( self, data: dict, table_name: Literal[ @@ -3938,7 +3940,7 @@ class PrismaClient: max_time=10, # maximum total time to retry for on_backoff=on_backoff, # specifying the function to call on backoff ) - async def update_data( # noqa: PLR0915 + async def update_data( self, token: Optional[str] = None, data: dict = {}, @@ -5514,7 +5516,7 @@ class ProxyUpdateSpend: return False -async def update_spend( # noqa: PLR0915 +async def update_spend( prisma_client: PrismaClient, db_writer_client: Optional[AsyncHTTPHandler], proxy_logging_obj: ProxyLogging, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 7031ecaa1a0..f6f0a92def0 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -284,7 +284,7 @@ async def arealtime_calls( @wrapper_client -async def _arealtime( # noqa: PLR0915 +async def _arealtime( model: str, websocket: Any, # fastapi websocket api_base: Optional[str] = None, diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index e27585116ce..e40e12e9197 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -75,7 +75,7 @@ async def arerank( @client -def rerank( # noqa: PLR0915 +def rerank( model: str, query: str, documents: List[Union[str, Dict[str, Any]]], diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index d3d30642216..5b5ff122c50 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -2,6 +2,7 @@ Handles transforming from Responses API -> LiteLLM completion (Chat Completion API) """ +import re from collections.abc import Sequence from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast @@ -1554,6 +1555,20 @@ class LiteLLMCompletionResponsesConfig: # Default to completed for unknown finish reasons return "completed" + @staticmethod + def _tool_call_id_from_responses_item( + item_id: Optional[str], call_id: Optional[str] + ) -> str: + """Bedrock Mantle returns a non-unique, index-based ``call_id`` (``call_0``, + ``call_1``, ... that resets every response) alongside a unique ``id`` + (``fc_...``). ``call_id`` is the canonical Responses API correlation key, so + prefer it; fall back to the unique ``id`` only when ``call_id`` is absent or + in that degenerate index form, otherwise multi-turn tool calls collide and an + agent cannot correlate its tool results.""" + if call_id and re.fullmatch(r"call_\d+", call_id) is None: + return call_id + return item_id or call_id or "" + @staticmethod def convert_response_function_tool_call_to_chat_completion_tool_call( tool_call_item: Any, @@ -1601,7 +1616,10 @@ class LiteLLMCompletionResponsesConfig: function_dict["provider_specific_fields"] = provider_specific_fields tool_call_dict: Dict[str, Any] = { - "id": tool_call_item.call_id, + "id": LiteLLMCompletionResponsesConfig._tool_call_id_from_responses_item( + getattr(tool_call_item, "id", None), + getattr(tool_call_item, "call_id", None), + ), "function": function_dict, "type": "function", "index": 0, diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 24b5db28571..acb7487f430 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -77,7 +77,7 @@ def _add_mcp_metadata_to_response( setattr(message, "provider_specific_fields", provider_fields) -async def acompletion_with_mcp( # noqa: PLR0915 +async def acompletion_with_mcp( model: str, messages: List, tools: Optional[List] = None, diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 94cff6922b5..df5de205d45 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -644,7 +644,7 @@ class LiteLLM_Proxy_MCP_Handler: return result_text or "Tool executed successfully" @staticmethod - async def _execute_tool_calls( # noqa: PLR0915 + async def _execute_tool_calls( tool_server_map: dict[str, str], tool_calls: List[Any], user_api_key_auth: Any, diff --git a/litellm/router.py b/litellm/router.py index 80584858311..5f26097443f 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -241,7 +241,7 @@ class Router: lowesttpm_logger: Optional[LowestTPMLoggingHandler] = None optional_callbacks: Optional[List[Union[CustomLogger, Callable, str]]] = None - def __init__( # noqa: PLR0915 + def __init__( self, model_list: Optional[ Union[List[DeploymentTypedDict], List[Dict[str, Any]]] @@ -2521,10 +2521,22 @@ class Router: from litellm.exceptions import MidStreamFallbackError from litellm.responses.streaming_iterator import ( BaseResponsesAPIStreamingIterator, + _get_openai_response_types, ) source_iterator = response + # Pre-resolve the set of terminal stream event types so the + # per-chunk type check inside FallbackResponsesStreamWrapper + # stays cheap; mirrors the source-iterator filter at + # responses/streaming_iterator.py:243-247. + _openai_types = _get_openai_response_types() + _RESPONSES_TERMINAL_EVENT_TYPES = ( + _openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + _openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + _openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED, + ) + class FallbackResponsesStreamWrapper(BaseResponsesAPIStreamingIterator): """ Subclasses BaseResponsesAPIStreamingIterator only for isinstance @@ -2550,9 +2562,16 @@ class Router: # is missing many of these attributes — use getattr fallbacks # so wrapper construction never raises AttributeError. The # bridge stores the logging object as `litellm_logging_obj`. - self.response = getattr(source_iterator, "response", None) - self.model = getattr(source_iterator, "model", None) - self.logging_obj = getattr( + # base class declares non-Optional types for these + # fields but the bridge path (LiteLLMCompletionStreamingIterator) + # can legitimately omit them at runtime — keep the None + # fallback. Same lines passed mypy on the pre-fix file + # because the surrounding function body wasn't fully + # type-narrowed; the new typed terminal-event tuple above + # is what made these surface. + self.response = getattr(source_iterator, "response", None) # type: ignore[assignment] + self.model = getattr(source_iterator, "model", None) # type: ignore[assignment] + self.logging_obj = getattr( # type: ignore[assignment] source_iterator, "logging_obj", getattr(source_iterator, "litellm_logging_obj", None), @@ -2587,7 +2606,23 @@ class Router: return self async def __anext__(self): - return await self._async_generator.__anext__() + chunk = await self._async_generator.__anext__() + # Sniff the terminal stream event off each forwarded chunk + # so ``self.completed_response`` is populated regardless of + # which inner iterator produced it (source_iterator, + # fallback_iterator, or any future wrapper). Without this + # the proxy's container-ownership hook (which reads + # ``getattr(stream_response, "completed_response", None)`` + # via _extract_completed_responses_response) silently + # records nothing on streaming /v1/responses calls — every + # follow-up /v1/containers//files call then 403s for + # the very key that created the container (#30210). + if ( + self.completed_response is None + and getattr(chunk, "type", None) in _RESPONSES_TERMINAL_EVENT_TYPES + ): + self.completed_response = chunk + return chunk async def aclose(self): # async generators always expose aclose — no defensive check needed. @@ -2691,7 +2726,7 @@ class Router: return FallbackResponsesStreamWrapper(stream_with_fallbacks()) - def _completion_streaming_iterator( # noqa: PLR0915 + def _completion_streaming_iterator( self, model_response: CustomStreamWrapper, messages: List[Dict[str, str]], @@ -2852,7 +2887,7 @@ class Router: f"Silent experiment failed for model {silent_model}: {str(e)}" ) - async def _acompletion( # noqa: PLR0915 + async def _acompletion( self, model: str, messages: List[Dict[str, str]], **kwargs ) -> Union[ ModelResponse, @@ -5123,7 +5158,7 @@ class Router: ) raise e - async def _acreate_file( # noqa: PLR0915 + async def _acreate_file( self, model: str, **kwargs, @@ -6432,7 +6467,7 @@ class Router: # propagate so they remain visible. return None - async def async_function_with_fallbacks_common_utils( # noqa: PLR0915 + async def async_function_with_fallbacks_common_utils( self, e: Exception, disable_fallbacks: Optional[bool], @@ -6808,7 +6843,7 @@ class Router: ) @tracer.wrap() - async def async_function_with_retries(self, *args, **kwargs): # noqa: PLR0915 + async def async_function_with_retries(self, *args, **kwargs): verbose_router_logger.debug("Inside async function with retries.") original_function = kwargs.pop("original_function") fallbacks = kwargs.pop("fallbacks", self.fallbacks) @@ -9289,7 +9324,7 @@ class Router: return model_info - def _set_model_group_info( # noqa: PLR0915 + def _set_model_group_info( self, model_group: str, user_facing_model_group_name: str ) -> Optional[ModelGroupInfo]: """ @@ -10531,7 +10566,7 @@ class Router: ) return client - def _pre_call_checks( # noqa: PLR0915 + def _pre_call_checks( self, model: str, healthy_deployments: List, diff --git a/litellm/router_strategy/complexity_router/evals/eval_complexity_router.py b/litellm/router_strategy/complexity_router/evals/eval_complexity_router.py index 939ecdd2d22..6c9318e83bb 100644 --- a/litellm/router_strategy/complexity_router/evals/eval_complexity_router.py +++ b/litellm/router_strategy/complexity_router/evals/eval_complexity_router.py @@ -250,10 +250,10 @@ def run_eval() -> Tuple[int, int, List[dict]]: total = len(EVAL_CASES) failures = [] - print("=" * 70) # noqa: T201 - print("COMPLEXITY ROUTER EVALUATION") # noqa: T201 - print("=" * 70) # noqa: T201 - print() # noqa: T201 + print("=" * 70) + print("COMPLEXITY ROUTER EVALUATION") + print("=" * 70) + print() for i, case in enumerate(EVAL_CASES, 1): tier, score, signals = router.classify(case.prompt, case.system_prompt) @@ -292,33 +292,33 @@ def run_eval() -> Tuple[int, int, List[dict]]: ) # Print result - print(f"[{i:2d}] {status} | {case.description}") # noqa: T201 + print(f"[{i:2d}] {status} | {case.description}") print( f" Expected: {case.expected_tier.value:10s} | Got: {tier.value:10s} | Score: {score:+.3f}" - ) # noqa: T201 + ) if signals: - print(f" Signals: {', '.join(signals)}") # noqa: T201 + print(f" Signals: {', '.join(signals)}") if not is_pass: - print(f" Prompt: {case.prompt[:60]}...") # noqa: T201 - print() # noqa: T201 + print(f" Prompt: {case.prompt[:60]}...") + print() # Summary - print("=" * 70) # noqa: T201 - print(f"RESULTS: {passed}/{total} passed ({100*passed/total:.1f}%)") # noqa: T201 - print("=" * 70) # noqa: T201 + print("=" * 70) + print(f"RESULTS: {passed}/{total} passed ({100*passed/total:.1f}%)") + print("=" * 70) if failures: - print("\nFAILURES:") # noqa: T201 - print("-" * 70) # noqa: T201 + print("\nFAILURES:") + print("-" * 70) for f in failures: - print(f"Case {f['case']}: {f['description']}") # noqa: T201 + print(f"Case {f['case']}: {f['description']}") print( f" Expected: {f['expected']}, Got: {f['actual']} (score: {f['score']})" - ) # noqa: T201 - print(f" Signals: {f['signals']}") # noqa: T201 + ) + print(f" Signals: {f['signals']}") if f["acceptable"]: - print(f" Acceptable: {f['acceptable']}") # noqa: T201 - print() # noqa: T201 + print(f" Acceptable: {f['acceptable']}") + print() return passed, total, failures @@ -330,17 +330,13 @@ def main(): # Exit with error code if too many failures pass_rate = passed / total if pass_rate < 0.80: - print( - f"\n❌ EVAL FAILED: Pass rate {pass_rate:.1%} is below 80% threshold" - ) # noqa: T201 + print(f"\n❌ EVAL FAILED: Pass rate {pass_rate:.1%} is below 80% threshold") sys.exit(1) elif pass_rate < 0.90: - print( - f"\n⚠️ EVAL WARNING: Pass rate {pass_rate:.1%} is below 90%" - ) # noqa: T201 + print(f"\n⚠️ EVAL WARNING: Pass rate {pass_rate:.1%} is below 90%") sys.exit(0) else: - print(f"\n✅ EVAL PASSED: Pass rate {pass_rate:.1%}") # noqa: T201 + print(f"\n✅ EVAL PASSED: Pass rate {pass_rate:.1%}") sys.exit(0) diff --git a/litellm/router_strategy/lowest_cost.py b/litellm/router_strategy/lowest_cost.py index 54498363f51..3f641d4f0fb 100644 --- a/litellm/router_strategy/lowest_cost.py +++ b/litellm/router_strategy/lowest_cost.py @@ -190,7 +190,7 @@ class LowestCostLoggingHandler(CustomLogger): ) pass - async def async_get_available_deployments( # noqa: PLR0915 + async def async_get_available_deployments( self, model_group: str, healthy_deployments: list, diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 870b3f29d48..3adb8d43920 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -35,9 +35,7 @@ class LowestLatencyLoggingHandler(CustomLogger): self.router_cache = router_cache self.routing_args = RoutingArgs(**routing_args) - def log_success_event( # noqa: PLR0915 - self, kwargs, response_obj, start_time, end_time - ): + def log_success_event(self, kwargs, response_obj, start_time, end_time): try: """ Update latency usage on success @@ -259,9 +257,7 @@ class LowestLatencyLoggingHandler(CustomLogger): ) pass - async def async_log_success_event( # noqa: PLR0915 - self, kwargs, response_obj, start_time, end_time - ): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: """ Update latency usage on success @@ -413,7 +409,7 @@ class LowestLatencyLoggingHandler(CustomLogger): ) pass - def _get_available_deployments( # noqa: PLR0915 + def _get_available_deployments( self, model_group: str, healthy_deployments: list, diff --git a/litellm/router_strategy/lowest_tpm_rpm.py b/litellm/router_strategy/lowest_tpm_rpm.py index 488f8450941..f807ba7232a 100644 --- a/litellm/router_strategy/lowest_tpm_rpm.py +++ b/litellm/router_strategy/lowest_tpm_rpm.py @@ -158,7 +158,7 @@ class LowestTPMLoggingHandler(CustomLogger): verbose_router_logger.debug(traceback.format_exc()) pass - def get_available_deployments( # noqa: PLR0915 + def get_available_deployments( self, model_group: str, healthy_deployments: list, diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index 23e8896cd5f..22664dcb704 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -380,7 +380,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): potential_deployments = [_deployment] return potential_deployments - def _common_checks_available_deployment( # noqa: PLR0915 + def _common_checks_available_deployment( self, model_group: str, healthy_deployments: list, diff --git a/litellm/secret_managers/main.py b/litellm/secret_managers/main.py index 5c31d81f04c..f4b1d4a1b69 100644 --- a/litellm/secret_managers/main.py +++ b/litellm/secret_managers/main.py @@ -156,7 +156,7 @@ def get_secret_bool( return str_to_bool(_secret_value) -def get_secret( # noqa: PLR0915 +def get_secret( secret_name: str, default_value: Optional[Union[str, bool]] = None, ): diff --git a/litellm/secret_managers/secret_manager_handler.py b/litellm/secret_managers/secret_manager_handler.py index 4ff94d18eff..3a3cf6272dc 100644 --- a/litellm/secret_managers/secret_manager_handler.py +++ b/litellm/secret_managers/secret_manager_handler.py @@ -23,7 +23,7 @@ def _is_base64(s): return False -def get_secret_from_manager( # noqa: PLR0915 +def get_secret_from_manager( client: Any, key_manager: str, secret_name: str, diff --git a/litellm/types/router.py b/litellm/types/router.py index 5047cee424b..1611f1e5538 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -166,6 +166,13 @@ class CredentialLiteLLMParams(BaseModel): api_key: Optional[str] = None api_base: Optional[str] = None api_version: Optional[str] = None + ## AZURE OAUTH ## + # Without this field, ``get_deployment_credentials_with_provider`` + # round-trips ``litellm_params`` through a strict Pydantic dump and + # silently drops the OAuth token before the files/batch/passthrough + # callers see it, breaking Azure deployments configured with + # ``azure_ad_token`` instead of a static ``api_key`` (#30235). + azure_ad_token: Optional[str] = None ## VERTEX AI ## vertex_project: Optional[str] = None vertex_location: Optional[str] = None diff --git a/litellm/types/utils.py b/litellm/types/utils.py index d3dc7eadb94..f2152577b4d 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1572,7 +1572,7 @@ class Usage(SafeAttributeModel, CompletionUsage): prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None """Breakdown of tokens used in the prompt.""" - def __init__( # noqa: PLR0915 + def __init__( self, prompt_tokens: Optional[int] = None, completion_tokens: Optional[int] = None, @@ -1908,7 +1908,7 @@ class ModelResponse(ModelResponseBase): choices: List[Choices] """The list of completion choices the model generated for the input prompt.""" - def __init__( # noqa: PLR0915 + def __init__( self, id=None, choices=None, @@ -3338,6 +3338,7 @@ class LlmProviders(str, Enum): CODESTRAL = "codestral" TEXT_COMPLETION_CODESTRAL = "text-completion-codestral" DASHSCOPE = "dashscope" + MODELSCOPE = "modelscope" MOONSHOT = "moonshot" PUBLICAI = "publicai" V0 = "v0" @@ -3420,6 +3421,7 @@ class LlmProviders(str, Enum): PARASAIL = "parasail" XIAOMI_MIMO = "xiaomi_mimo" TENSORMESH = "tensormesh" + LIBERTAI = "libertai" LITELLM_AGENT = "litellm_agent" CURSOR = "cursor" BEDROCK_MANTLE = "bedrock_mantle" @@ -3455,6 +3457,7 @@ class SearchProviders(str, Enum): GOOGLE_PSE = "google_pse" DATAFORSEO = "dataforseo" FIRECRAWL = "firecrawl" + FASTCRW = "fastcrw" SEARXNG = "searxng" LINKUP = "linkup" DUCKDUCKGO = "duckduckgo" diff --git a/litellm/utils.py b/litellm/utils.py index 4c67abdf937..9c5989a11d3 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -488,7 +488,7 @@ def print_verbose( elif log_level == "ERROR": verbose_logger.error(print_statement) if litellm.set_verbose is True and logger_only is False: - print(print_statement) # noqa + print(print_statement) # noqa: T201 except Exception: pass @@ -760,7 +760,7 @@ def _remove_thought_signatures_from_messages( return processed_messages -def function_setup( # noqa: PLR0915 +def function_setup( original_function: str, rules_obj, start_time, *args, **kwargs ): # just run once to check if user wants to send their data anywhere - PostHog/Sentry/Slack/etc. ### NOTICES ### @@ -1422,12 +1422,12 @@ def post_call_processing( raise e -def client(original_function): # noqa: PLR0915 +def client(original_function): Rules = getattr(sys.modules[__name__], "Rules") rules_obj = Rules() @wraps(original_function) - def wrapper(*args, **kwargs): # noqa: PLR0915 + def wrapper(*args, **kwargs): # DO NOT MOVE THIS. It always needs to run first # Check if this is an async function. If so only execute the async function call_type = original_function.__name__ @@ -1775,7 +1775,7 @@ def client(original_function): # noqa: PLR0915 raise e @wraps(original_function) - async def wrapper_async(*args, **kwargs): # noqa: PLR0915 + async def wrapper_async(*args, **kwargs): print_args_passed_to_litellm(original_function, args, kwargs) start_time = datetime.datetime.now() result = None @@ -2942,7 +2942,7 @@ def _resolve_builtin_model_cost_entry( return None -def register_model(model_cost: Union[str, dict]): # noqa: PLR0915 +def register_model(model_cost: Union[str, dict]): """ Register new / Override existing models (and their pricing) to specific providers. Provide EITHER a model cost dictionary or a url to a hosted json blob @@ -3015,6 +3015,21 @@ def register_model(model_cost: Union[str, dict]): # noqa: PLR0915 # custom pricing on subsequent cost lookups. if existing_model.get("litellm_provider") is None: existing_model.pop("litellm_provider", None) + # Same pattern for cost fields (#30198): ``_get_model_info_helper`` + # synthesizes ``input_cost_per_token`` / ``output_cost_per_token`` + # = 0 when they are absent from the raw entry. Writing those zeros + # back flips a sparse entry from "no cost keys" (priced via name) + # to "cost keys = 0" (free), which makes + # ``_is_cost_explicitly_configured`` return True and silently + # disables budget enforcement on the next re-registration. + _raw_entry = litellm.model_cost.get(model_cost_key) + if _raw_entry is None: + _raw_entry = litellm.model_cost.get(key) + if _raw_entry is None: + _raw_entry = {} + for _cost_field in ("input_cost_per_token", "output_cost_per_token"): + if _cost_field not in _raw_entry and _cost_field not in value: + existing_model.pop(_cost_field, None) ## override / add new keys to the existing model cost dictionary updated_dictionary = _update_dictionary(existing_model, value) litellm.model_cost.setdefault(model_cost_key, {}).update(updated_dictionary) @@ -3350,7 +3365,7 @@ def get_optional_params_image_gen( return optional_params -def get_optional_params_embeddings( # noqa: PLR0915 +def get_optional_params_embeddings( # 2 optional params model: str, user: Optional[str] = None, @@ -4097,7 +4112,7 @@ def pre_process_optional_params( return optional_params -def get_optional_params( # noqa: PLR0915 +def get_optional_params( # use the openai defaults # https://platform.openai.com/docs/api-reference/chat/create model: str, @@ -5827,7 +5842,7 @@ def _is_potential_model_name_in_model_cost( ) -def _get_model_info_helper( # noqa: PLR0915 +def _get_model_info_helper( model: str, custom_llm_provider: Optional[str] = None, api_base: Optional[str] = None, @@ -6551,7 +6566,7 @@ def create_proxy_transport_and_mounts(): return sync_proxy_mounts, async_proxy_mounts -def validate_environment( # noqa: PLR0915 +def validate_environment( model: Optional[str] = None, api_key: Optional[str] = None, api_base: Optional[str] = None, @@ -6833,6 +6848,11 @@ def validate_environment( # noqa: PLR0915 keys_in_environment = True else: missing_keys.append("DASHSCOPE_API_KEY") + elif custom_llm_provider == "modelscope": + if "MODELSCOPE_API_KEY" in os.environ: + keys_in_environment = True + else: + missing_keys.append("MODELSCOPE_API_KEY") elif custom_llm_provider == "moonshot": if "MOONSHOT_API_KEY" in os.environ: keys_in_environment = True @@ -8504,6 +8524,7 @@ class ProviderConfigManager: LlmProviders.NEBIUS: (lambda: litellm.NebiusConfig(), False), LlmProviders.WANDB: (lambda: litellm.WandbConfig(), False), LlmProviders.DASHSCOPE: (lambda: litellm.DashScopeChatConfig(), False), + LlmProviders.MODELSCOPE: (lambda: litellm.ModelScopeChatConfig(), False), LlmProviders.MOONSHOT: (lambda: litellm.MoonshotChatConfig(), False), LlmProviders.DOCKER_MODEL_RUNNER: ( lambda: litellm.DockerModelRunnerChatConfig(), @@ -8609,7 +8630,7 @@ class ProviderConfigManager: return LangFlowConfig() @staticmethod - def get_provider_chat_config( # noqa: PLR0915 + def get_provider_chat_config( model: str, provider: LlmProviders, base_model: Optional[str] = None, @@ -9442,6 +9463,12 @@ class ProviderConfigManager: ) return get_dashscope_image_generation_config(model) + elif LlmProviders.MODELSCOPE == provider: + from litellm.llms.modelscope.image_generation import ( + get_modelscope_image_generation_config, + ) + + return get_modelscope_image_generation_config(model) return None @staticmethod @@ -9647,6 +9674,7 @@ class ProviderConfigManager: from litellm.llms.dataforseo.search.transformation import DataForSEOSearchConfig from litellm.llms.duckduckgo.search.transformation import DuckDuckGoSearchConfig from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig + from litellm.llms.fastcrw.search.transformation import FastCRWSearchConfig from litellm.llms.firecrawl.search.transformation import FirecrawlSearchConfig from litellm.llms.google_pse.search.transformation import GooglePSESearchConfig from litellm.llms.linkup.search.transformation import LinkupSearchConfig @@ -9669,6 +9697,7 @@ class ProviderConfigManager: SearchProviders.GOOGLE_PSE: GooglePSESearchConfig, SearchProviders.DATAFORSEO: DataForSEOSearchConfig, SearchProviders.FIRECRAWL: FirecrawlSearchConfig, + SearchProviders.FASTCRW: FastCRWSearchConfig, SearchProviders.SEARXNG: SearXNGSearchConfig, SearchProviders.LINKUP: LinkupSearchConfig, SearchProviders.DUCKDUCKGO: DuckDuckGoSearchConfig, diff --git a/litellm/videos/main.py b/litellm/videos/main.py index a61fe99d584..b087f1e88d8 100644 --- a/litellm/videos/main.py +++ b/litellm/videos/main.py @@ -159,7 +159,7 @@ def video_generation( @client -def video_generation( # noqa: PLR0915 +def video_generation( prompt: str, model: Optional[str] = None, input_reference: Optional[FileTypes] = None, @@ -569,7 +569,7 @@ def video_remix( @client -def video_remix( # noqa: PLR0915 +def video_remix( video_id: str, prompt: str, timeout=600, # default to 10 minutes @@ -790,7 +790,7 @@ def video_list( @client -def video_list( # noqa: PLR0915 +def video_list( after: Optional[str] = None, limit: Optional[int] = None, order: Optional[str] = None, @@ -993,7 +993,7 @@ def video_status( @client -def video_status( # noqa: PLR0915 +def video_status( video_id: str, timeout=600, # default to 10 minutes custom_llm_provider=None, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b181df94131..f0c15654cfe 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -18822,6 +18822,38 @@ "supports_response_schema": true, "supports_vision": true }, + "github_copilot/mai-code-1-flash": { + "cache_read_input_token_cost": 7.5e-08, + "input_cost_per_token": 7.5e-07, + "litellm_provider": "github_copilot", + "max_input_tokens": 128000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 4.5e-06, + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true + }, + "github_copilot/mai-code-1-flash-internal": { + "cache_read_input_token_cost": 7.5e-08, + "input_cost_per_token": 7.5e-07, + "litellm_provider": "github_copilot", + "max_input_tokens": 128000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 4.5e-06, + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true + }, "github_copilot/text-embedding-3-small": { "litellm_provider": "github_copilot", "max_input_tokens": 8191, @@ -40986,6 +41018,174 @@ "litellm_provider": "llamagate", "mode": "embedding" }, + "libertai/hermes-3-8b-tee": { + "max_tokens": 16000, + "max_input_tokens": 16000, + "max_output_tokens": 16000, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": false, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/gemma-4-31b-it": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 4e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/gemma-4-31b-it-thinking": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 4e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_reasoning": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.6-27b": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.6-27b-thinking": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_reasoning": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.6-35b-a3b": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.6-35b-a3b-thinking": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_reasoning": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.5-122b-a10b": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.75e-06, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/qwen3.5-122b-a10b-thinking": { + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.75e-06, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_reasoning": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/deepseek-v4-flash": { + "max_tokens": 200000, + "max_input_tokens": 200000, + "max_output_tokens": 200000, + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.75e-06, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": false, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/deepseek-v4-flash-thinking": { + "max_tokens": 200000, + "max_input_tokens": 200000, + "max_output_tokens": 200000, + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.75e-06, + "litellm_provider": "libertai", + "mode": "chat", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_vision": false, + "supports_reasoning": true, + "source": "https://docs.libertai.io/apis/text/" + }, + "libertai/bge-m3": { + "max_tokens": 8192, + "max_input_tokens": 8192, + "input_cost_per_token": 1e-08, + "output_cost_per_token": 0.0, + "litellm_provider": "libertai", + "mode": "embedding", + "source": "https://docs.libertai.io/apis/text/" + }, "sarvam/sarvam-m": { "cache_creation_input_token_cost": 0, "cache_creation_input_token_cost_above_1hr": 0, diff --git a/mypy-code-budget.json b/mypy-code-budget.json new file mode 100644 index 00000000000..2cae0d661e9 --- /dev/null +++ b/mypy-code-budget.json @@ -0,0 +1,18 @@ +{ + "import-not-found": { + "baseline": 8, + "slack": 3 + }, + "no-any-return": { + "baseline": 902, + "slack": 10 + }, + "no-untyped-def": { + "baseline": 4888, + "slack": 10 + }, + "valid-type": { + "baseline": 1, + "slack": 3 + } +} diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 2ad2b3ec982..b90e5d2698d 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -972,6 +972,23 @@ "search": true } }, + "fastcrw": { + "display_name": "fastCRW (`fastcrw`)", + "url": "https://docs.litellm.ai/docs/search/fastcrw", + "endpoints": { + "chat_completions": false, + "messages": false, + "responses": false, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "search": true + } + }, "linkup": { "display_name": "Linkup (`linkup`)", "url": "https://docs.litellm.ai/docs/search/linkup", @@ -1359,6 +1376,23 @@ "interactions": true } }, + "libertai": { + "display_name": "LibertAI (`libertai`)", + "url": "https://docs.litellm.ai/docs/providers/libertai", + "endpoints": { + "chat_completions": true, + "messages": true, + "responses": false, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": false + } + }, "litellm_proxy": { "display_name": "LiteLLM Proxy (`litellm_proxy`)", "url": "https://docs.litellm.ai/docs/providers/litellm_proxy", @@ -1468,6 +1502,24 @@ "interactions": true } }, + "modelscope": { + "display_name": "ModelScope (`modelscope`)", + "url": "https://docs.litellm.ai/docs/providers/modelscope", + "endpoints": { + "chat_completions": true, + "messages": true, + "responses": true, + "embeddings": false, + "image_generations": true, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": false, + "interactions": false + } + }, "moonshot": { "display_name": "Moonshot (`moonshot`)", "url": "https://docs.litellm.ai/docs/providers/moonshot", diff --git a/pyproject.toml b/pyproject.toml index 6429b810969..8b1386aaf87 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -149,6 +149,7 @@ dev = [ "flake8==7.3.0", "black==26.3.1", "mypy==1.19.0", + "basedpyright==1.39.7", "pytest==9.0.3", "pytest-mock==3.15.1", "pytest-asyncio==1.3.0", @@ -220,7 +221,6 @@ ci = [ "blockbuster==1.5.26", "beautifulsoup4==4.14.3", "pylint==4.0.5", - "pyright==1.1.408", "langchain-mcp-adapters==0.2.1", "langchain-openai==1.1.14", "langgraph==1.0.10", diff --git a/pyrightconfig.json b/pyrightconfig.json index f930e44d305..97f099d5b2c 100644 --- a/pyrightconfig.json +++ b/pyrightconfig.json @@ -1,7 +1,12 @@ { + "include": ["litellm"], "ignore": [], "exclude": ["**/node_modules", "**/__pycache__", "litellm/types/utils.py", "litellm/proxy/_types.py"], + "pythonVersion": "3.12", + "typeCheckingMode": "strict", + "enableTypeIgnoreComments": false, "reportMissingImports": false, - "reportPrivateImportUsage": false + "reportPrivateImportUsage": false, + "reportExplicitAny": "error", + "reportAny": "error" } - \ No newline at end of file diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 6363b72353f..62ebdb559fc 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -1,12 +1,494 @@ { - "ANN001": { "baseline": 2865, "slack": 10 }, - "ANN002": { "baseline": 64, "slack": 3 }, - "ANN003": { "baseline": 759, "slack": 10 }, - "ANN401": { "baseline": 1885, "slack": 10 }, - "B006": { "baseline": 180, "slack": 3 }, - "C901": { "baseline": 301, "slack": 3 }, - "PLR0913": { "baseline": 1813, "slack": 3 }, - "PLW0603": { "baseline": 183, "slack": 3 }, - "RUF012": { "baseline": 158, "slack": 3 }, - "TID251": { "baseline": 2404, "slack": 10 } + "ANN001": { + "baseline": 2865, + "slack": 50 + }, + "ANN002": { + "baseline": 64, + "slack": 5 + }, + "ANN003": { + "baseline": 759, + "slack": 30 + }, + "ANN201": { + "baseline": 1944, + "slack": 50 + }, + "ANN202": { + "baseline": 858, + "slack": 30 + }, + "ANN204": { + "baseline": 658, + "slack": 20 + }, + "ANN205": { + "baseline": 117, + "slack": 10 + }, + "ANN206": { + "baseline": 120, + "slack": 10 + }, + "ANN401": { + "baseline": 1886, + "slack": 50 + }, + "ASYNC230": { + "baseline": 11, + "slack": 3 + }, + "B004": { + "baseline": 1, + "slack": 3 + }, + "B006": { + "baseline": 180, + "slack": 10 + }, + "B008": { + "baseline": 490, + "slack": 15 + }, + "B009": { + "baseline": 79, + "slack": 5 + }, + "B010": { + "baseline": 187, + "slack": 10 + }, + "B018": { + "baseline": 2, + "slack": 3 + }, + "B019": { + "baseline": 1, + "slack": 3 + }, + "B021": { + "baseline": 1, + "slack": 3 + }, + "B026": { + "baseline": 3, + "slack": 3 + }, + "B033": { + "baseline": 1, + "slack": 3 + }, + "BLE001": { + "baseline": 2854, + "slack": 50 + }, + "C401": { + "baseline": 8, + "slack": 3 + }, + "C404": { + "baseline": 1, + "slack": 3 + }, + "C405": { + "baseline": 20, + "slack": 3 + }, + "C408": { + "baseline": 11, + "slack": 3 + }, + "C414": { + "baseline": 4, + "slack": 3 + }, + "C419": { + "baseline": 1, + "slack": 3 + }, + "C901": { + "baseline": 301, + "slack": 15 + }, + "D419": { + "baseline": 6, + "slack": 3 + }, + "DTZ001": { + "baseline": 2, + "slack": 3 + }, + "DTZ003": { + "baseline": 30, + "slack": 3 + }, + "DTZ005": { + "baseline": 229, + "slack": 15 + }, + "DTZ006": { + "baseline": 10, + "slack": 3 + }, + "DTZ007": { + "baseline": 20, + "slack": 3 + }, + "DTZ011": { + "baseline": 3, + "slack": 3 + }, + "EXE001": { + "baseline": 4, + "slack": 3 + }, + "EXE002": { + "baseline": 3, + "slack": 3 + }, + "F401": { + "baseline": 20, + "slack": 3 + }, + "FURB136": { + "baseline": 1, + "slack": 3 + }, + "FURB168": { + "baseline": 1, + "slack": 3 + }, + "FURB188": { + "baseline": 49, + "slack": 3 + }, + "I001": { + "baseline": 258, + "slack": 15 + }, + "LOG015": { + "baseline": 5, + "slack": 3 + }, + "N999": { + "baseline": 1, + "slack": 3 + }, + "PERF102": { + "baseline": 27, + "slack": 3 + }, + "PERF401": { + "baseline": 136, + "slack": 10 + }, + "PERF402": { + "baseline": 6, + "slack": 3 + }, + "PERF403": { + "baseline": 69, + "slack": 5 + }, + "PIE790": { + "baseline": 263, + "slack": 15 + }, + "PIE800": { + "baseline": 1, + "slack": 3 + }, + "PIE804": { + "baseline": 21, + "slack": 3 + }, + "PIE810": { + "baseline": 41, + "slack": 3 + }, + "PLC0206": { + "baseline": 28, + "slack": 3 + }, + "PLC0208": { + "baseline": 1, + "slack": 3 + }, + "PLC0414": { + "baseline": 35, + "slack": 3 + }, + "PLR0124": { + "baseline": 1, + "slack": 3 + }, + "PLR0206": { + "baseline": 1, + "slack": 3 + }, + "PLR0402": { + "baseline": 6, + "slack": 3 + }, + "PLR0913": { + "baseline": 1813, + "slack": 50 + }, + "PLR1704": { + "baseline": 3, + "slack": 3 + }, + "PLR1711": { + "baseline": 31, + "slack": 3 + }, + "PLR1714": { + "baseline": 252, + "slack": 15 + }, + "PLR1730": { + "baseline": 7, + "slack": 3 + }, + "PLR2044": { + "baseline": 1, + "slack": 3 + }, + "PLW0127": { + "baseline": 41, + "slack": 3 + }, + "PLW0133": { + "baseline": 1, + "slack": 3 + }, + "PLW0602": { + "baseline": 215, + "slack": 15 + }, + "PLW0603": { + "baseline": 183, + "slack": 10 + }, + "PLW1508": { + "baseline": 188, + "slack": 10 + }, + "PLW1510": { + "baseline": 2, + "slack": 3 + }, + "PYI030": { + "baseline": 2, + "slack": 3 + }, + "PYI036": { + "baseline": 2, + "slack": 3 + }, + "PYI041": { + "baseline": 9, + "slack": 3 + }, + "PYI064": { + "baseline": 2, + "slack": 3 + }, + "RET501": { + "baseline": 35, + "slack": 3 + }, + "RET504": { + "baseline": 709, + "slack": 20 + }, + "RUF010": { + "baseline": 844, + "slack": 30 + }, + "RUF012": { + "baseline": 158, + "slack": 10 + }, + "RUF015": { + "baseline": 8, + "slack": 3 + }, + "RUF019": { + "baseline": 38, + "slack": 3 + }, + "RUF022": { + "baseline": 80, + "slack": 5 + }, + "RUF023": { + "baseline": 2, + "slack": 3 + }, + "RUF046": { + "baseline": 5, + "slack": 3 + }, + "RUF051": { + "baseline": 3, + "slack": 3 + }, + "RUF059": { + "baseline": 69, + "slack": 5 + }, + "RUF100": { + "baseline": 465, + "slack": 15 + }, + "S110": { + "baseline": 222, + "slack": 15 + }, + "S112": { + "baseline": 21, + "slack": 3 + }, + "SIM101": { + "baseline": 58, + "slack": 5 + }, + "SIM102": { + "baseline": 311, + "slack": 15 + }, + "SIM103": { + "baseline": 119, + "slack": 10 + }, + "SIM113": { + "baseline": 3, + "slack": 3 + }, + "SIM114": { + "baseline": 103, + "slack": 10 + }, + "SIM115": { + "baseline": 2, + "slack": 3 + }, + "SIM117": { + "baseline": 7, + "slack": 3 + }, + "SIM118": { + "baseline": 104, + "slack": 10 + }, + "SIM201": { + "baseline": 1, + "slack": 3 + }, + "SIM210": { + "baseline": 9, + "slack": 3 + }, + "SIM211": { + "baseline": 1, + "slack": 3 + }, + "SIM222": { + "baseline": 1, + "slack": 3 + }, + "SIM401": { + "baseline": 9, + "slack": 3 + }, + "TC004": { + "baseline": 5, + "slack": 3 + }, + "TC005": { + "baseline": 6, + "slack": 3 + }, + "TID251": { + "baseline": 2664, + "slack": 50 + }, + "TRY002": { + "baseline": 528, + "slack": 20 + }, + "TRY004": { + "baseline": 93, + "slack": 5 + }, + "TRY201": { + "baseline": 409, + "slack": 15 + }, + "TRY203": { + "baseline": 113, + "slack": 10 + }, + "TRY300": { + "baseline": 853, + "slack": 30 + }, + "UP006": { + "baseline": 12941, + "slack": 100 + }, + "UP007": { + "baseline": 2520, + "slack": 50 + }, + "UP008": { + "baseline": 2, + "slack": 3 + }, + "UP012": { + "baseline": 4, + "slack": 3 + }, + "UP018": { + "baseline": 18, + "slack": 3 + }, + "UP024": { + "baseline": 12, + "slack": 3 + }, + "UP028": { + "baseline": 2, + "slack": 3 + }, + "UP031": { + "baseline": 2, + "slack": 3 + }, + "UP032": { + "baseline": 609, + "slack": 20 + }, + "UP034": { + "baseline": 1, + "slack": 3 + }, + "UP035": { + "baseline": 2250, + "slack": 50 + }, + "UP036": { + "baseline": 1, + "slack": 3 + }, + "UP037": { + "baseline": 100, + "slack": 5 + }, + "UP045": { + "baseline": 18417, + "slack": 100 + } } diff --git a/ruff-strict.toml b/ruff-strict.toml index 03145255ebf..1caa3567872 100644 --- a/ruff-strict.toml +++ b/ruff-strict.toml @@ -1,7 +1,8 @@ extend = "ruff.toml" [lint] -select = ["ANN001", "ANN002", "ANN003", "ANN401", "B006", "C901", "PLR0913", "PLW0603", "RUF012", "TID251"] +preview = true +select = ["ANN", "ASYNC230", "B004", "B006", "B008", "B009", "B010", "B018", "B019", "B021", "B026", "B033", "BLE", "C401", "C404", "C405", "C408", "C414", "C419", "C901", "D419", "DTZ001", "DTZ003", "DTZ005", "DTZ006", "DTZ007", "DTZ011", "EXE001", "EXE002", "F401", "FURB136", "FURB168", "FURB188", "I001", "LOG015", "N999", "PERF102", "PERF401", "PERF402", "PERF403", "PIE790", "PIE800", "PIE804", "PIE810", "PLC0206", "PLC0208", "PLC0414", "PLR0124", "PLR0206", "PLR0402", "PLR0913", "PLR1704", "PLR1711", "PLR1714", "PLR1730", "PLR2044", "PLW0127", "PLW0133", "PLW0602", "PLW0603", "PLW1508", "PLW1510", "PYI030", "PYI036", "PYI041", "PYI064", "RET501", "RET504", "RUF010", "RUF012", "RUF015", "RUF019", "RUF022", "RUF023", "RUF046", "RUF051", "RUF059", "RUF100", "S110", "S112", "SIM101", "SIM102", "SIM103", "SIM113", "SIM114", "SIM115", "SIM117", "SIM118", "SIM201", "SIM210", "SIM211", "SIM222", "SIM401", "TC004", "TC005", "TID251", "TRY002", "TRY004", "TRY201", "TRY203", "TRY300", "UP006", "UP007", "UP008", "UP012", "UP018", "UP024", "UP028", "UP031", "UP032", "UP034", "UP035", "UP036", "UP037", "UP045"] extend-select = [] [lint.mccabe] @@ -17,4 +18,16 @@ max-args = 5 "typing.Dict".msg = "Frozen dataclass / NamedTuple / ReadOnly TypedDict; create a Mapping alias with concrete value types if truly dynamic." "typing.Set".msg = "frozenset[X] or AbstractSet[X]." "typing.MutableSequence".msg = "Sequence[X]." -"typing.MutableMapping".msg = "See typing.Dict." \ No newline at end of file +"typing.MutableMapping".msg = "See typing.Dict." +# Unchecked casts: cast() lies to the type checker with no runtime guarantee. +# Validate into a concrete frozen type at the boundary (pydantic) instead. +# Per-call-site coverage lives in check_type_discipline.py (LIT006); this freezes +# new cast imports. Suppress (with a reason) via `# noqa: TID251 # `. +"typing.cast".msg = "No unchecked casts: validate into a frozen dataclass/NamedTuple/ReadOnly TypedDict at the boundary (pydantic)." +"typing_extensions.cast".msg = "Same as typing.cast." +# Unverified narrowing predicates: the checker never validates the guard body, so a +# wrong guard silently corrupts types. Banned outright (there are none today). +"typing.TypeGuard".msg = "Unverified narrowing. Parse into a concrete type, or use isinstance for a runtime-checked narrowing." +"typing_extensions.TypeGuard".msg = "Same as typing.TypeGuard." +"typing.TypeIs".msg = "Unverified narrowing (the body is trusted). Parse into a concrete type instead." +"typing_extensions.TypeIs".msg = "Same as typing.TypeIs." diff --git a/ruff.toml b/ruff.toml index 6c854b7ad03..2db4122a30e 100644 --- a/ruff.toml +++ b/ruff.toml @@ -1,5 +1,15 @@ lint.ignore = ["F405", "E402", "E501", "F403"] -lint.extend-select = ["E501", "PLR0915", "T20"] +lint.extend-select = ["E501", "T20", "PGH004", "RUF008", "RUF009", "RUF100"] +# RUF100 (unused-noqa) only knows the rules enabled in THIS config, so it would strip +# `# noqa` directives that protect rules enforced elsewhere. List those codes as external +# so RUF100 leaves their directives alone: the strict gate (ruff-strict.toml) and upstream +# litellm's own ruff config both rely on suppressions this config can't see. +lint.external = [ + # Enforced by the strict-rule gate (scripts/ruff_strict_gate.py + ruff-strict.toml) + "C901", + # Enforced by upstream litellm's ruff config, but not run in this repo's CI + "PLC0415", "E402", "BLE001", "ARG002", "S102", "S324", "S606", "D401", "F403", "F405", +] line-length = 120 exclude = ["litellm/types/*", "litellm/__init__.py", "litellm/proxy/example_config_yaml/*", "tests/*"] @@ -13,9 +23,4 @@ exclude = ["litellm/types/*", "litellm/__init__.py", "litellm/proxy/example_conf "litellm/llms/azure_ai/embed/__init__.py" = ["F401"] "litellm/llms/azure_ai/rerank/__init__.py" = ["F401"] "litellm/llms/bedrock/chat/__init__.py" = ["F401"] -"litellm/proxy/utils.py" = ["F401", "PLR0915"] -"litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py" = ["PLR0915"] -"litellm/proxy/guardrails/guardrail_hooks/guardrail_benchmarks/test_eval.py" = ["PLR0915"] -"litellm/responses/streaming_iterator.py" = ["PLR0915"] -"litellm/files/main.py" = ["PLR0915"] -"litellm/llms/litellm_proxy/skills/sandbox_executor.py" = ["PLR0915"] +"litellm/proxy/utils.py" = ["F401"] diff --git a/scripts/budget_ratchet_check.py b/scripts/budget_ratchet_check.py new file mode 100644 index 00000000000..c4b2c3ee655 --- /dev/null +++ b/scripts/budget_ratchet_check.py @@ -0,0 +1,157 @@ +#!/usr/bin/env python3 +"""Non-gating ratchet guard: budget ceilings may only fall, never rise. + +Every `*-budget.json` file (ruff-strict, type-discipline, mypy-code, basedpyright-code) is a +one-way ratchet: each rule's ceiling is `baseline + slack`, and the whole point is +to drive that number DOWN over time. This check compares every budget file against +its own content at the merge-base with the target branch and fails (exits 1, red) if: + + * a rule's ceiling went up, + * a rule was dropped from a budget (its ceiling effectively became infinite), or + * an entire budget file was deleted. + +New rules and lowered/equal ceilings are fine. + +This is deliberately NOT a gating check. It should turn the run red so that a +loosening is impossible to miss in review, but it must stay OUT of the +branch-protection required-checks list: a justified bump (e.g. banning a new API, +which mechanically raises a baseline) can then still be merged by a human who has +seen the red and accepted it. + +Usage: + python scripts/budget_ratchet_check.py [--base REF] [budget.json ...] + +Stdlib only. +""" + +from __future__ import annotations + +import argparse +import json +import subprocess +import sys +from pathlib import Path +from typing import NamedTuple + +REPO_ROOT = Path(__file__).resolve().parent.parent +DEFAULT_BASE = "origin/litellm_internal_staging" +DEFAULT_BUDGETS: tuple[str, ...] = ( + "ruff-strict-budget.json", + "type-discipline-budget.json", + "mypy-code-budget.json", + "basedpyright-code-budget.json", +) + + +class Regression(NamedTuple): + budget: str + rule: str + detail: str + + +def _run(cmd: list[str]) -> subprocess.CompletedProcess[str]: + return subprocess.run(cmd, cwd=REPO_ROOT, capture_output=True, text=True) + + +def _merge_base(base: str) -> str: + """The common ancestor of `base` and HEAD, so unrelated base drift is ignored.""" + proc = _run(["git", "merge-base", base, "HEAD"]) + return proc.stdout.strip() or base + + +def _load_head(rel: str) -> dict | None: + path = REPO_ROOT / rel + if not path.exists(): + return None + return json.loads(path.read_text()) + + +def _ref_is_commit(ref: str) -> bool: + return _run(["git", "rev-parse", "--verify", "--quiet", f"{ref}^{{commit}}"]).returncode == 0 + + +def _load_base(rel: str, ref: str) -> dict | None: + """Budget content at `ref`, or None when the file did not exist there. + + `ref` is verified as a real commit by the caller, so a non-zero `git show` here means + the path was absent at that commit, not that the ref itself is unresolvable. + """ + proc = _run(["git", "show", f"{ref}:{rel}"]) + if proc.returncode != 0: + return None + return json.loads(proc.stdout) + + +def _caps(budget: dict) -> dict[str, int]: + """Map each rule to its ceiling (baseline + slack); skip malformed specs.""" + caps: dict[str, int] = {} + for rule, spec in budget.items(): + if isinstance(spec, dict): + caps[rule] = int(spec.get("baseline", 0)) + int(spec.get("slack", 0)) + return caps + + +def regressions_for(rel: str, base: dict | None, head: dict | None) -> list[Regression]: + if base is None: + return [] # new budget file: nothing to ratchet against yet + if head is None: + return [Regression(rel, "*", "budget file was deleted (every ceiling removed)")] + + base_caps = _caps(base) + head_caps = _caps(head) + out: list[Regression] = [] + for rule, base_cap in sorted(base_caps.items()): + if rule not in head_caps: + out.append(Regression(rel, rule, f"rule dropped (ceiling {base_cap} -> removed)")) + elif head_caps[rule] > base_cap: + out.append(Regression(rel, rule, f"ceiling raised {base_cap} -> {head_caps[rule]}")) + return out + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base", default=DEFAULT_BASE) + parser.add_argument("budgets", nargs="*", help="budget files to check") + args = parser.parse_args() + budgets = args.budgets or list(DEFAULT_BUDGETS) + + ref = _merge_base(args.base) + if not _ref_is_commit(ref): + print( + f"FAIL: base ref {ref!r} does not resolve to a commit, so the ratchet has nothing " + f"to compare against; refusing to pass vacuously (check the --base / BASE_SHA value)", + file=sys.stderr, + ) + return 1 + + regressions: list[Regression] = [] + checked: list[str] = [] + for rel in budgets: + base = _load_base(rel, ref) + head = _load_head(rel) + if base is None and head is None: + continue + if base is None: + print(f"skip {rel}: new file (no base at {args.base} to ratchet against)") + continue + checked.append(rel) + regressions.extend(regressions_for(rel, base, head)) + + if regressions: + print(f"FAIL: budget ceiling(s) loosened vs base {args.base} (merge-base {ref[:12]}):") + for reg in regressions: + print(f" {reg.budget} {reg.rule}: {reg.detail}") + print( + "Budgets are one-way ratchets and may only go down or stay flat. This " + "check is non-gating: if the increase is justified (e.g. a newly banned " + "API), a human can merge over the red after acknowledging it." + ) + return 1 + + suffix = f" ({', '.join(checked)})" if checked else "" + print(f"OK: no budget ceiling increased vs base {args.base}{suffix}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/check_any_discipline.py b/scripts/check_any_discipline.py new file mode 100644 index 00000000000..3185953d473 --- /dev/null +++ b/scripts/check_any_discipline.py @@ -0,0 +1,556 @@ +#!/usr/bin/env python3 +"""Any-discipline gate: fail when a *changed* file holds a value typed `Any`. + +Where ruff, `mypy --strict`, and even basedpyright's `reportAny` stop short, this +catches the case that actually bites: a *union* hiding an `Any`. For example +`re.Match.group()` -> `str | Any`, `json.loads()` -> `Any`, and bare `list`/`dict` +-> `list[Any]`/`dict[..., Any]`. Any value whose inferred type *contains* `Any` +(recursively, through unions / generics / tuples) is reported. + +Scope: changed-only, changed-lines +---------------------------------- +litellm already contains a large amount of pre-existing `Any` (a single legacy +file can have >100 findings), and a whole-tree scan would have to re-export types +for litellm's entire import closure on every run (~2 min, ~3 GB). So this gate is +*changed-only* and reports a finding only on a line that the diff against +`--base` actually adds or edits (untracked files count as wholly new). A brand +new file is therefore checked in full, while editing a legacy file only requires +*your* lines to be clean -- you can't introduce an `X | Any`, but you aren't +forced to clean the file's existing debt. This mirrors how `ruff_strict_gate.py` +blames a change only for the violations it introduces; cold legacy code is left +to the ratchet gates (mypy/basedpyright/ruff budgets). + +How it works +------------ +It loads `litellm/mypy.ini` (the same config `make lint-mypy` uses, so findings +match what developers already see), builds the changed files with mypy asking for +its exported expression->type map, and walks each file's AST applying a recursive +"contains Any" predicate -- the test `mypy --disallow-any-expr` uses internally +but applies inconsistently (python/mypy#12856). + +mypy only re-exports types for modules it re-type-checks, so for each target we +invalidate just its cached hash (deps stay warm) to force a fast re-check against +a persisted incremental cache (.mypy_cache_any). + +Rules +----- +Codes share the `LIT***` namespace with `scripts/check_type_discipline.py` (PR +#30500), which owns LIT001/002/003/004/006/007/008. This gate claims the rest: +LIT009 A value expression's inferred type is, or contains, `Any`. + Suppress with `# any-ok: ` on the offending line. +LIT005 An `# any-ok` suppression without a reason (the shared + suppression-needs-a-reason code, same as `# cast-ok` / `# guard-ok`). +LIT000 Setup failure: mypy could not build, or a target file could not be read. + +`Any`s produced purely by an already-reported error, and the special-form / +implementation-artifact internal `Any`s, are ignored. A bound method *reference* +whose signature mentions `Any` is not flagged -- only the value its call produces. + +Usage +----- + # gate mode (CI / pre-push): check changed lines under litellm/ + uv run --no-sync python scripts/check_any_discipline.py --changed --base origin/litellm_internal_staging + + # whole-file spot-check (no line filter), paths relative to repo root + uv run --no-sync python scripts/check_any_discipline.py litellm/budget_manager.py + +Exit code 1 if any Any-tainted value is found, 2 on a setup/usage error. +""" + +from __future__ import annotations + +import argparse +import json +import os +import re +import subprocess +import sys +import tokenize +from collections.abc import Iterable, Sequence +from pathlib import Path +from typing import NamedTuple + +try: + from mypy import build + from mypy.config_parser import parse_config_file + from mypy.find_sources import create_source_list + from mypy.fscache import FileSystemCache + from mypy.modulefinder import BuildSource + from mypy.nodes import AssignmentStmt, Expression, NameExpr, Node + from mypy.options import Options + from mypy.types import ( + AnyType, + CallableType, + Instance, + Overloaded, + TupleType, + Type, + TypeOfAny, + UnionType, + get_proper_type, + ) +except ImportError: # pragma: no cover - environment guard + sys.stderr.write( + "check_any_discipline: mypy is not importable in this interpreter.\n" + "Run it through the project environment, e.g.\n" + " uv run --no-sync python scripts/check_any_discipline.py --changed\n" + ) + raise SystemExit(2) + + +REPO_ROOT = Path(__file__).resolve().parent.parent +LITELLM_DIR = REPO_ROOT / "litellm" +MYPY_INI = LITELLM_DIR / "mypy.ini" +CACHE_DIR = REPO_ROOT / ".mypy_cache_any" +PY_TAG = f"{sys.version_info.major}.{sys.version_info.minor}" +DEFAULT_BASE = "origin/litellm_internal_staging" + +MIN_REASON_LEN = 3 +ANY_OK_RE = re.compile(r"#\s*any-ok(?::\s*(?P.*))?") +_HUNK_RE = re.compile(r"^@@ -\d+(?:,\d+)? \+(\d+)(?:,(\d+))? @@") + +# Files allowed to surface `Any` (the typed/untyped boundary). A finding is +# skipped if any fragment below is a substring of the file's posix path. Keep +# this tight -- prefer a line-level `# any-ok: ` over a blanket exemption. +BOUNDARY_PATHS: frozenset[str] = frozenset() + +# `Any` kinds that are not actionable: produced by an already-reported error, or +# an internal placeholder that never corresponds to a concrete runtime value. +# NOTE: `special_form` is deliberately NOT here. In mypy 1.19 the `Any` in +# typeshed unions like `re.Match.group() -> str | Any` is tagged `special_form`, +# and that union is the headline case this gate exists to catch. +_HARMLESS_ANY = frozenset( + kind + for kind in ( + TypeOfAny.from_error, + getattr(TypeOfAny, "implementation_artifact", None), + ) + if kind is not None +) + +# AST attributes that point OUTSIDE the syntactic subtree (a RefExpr's resolved +# definition, a node's TypeInfo). Skipping exactly these two makes a generic +# child-walk equivalent to mypy's TraverserVisitor -- validated to the node +# against ExtendedTraverserVisitor across the full grammar (see commit notes). +_NON_SYNTACTIC_ATTRS = frozenset({"node", "info"}) + + +class Violation(NamedTuple): + path: Path + line: int + col: int + code: str + message: str + + def render(self) -> str: + return f"{self.path}:{self.line}:{self.col}: {self.code} {self.message}" + + +# --------------------------------------------------------------------------- # +# The "contains Any" predicate +# --------------------------------------------------------------------------- # + + +def contains_any(t: Type, _seen: set[int] | None = None) -> bool: + """True if a *value* of type ``t`` carries `Any` anywhere meaningful.""" + seen = _seen if _seen is not None else set() + p = get_proper_type(t) + if id(p) in seen: + return False + seen.add(id(p)) + + # A function/method *reference* whose signature mentions Any is not itself an + # unsafe value -- only its eventual call result is. Don't recurse into it. + if isinstance(p, (CallableType, Overloaded)): + return False + if isinstance(p, AnyType): + return p.type_of_any not in _HARMLESS_ANY + if isinstance(p, UnionType): + return any(contains_any(item, seen) for item in p.items) + if isinstance(p, Instance): + return any(contains_any(arg, seen) for arg in p.args) + if isinstance(p, TupleType): + return any(contains_any(item, seen) for item in p.items) + return False + + +# --------------------------------------------------------------------------- # +# Generic, leak-free AST walk (works under a mypyc-compiled mypy, which forbids +# subclassing TraverserVisitor) +# --------------------------------------------------------------------------- # + + +def _walk_file(tree: Node) -> tuple[list[Expression], set[int]]: + """Return (every Expression in `tree`, ids of simple assignment-target names). + + The walk follows only syntactic children (every attribute except the two + non-syntactic back-references), so it never escapes the module. Simple + ``x = `` name targets are collected separately so we don't double-report + the assigned name as an echo of an Any rvalue. + """ + exprs: list[Expression] = [] + skip_lvalues: set[int] = set() + stack: list[object] = [tree] + seen: set[int] = set() + while stack: + n = stack.pop() + if isinstance(n, Node): + if id(n) in seen: + continue + seen.add(id(n)) + if isinstance(n, Expression): + exprs.append(n) + if isinstance(n, AssignmentStmt): + for lvalue in n.lvalues: + if isinstance(lvalue, NameExpr): + skip_lvalues.add(id(lvalue)) + for name in dir(n): + if name.startswith("__") or name in _NON_SYNTACTIC_ATTRS: + continue + try: + val = getattr(n, name) + except Exception: + continue + if callable(val): + continue + if isinstance(val, (Node, list, tuple)): + stack.append(val) + elif isinstance(n, (list, tuple)): + stack.extend(n) + return exprs, skip_lvalues + + +def find_any_in_tree(tree: Node, idmap: dict[int, Type]) -> list[tuple[int, int, str]]: + exprs, skip_lvalues = _walk_file(tree) + findings: list[tuple[int, int, str]] = [] + for expr in exprs: + if id(expr) in skip_lvalues: + continue + t = idmap.get(id(expr)) + if t is not None and contains_any(t): + findings.append((expr.line, expr.column, str(get_proper_type(t)))) + + out: list[tuple[int, int, str]] = [] + seen_pos: set[tuple[int, int]] = set() + for line, col, typ in sorted(findings): + if line < 1 or (line, col) in seen_pos: + continue + seen_pos.add((line, col)) + out.append((line, col, typ)) + return out + + +# --------------------------------------------------------------------------- # +# Comment scanning (LIT005 + any-ok suppression) +# --------------------------------------------------------------------------- # + + +def _reason_ok(reason: str | None) -> bool: + return reason is not None and len(reason.strip()) >= MIN_REASON_LEN + + +def scan_any_ok( + path: Path, source: str +) -> tuple[frozenset[int], tuple[Violation, ...]]: + """Return (lines with a valid any-ok suppression, LIT005 violations).""" + try: + tokens = tokenize.generate_tokens( + iter(source.splitlines(keepends=True)).__next__ + ) + comments = tuple( + (t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT + ) + except tokenize.TokenError: + return frozenset(), () + + ok_lines: set[int] = set() + violations: list[Violation] = [] + for line, text in comments: + m = ANY_OK_RE.search(text) + if m is None: + continue + if _reason_ok(m.group("reason")): + ok_lines.add(line) + else: + violations.append( + Violation( + path, + line, + 0, + "LIT005", + "any-ok requires a reason: `# any-ok: `", + ) + ) + return frozenset(ok_lines), tuple(violations) + + +# --------------------------------------------------------------------------- # +# mypy build (parity with `make lint-mypy`) + forced target re-check +# --------------------------------------------------------------------------- # + + +def _build_options() -> Options: + opts = Options() + if MYPY_INI.exists(): + parse_config_file(opts, lambda: None, str(MYPY_INI), sys.stdout, sys.stderr) + opts.export_types = True + opts.preserve_asts = True + opts.incremental = True + opts.cache_dir = str(CACHE_DIR) + opts.show_traceback = False + return opts + + +def _meta_path(module: str) -> Path: + return CACHE_DIR / PY_TAG / (module.replace(".", os.sep) + ".meta.json") + + +def _force_recheck(sources: Sequence[BuildSource]) -> None: + """Invalidate each target's cached entry so mypy re-type-checks (and thus + re-exports types + preserves the AST for) exactly these modules, while their + dependencies stay warm. A missing entry is a cold build for that module. + + mypy trusts a cache entry whenever the source mtime matches the cached one + (it never re-hashes on that fast path), so we must break BOTH: zero the + cached mtime to force a re-hash, and corrupt the cached hash so the re-hash + mismatches and the module is treated as changed.""" + for src in sources: + if not src.module: + continue + meta = _meta_path(src.module) + if not meta.exists(): + continue + try: + data = json.loads(meta.read_text()) + data["hash"] = "0" * 40 + data["mtime"] = 0 + meta.write_text(json.dumps(data)) + except (OSError, ValueError): + continue + + +def check_files(rel_paths: Sequence[str]) -> tuple[Violation, ...]: + """`rel_paths` are relative to the litellm package dir (the build cwd).""" + prev_cwd = Path.cwd() + os.chdir(LITELLM_DIR) + try: + opts = _build_options() + fscache = FileSystemCache() + sources = create_source_list(list(rel_paths), opts, fscache) + _force_recheck(sources) + try: + res = build.build(sources, options=opts, fscache=fscache) + except build.CompileError as exc: + joined = "; ".join(exc.messages[:3]) or "blocking error" + return ( + Violation( + Path(rel_paths[0]), + 0, + 0, + "LIT000", + f"mypy could not build: {joined}", + ), + ) + idmap = {id(expr): t for expr, t in res.types.items()} + # Resolve trees to absolute source paths while cwd is the build dir, since + # mypy stores the paths it was given (relative to this cwd). + trees: dict[str, Node] = {} + for state in res.graph.values(): + if state.path and state.tree is not None: + trees[os.path.realpath(state.path)] = state.tree + finally: + os.chdir(prev_cwd) + + out: list[Violation] = [] + for rel in rel_paths: + abs_path = (LITELLM_DIR / rel).resolve() + report_path = abs_path.relative_to(REPO_ROOT) + if _is_boundary(report_path): + continue + try: + source = abs_path.read_text(encoding="utf-8") + except (OSError, UnicodeDecodeError) as exc: + out.append( + Violation(report_path, 0, 0, "LIT000", f"could not read file: {exc}") + ) + continue + + ok_lines, ok_violations = scan_any_ok(report_path, source) + out.extend(ok_violations) + tree = trees.get(os.path.realpath(abs_path)) + if tree is None: + continue + for line, col, typ in find_any_in_tree(tree, idmap): + if line in ok_lines: + continue + out.append( + Violation( + report_path, + line, + col, + "LIT009", + f"value type contains Any -> {typ}", + ) + ) + return tuple(out) + + +# --------------------------------------------------------------------------- # +# File selection (changed-only, changed-lines) + driver +# --------------------------------------------------------------------------- # + + +class _AllLines: + """Sentinel: a wholly new / untracked file -- every line is in scope. + + A distinct object, not None, so that `line_map.get(path)` returning None for + a path absent from the map is never mistaken for "whole file in scope".""" + + +# A changed file's in-scope lines: a specific set, or every line. +LineScope = set[int] | _AllLines +ALL_LINES = _AllLines() + + +def _is_boundary(path: Path) -> bool: + posix = path.as_posix() + return any(frag in posix for frag in BOUNDARY_PATHS) + + +def _git(*args: str) -> list[str]: + result = subprocess.run( + ["git", "-C", str(REPO_ROOT), *args], + capture_output=True, + text=True, + check=True, + ) + return result.stdout.splitlines() + + +def _parse_added_lines(diff_text: str) -> dict[str, set[int]]: + """Map repo-relative path -> set of new-file line numbers the diff adds/edits.""" + changed: dict[str, set[int]] = {} + path: str | None = None + for line in diff_text.splitlines(): + if line.startswith("+++ b/"): + path = line[6:] + elif path and (m := _HUNK_RE.match(line)): + start = int(m.group(1)) + count = int(m.group(2)) if m.group(2) is not None else 1 + if count: + changed.setdefault(path, set()).update(range(start, start + count)) + return changed + + +def changed_line_map(base: str) -> dict[str, LineScope] | None: + """Repo-relative `.py` path under litellm/ -> changed line numbers (or + ALL_LINES for untracked files). Compares the working tree to the merge-base + with `base`, so it covers committed-on-branch + unstaged edits. None if git + is unavailable / not a repo.""" + try: + merge_base = _git("merge-base", base, "HEAD") + point = merge_base[0].strip() if merge_base else base + diff = "\n".join( + _git( + "diff", + "--unified=0", + "--no-color", + "--diff-filter=d", + point, + "--", + "litellm", + ) + ) + untracked = _git("ls-files", "--others", "--exclude-standard", "--", "litellm") + except (subprocess.CalledProcessError, FileNotFoundError): + return None + + out: dict[str, LineScope] = {} + for name, lines in _parse_added_lines(diff).items(): + if name.endswith(".py") and (REPO_ROOT / name).exists(): + out[name] = lines + for name in untracked: + if name.endswith(".py") and (REPO_ROOT / name).exists(): + out[name] = ALL_LINES + return out + + +def _to_litellm_relative(paths: Iterable[Path]) -> list[str]: + rels: list[str] = [] + for p in sorted(paths): + try: + rels.append(p.resolve().relative_to(LITELLM_DIR).as_posix()) + except ValueError: + continue + return rels + + +def _in_scope(v: Violation, line_map: dict[str, LineScope] | None) -> bool: + """A finding survives if line filtering is off (explicit paths), it's a build + error, or its line is one the diff added/edited.""" + if line_map is None or v.code == "LIT000": + return True + lines = line_map.get(v.path.as_posix()) + return lines is ALL_LINES or (lines is not None and v.line in lines) + + +def main(argv: Sequence[str]) -> int: + parser = argparse.ArgumentParser( + description="Any-discipline gate (changed-only, changed-lines)." + ) + parser.add_argument( + "paths", + nargs="*", + help="explicit files (repo-root relative); whole-file, no line filter", + ) + parser.add_argument( + "--changed", + action="store_true", + help="check changed lines under litellm/ vs --base", + ) + parser.add_argument("--base", default=os.environ.get("ANY_GATE_BASE", DEFAULT_BASE)) + args = parser.parse_args(list(argv)) + + line_map: dict[str, LineScope] | None = None + if args.changed: + line_map = changed_line_map(args.base) + if line_map is None: + print( + "check_any_discipline: not a git repository; nothing to check", + file=sys.stderr, + ) + return 0 + rel_paths = _to_litellm_relative( + (REPO_ROOT / name).resolve() for name in line_map + ) + elif args.paths: + rel_paths = _to_litellm_relative((REPO_ROOT / p).resolve() for p in args.paths) + else: + parser.error("pass --changed or explicit file paths") + return 2 + + if not rel_paths: + print("OK: no changed Python lines under litellm/ to check") + return 0 + + violations = tuple(v for v in check_files(rel_paths) if _in_scope(v, line_map)) + + for v in sorted(violations): + print(v.render()) + + if violations: + n = len(violations) + print( + f"\nFAIL: {n} Any-discipline violation(s) on changed lines.\n" + "Give the value a concrete type, or annotate the line `# any-ok: `.", + file=sys.stderr, + ) + return 1 + print( + f"OK: {len(rel_paths)} changed file(s) under litellm/ have no Any-typed values on changed lines" + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main(sys.argv[1:])) diff --git a/scripts/check_type_discipline.py b/scripts/check_type_discipline.py new file mode 100644 index 00000000000..d83a1a7512f --- /dev/null +++ b/scripts/check_type_discipline.py @@ -0,0 +1,476 @@ +#!/usr/bin/env python3 +"""Type-discipline checker: the rules ruff can't enforce. + +Rules +----- +LIT001 Mutable collection in a type annotation, anywhere it appears: function + parameters, return types, class attributes, locals, and module globals. + Covers the builtins (dict/list/set, bare or parameterized), their typing + aliases (Dict/List/...), the collections concretes (deque/defaultdict/...), + and the mutable ABCs (MutableMapping/MutableSequence/MutableSet). A mutable + collection lets whoever holds it grow or rewrite it after the fact; annotate + a read-only view instead (Mapping/Sequence/AbstractSet/tuple[X, ...]/ + frozenset[X], or a frozen dataclass / NamedTuple / ReadOnly TypedDict) and + build it functionally (comprehension / map, not append-in-a-loop). + Suppress with `# mutable-ok: ` on the offending line. +LIT002 Mutable-collection *construction*: a list/dict/set literal or comprehension, or + a call to a mutable constructor (list/dict/set/deque/defaultdict/Counter/...). + Catches the unannotated seed-then-mutate pattern LIT001 cannot see (`acc = []`). + Build the value in one shot and freeze it: a `tuple`/`frozenset` wrapping a + generator (`tuple(f(x) for x in xs)`), a tuple literal, or a frozen dataclass / + NamedTuple / ReadOnly TypedDict. Generator expressions and `tuple`/`frozenset` + calls are not construction and pass. Annotation-internal lists (`Callable[[int], + str]`) are exempt. Suppress with `# mutable-ok: `. +LIT003 noqa suppression without rule codes or without a reason. + Required shape: `# noqa: TID251 # ` +LIT004 type/pyright/mypy ignore without bracketed codes or without a reason. + Required shape: `# pyright: ignore[reportArgumentType] # ` +LIT005 A `# mutable-ok` / `# cast-ok` / `# guard-ok` / `# kwargs-ok` / `# any-ok` + suppression without a reason. (`any-ok` belongs to check_any_discipline.py; + it is enumerated here so the reason requirement holds even when only this + stdlib checker runs.) +LIT006 `cast(...)` call. typing.cast is an unchecked assertion (the moral equivalent + of TypeScript's `as`); it lies to the type checker with zero runtime guarantee. + Validate into a concrete frozen type at the boundary instead. + Suppress with `# cast-ok: ` on the call's first line. +LIT007 `TypeGuard[...]` / `TypeIs[...]` annotation. The narrowing predicate's body is + never verified by the checker, so a wrong guard silently corrupts types. + Prefer parsing into a concrete type. Suppress with `# guard-ok: `. +LIT008 `**kwargs` parameter. The keyword contract is erased and everything it carries + is effectively Any. ruff can force it to be typed (ANN003) but can't ban the + syntax. Declare explicit keyword params, or accept one frozen payload. `*args`, + by contrast, is fine when typed (it's just a tuple). Suppress: `# kwargs-ok: `. + +LIT000 and LIT009 are the sibling Any gate's (check_any_discipline.py, #30379): a mypy +build/read failure and an Any-typed value. They share this LIT namespace but are emitted +by that checker, not this one. + +Usage +----- + python check_type_discipline.py litellm/ tests/ + +Exit code 1 if any violation is found. Stdlib only. +""" + +from __future__ import annotations + +import ast +import io +import re +import sys +import tokenize +from dataclasses import dataclass +from pathlib import Path +from collections.abc import Iterable, Iterator, Sequence +from typing import NamedTuple + +# Mutable collection types, banned in *every* annotation. Name-based, so `dict`, +# `typing.Dict`, `collections.deque`, and `collections.abc.MutableMapping` all match +# however they were imported. The read-only interfaces (Mapping, Sequence, the +# immutable AbstractSet / `abc.Set`, Collection) and the immutable concretes (tuple, +# frozenset) are the escape hatch and are deliberately absent -- as is the bare name +# `Set`, which collides with the read-only `collections.abc.Set`. +MUTABLE_COLLECTIONS = frozenset(( + "dict", "list", "set", + "Dict", "List", "DefaultDict", "OrderedDict", "Counter", "Deque", "ChainMap", + "deque", "defaultdict", + "MutableMapping", "MutableSequence", "MutableSet", +)) + +# Callables whose result is a fresh *mutable* collection (LIT002). `tuple` and +# `frozenset` are deliberately absent -- they are the wrappers you reach for, and +# a generator expression fed to them is the blessed one-shot build. +MUTABLE_CONSTRUCTORS = frozenset(( + "dict", "list", "set", + "deque", "defaultdict", "OrderedDict", "Counter", "ChainMap", +)) +# A *qualified* call (`x.deque()`) counts as construction only for names that are rarely +# method names; `dict`/`list`/`set` are dropped here because `.dict()` / `.set()` / `.list()` +# are common methods (e.g. pydantic's `model.dict()`), not collection construction. A +# qualified `collections.deque(...)` still counts. +QUALIFIED_CONSTRUCTORS = MUTABLE_CONSTRUCTORS - frozenset(("dict", "list", "set")) +UNSAFE_GUARDS = frozenset(("TypeGuard", "TypeIs")) +MIN_REASON_LEN = 3 + +NOQA_RE = re.compile( + r"#\s*noqa" + r"(?P:\s*(?P[A-Z]+[0-9]+(?:\s*,\s*[A-Z]+[0-9]+)*))?" + r"(?P.*)", + re.IGNORECASE, +) +IGNORE_RE = re.compile( + r"#\s*(?:type|pyright|mypy):\s*ignore(?P\[[^\]]*\])?(?P.*)" +) +MUTABLE_OK_RE = re.compile(r"#\s*mutable-ok(?::\s*(?P.*))?") +CAST_OK_RE = re.compile(r"#\s*cast-ok(?::\s*(?P.*))?") +GUARD_OK_RE = re.compile(r"#\s*guard-ok(?::\s*(?P.*))?") +KWARGS_OK_RE = re.compile(r"#\s*kwargs-ok(?::\s*(?P.*))?") +ANY_OK_RE = re.compile(r"#\s*any-ok(?::\s*(?P.*))?") + +# Suppression tokens that must each carry a reason (LIT005). `any-ok` is owned by +# check_any_discipline.py but listed here so the reason requirement is enforced even +# when only this stdlib checker runs. +OK_SUPPRESSIONS: tuple[tuple[str, re.Pattern[str]], ...] = ( + ("mutable-ok", MUTABLE_OK_RE), + ("cast-ok", CAST_OK_RE), + ("guard-ok", GUARD_OK_RE), + ("kwargs-ok", KWARGS_OK_RE), + ("any-ok", ANY_OK_RE), +) + + +class Violation(NamedTuple): + path: Path + line: int + code: str + message: str + + def render(self) -> str: + return f"{self.path}:{self.line}: {self.code} {self.message}" + + +@dataclass(frozen=True, slots=True) +class Comments: + """The lines carrying each valid `*-ok` suppression.""" + + mutable_ok_lines: frozenset[int] + cast_ok_lines: frozenset[int] + guard_ok_lines: frozenset[int] + kwargs_ok_lines: frozenset[int] + + +# --------------------------------------------------------------------------- # +# Comment scanning (LIT003 / LIT004 / LIT005) +# --------------------------------------------------------------------------- # + + +def _reason_of(rest: str) -> str: + return rest.strip().lstrip("#-").strip() + + +def _valid_ok(regex: re.Pattern[str], text: str) -> bool: + """True iff `text` carries this suppression with a reason of usable length.""" + m = regex.search(text) + return bool(m) and len((m.group("reason") or "").strip()) >= MIN_REASON_LEN + + +def _comment_violations(path: Path, line_no: int, text: str) -> Iterator[Violation]: + """Pure: all LIT003/004/005 findings for one comment.""" + for token, regex in OK_SUPPRESSIONS: + m = regex.search(text) + if m and len((m.group("reason") or "").strip()) < MIN_REASON_LEN: + yield Violation(path, line_no, "LIT005", f"{token} requires a reason: `# {token}: `") + + m = NOQA_RE.search(text) + if m: + if not m.group("codes"): + yield Violation(path, line_no, "LIT003", "noqa requires rule codes: `# noqa: XXX123 # `") + elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN: + yield Violation(path, line_no, "LIT003", "noqa requires a reason: `# noqa: XXX123 # `") + + m = IGNORE_RE.search(text) + if m: + codes = m.group("codes") + if not codes or codes == "[]": + yield Violation(path, line_no, "LIT004", + "ignore requires codes: `# pyright: ignore[ruleName] # `") + elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN: + yield Violation(path, line_no, "LIT004", + "ignore requires a reason: `# pyright: ignore[ruleName] # `") + + +def scan_comments(path: Path, source: str) -> tuple[Comments, tuple[Violation, ...]]: + try: + tokens = tokenize.generate_tokens(io.StringIO(source).readline) + comment_toks = tuple((t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT) + except (tokenize.TokenError, SyntaxError): + # tokenize raises TokenError (EOF mid-construct) or a SyntaxError subclass + # (IndentationError / TabError) on malformed source; defer to ast.parse below, + # which re-raises and is reported as LIT000 rather than crashing the run. + return Comments(frozenset(), frozenset(), frozenset(), frozenset()), () + + def _lines_with(regex: re.Pattern[str]) -> frozenset[int]: + return frozenset(line for line, text in comment_toks if _valid_ok(regex, text)) + + return ( + Comments( + mutable_ok_lines=_lines_with(MUTABLE_OK_RE), + cast_ok_lines=_lines_with(CAST_OK_RE), + guard_ok_lines=_lines_with(GUARD_OK_RE), + kwargs_ok_lines=_lines_with(KWARGS_OK_RE), + ), + tuple(v for line, text in comment_toks for v in _comment_violations(path, line, text)), + ) + + +# --------------------------------------------------------------------------- # + + +def mutable_names_in(annotation: ast.expr) -> Iterator[str]: + """Yield mutable-collection names anywhere inside an annotation expression. + + Matches bare names (`dict`, `MutableMapping`) and dotted access (`typing.Dict`, + `collections.deque`, `collections.abc.MutableMapping`), descends through nesting + (`Mapping[str, list[int]]`, `tuple[set[int], ...]`) and string forward references. + """ + for node in ast.walk(annotation): + if isinstance(node, ast.Name) and node.id in MUTABLE_COLLECTIONS: + yield node.id + elif isinstance(node, ast.Attribute) and node.attr in MUTABLE_COLLECTIONS: + yield node.attr + elif isinstance(node, ast.Constant): + value: object = node.value # forward references arrive as string constants + if isinstance(value, str): + try: + inner = ast.parse(value, mode="eval").body + except SyntaxError: + continue + yield from mutable_names_in(inner) + + +def _mutable_ann(path: Path, line: int, name: str, where: str) -> Violation: + return Violation( + path, line, "LIT001", + f"mutable `{name}` in {where}: a mutable collection can be grown or rewritten " + f"by whoever holds it. Annotate a read-only view -- Mapping[...], Sequence[...], " + f"AbstractSet[...], tuple[X, ...], frozenset[X], or a frozen dataclass / " + f"NamedTuple / ReadOnly TypedDict -- and build it functionally, not by " + f"append-in-a-loop (suppress: `# mutable-ok: `)", + ) + + +def _annotation_violations( + path: Path, annotation: ast.expr | None, line: int, where: str, ok_lines: frozenset[int] +) -> Iterator[Violation]: + if annotation is None or line in ok_lines: + return + yield from (_mutable_ann(path, line, name, where) for name in mutable_names_in(annotation)) + + +def _function_violations( + path: Path, node: ast.FunctionDef | ast.AsyncFunctionDef, comments: Comments +) -> Iterator[Violation]: + mutable_ok = comments.mutable_ok_lines + args = node.args + for arg in (*args.posonlyargs, *args.args, *args.kwonlyargs): + yield from _annotation_violations( + path, arg.annotation, arg.lineno, f"parameter `{arg.arg}` of `{node.name}`", mutable_ok + ) + + # *args is allowed when typed (it's just a tuple); ruff ANN002 forces the + # annotation, so here we only add the LIT001 mutable-collection check on the element type. + if args.vararg is not None: + yield from _annotation_violations( + path, args.vararg.annotation, args.vararg.lineno, f"`*args` of `{node.name}`", mutable_ok + ) + + # **kwargs is banned outright (LIT008): it erases the keyword contract and forces + # Any-typing on everything it carries. ruff can require it be typed (ANN003) but + # cannot ban the syntax, so this rule does. + if args.kwarg is not None and args.kwarg.lineno not in comments.kwargs_ok_lines: + yield Violation( + path, args.kwarg.lineno, "LIT008", + f"`**{args.kwarg.arg}` is banned: it erases the keyword contract and forces " + f"Any-typing; declare explicit keyword parameters, or accept one frozen payload " + f"(frozen dataclass / NamedTuple / ReadOnly TypedDict) " + f"(suppress: `# kwargs-ok: `)", + ) + + if node.returns is not None: + yield from _annotation_violations( + path, node.returns, node.returns.lineno, f"return type of `{node.name}`", mutable_ok + ) + + +def iter_annotation_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: + # Every annotation is in scope: signatures (params / *args / return) plus every + # `x: T` -- class attribute, local, or module global. The latter three are all + # ast.AnnAssign, so one walk covers them; only the signature annotations (which + # are not AnnAssign) need the dedicated helper. + for node in ast.walk(tree): + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + yield from _function_violations(path, node, comments) + elif isinstance(node, ast.AnnAssign): + target = node.target.id if isinstance(node.target, ast.Name) else "" + yield from _annotation_violations( + path, node.annotation, node.lineno, + f"the type of `{target}`", comments.mutable_ok_lines, + ) + + +# --------------------------------------------------------------------------- # +# Unchecked casts (LIT006) and unverified narrowing predicates (LIT007) +# --------------------------------------------------------------------------- # + + +def _is_cast_call(node: ast.Call) -> bool: + """`cast(...)` or `typing.cast(...)`, however the name was imported/aliased. + + Name-based like MUTABLE_COLLECTIONS: a stray method called `.cast()` is a rare + false positive, suppressible with `# cast-ok: `. + """ + func = node.func + return (isinstance(func, ast.Name) and func.id == "cast") or ( + isinstance(func, ast.Attribute) and func.attr == "cast" + ) + + +def iter_cast_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: + for node in ast.walk(tree): + if isinstance(node, ast.Call) and _is_cast_call(node) and node.lineno not in comments.cast_ok_lines: + yield Violation( + path, node.lineno, "LIT006", + "cast() is an unchecked assertion (the type checker takes it on faith); " + "validate into a frozen dataclass/NamedTuple/ReadOnly TypedDict at the " + "boundary instead (suppress: `# cast-ok: `)", + ) + + +def iter_guard_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: + # TypeGuard/TypeIs are legal only as a function's return annotation (`-> TypeGuard[int]`), + # so the walk is confined to `node.returns`; a runtime name that merely happens to read + # `TypeGuard` is not a narrowing predicate. ruff bans the import; this flags the use. + for node in ast.walk(tree): + if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) or node.returns is None: + continue + for sub in ast.walk(node.returns): + name = ( + sub.id if isinstance(sub, ast.Name) + else sub.attr if isinstance(sub, ast.Attribute) + else None + ) + if name in UNSAFE_GUARDS and sub.lineno not in comments.guard_ok_lines: + yield Violation( + path, sub.lineno, "LIT007", + f"`{name}` narrowing predicate: the checker never verifies the body, so a " + f"wrong guard silently corrupts types; parse into a concrete type instead " + f"(suppress: `# guard-ok: `)", + ) + + +# --------------------------------------------------------------------------- # +# Mutable-collection construction (LIT002) +# --------------------------------------------------------------------------- # + + +def _annotations_of(node: ast.AST) -> tuple[ast.expr | None, ...]: + """The annotation expressions a node carries (signatures and `x: T`).""" + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + a = node.args + params = (*a.posonlyargs, *a.args, *a.kwonlyargs, a.vararg, a.kwarg) + return (*(p.annotation for p in params if p is not None), node.returns) + if isinstance(node, ast.AnnAssign): + return (node.annotation,) + return () + + +def _annotation_node_ids(tree: ast.AST) -> frozenset[int]: + """ids() of every node living inside an annotation. + + A list display inside an annotation (`Callable[[int], str]`) is type syntax, + not construction, so the LIT002 walk must skip those subtrees. + """ + return frozenset( + id(sub) + for node in ast.walk(tree) + for ann in _annotations_of(node) + if ann is not None + for sub in ast.walk(ann) + ) + + +def _construction_kind(node: ast.expr) -> str | None: + """Human label if `node` builds a mutable collection, else None.""" + if isinstance(node, ast.List): + return "list literal" + if isinstance(node, ast.ListComp): + return "list comprehension" + if isinstance(node, ast.Set): + return "set literal" + if isinstance(node, ast.SetComp): + return "set comprehension" + if isinstance(node, ast.Dict): + return "dict literal" + if isinstance(node, ast.DictComp): + return "dict comprehension" + if isinstance(node, ast.Call): + func = node.func + if isinstance(func, ast.Name) and func.id in MUTABLE_CONSTRUCTORS: + return f"`{func.id}()` constructor" + if isinstance(func, ast.Attribute) and func.attr in QUALIFIED_CONSTRUCTORS: + return f"`{func.attr}()` constructor" + return None + + +def iter_construction_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: + in_annotation = _annotation_node_ids(tree) + for node in ast.walk(tree): + if not isinstance(node, ast.expr) or id(node) in in_annotation: + continue + kind = _construction_kind(node) + if kind is None or node.lineno in comments.mutable_ok_lines: + continue + yield Violation( + path, node.lineno, "LIT002", + f"mutable {kind}: this builds a collection that can be grown or rewritten. " + f"Build it in one shot and freeze it -- a tuple/frozenset wrapping a generator " + f"(`tuple(f(x) for x in xs)`), a tuple literal, or a frozen dataclass / NamedTuple " + f"/ ReadOnly TypedDict (suppress: `# mutable-ok: `)", + ) + + +# --------------------------------------------------------------------------- # +# Driver +# --------------------------------------------------------------------------- # + + +def check_file(path: Path) -> tuple[Violation, ...]: + try: + source = path.read_text(encoding="utf-8") + except (OSError, UnicodeDecodeError) as exc: + return (Violation(path, 0, "LIT000", f"could not read file: {exc}"),) + + comments, violations = scan_comments(path, source) + + try: + tree = ast.parse(source, filename=str(path)) + except SyntaxError as exc: + return (*violations, Violation(path, exc.lineno or 0, "LIT000", f"syntax error: {exc.msg}")) + + return ( + *violations, + *iter_annotation_violations(path, tree, comments), + *iter_cast_violations(path, tree, comments), + *iter_guard_violations(path, tree, comments), + *iter_construction_violations(path, tree, comments), + ) + + +def collect_paths(raw: Iterable[str]) -> Iterator[Path]: + for item in raw: + p = Path(item) + if p.is_dir(): + yield from sorted(p.rglob("*.py")) + elif p.suffix == ".py": + yield p + + +def main(argv: Sequence[str]) -> int: + paths = tuple(a for a in argv if not a.startswith("-")) + if not paths: + print("usage: check_type_discipline.py ...", file=sys.stderr) + return 2 + + violations = sorted(v for path in collect_paths(paths) for v in check_file(path)) + for v in violations: + print(v.render()) + + if violations: + print(f"\n{len(violations)} violation(s).", file=sys.stderr) + return 1 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main(sys.argv[1:])) + \ No newline at end of file diff --git a/scripts/type_check_gate.py b/scripts/type_check_gate.py new file mode 100644 index 00000000000..5ff485f0b0f --- /dev/null +++ b/scripts/type_check_gate.py @@ -0,0 +1,196 @@ +#!/usr/bin/env python3 +"""Per-rule count gate for mypy and basedpyright. + +Each tool's output is reduced to a count of errors per *rule* (mypy error codes +like ``arg-type``, basedpyright rules like ``reportAny``) and checked against a +committed budget of the form ``{rule: {baseline, slack}}``, the same shape as +``ruff-strict-budget.json``. A rule fails when its codebase-wide total exceeds +``baseline + slack``. Counts ignore file, line, and column, so a violation +moving anywhere in the tree is invisible; only the per-rule total moves the +needle. + +Unlike ``ruff_strict_gate.py`` this does *not* re-run the tool on the merge base +to compute a delta: a second mypy/basedpyright pass is minutes and gigabytes, +whereas ruff is milliseconds. The committed budget is the baseline instead -- +exactly how the previous per-file gate worked -- so keep it fresh with +``--update`` (ratchet), which re-captures every rule's count from the current +tree while preserving each rule's slack. Tool output is read from stdin, so the +caller decides how to invoke the tool (and from which cwd). + +mypy is parsed from its text output (one error per line, the rule code in a +trailing ``[bracket]``). basedpyright is parsed from ``--outputjson``: its text +diagnostics routinely wrap across lines, leaving the ``(reportRule)`` on a +continuation line away from the ``- error:`` marker, so line parsing +mis-attributes ~60% of errors -- the JSON carries an unambiguous ``rule`` field. +""" + +import argparse +import json +import re +import sys +from collections import Counter +from pathlib import Path +from typing import Iterable, Mapping, NamedTuple + +REPO_ROOT = Path(__file__).resolve().parent.parent + +# mypy: one error per line, e.g. `path:12: error: msg [arg-type]`. ERROR_LINE +# recognizes the line; MYPY_CODE pulls the trailing [code]. Kept separate so an +# error emitted without a code is still counted (under UNCODED), never dropped. +MYPY_ERROR = re.compile(r"^(?P.+?):\d+: error:") +MYPY_CODE = re.compile(r"\[(?P[a-z][a-z0-9-]*)\]\s*$") + +# Bucket for an error whose rule code we couldn't read (a mypy error with no +# code, or a basedpyright diagnostic with no `rule`). Counted so it's gated. +UNCODED = "" + +# Ceiling for a rule that shows up at HEAD but isn't in the budget at all -- a +# brand-new error category (new construct, or a tool/version change). baseline +# is treated as 0, so the rule fails once it clears this much slack. +DEFAULT_SLACK = 10 + + +class Breach(NamedTuple): + code: str + total: int + cap: int + + +def _seed_slack(baseline: int) -> int: + """Slack written for a rule first captured into a budget; busy rules get + more headroom, mirroring the tiering in ruff-strict-budget.json. Existing + rules keep whatever slack their JSON already declares.""" + return 10 if baseline >= 50 else 3 + + +def _to_repo_relative(raw: str) -> str | None: + path = Path(raw) + absolute = path if path.is_absolute() else Path.cwd() / path + try: + return absolute.resolve().relative_to(REPO_ROOT).as_posix() + except ValueError: + return None + + +def count_mypy(lines: Iterable[str]) -> dict[str, int]: + """Count in-repo mypy errors per rule code from text output. Errors for + files outside the repo (third-party stubs) are ignored, as before.""" + counts: Counter[str] = Counter() + for raw in lines: + line = raw.rstrip("\n") + match = MYPY_ERROR.match(line) + if match is None or _to_repo_relative(match.group("file")) is None: + continue + code = MYPY_CODE.search(line) + counts[code.group("code") if code else UNCODED] += 1 + return dict(counts) + + +def count_basedpyright(payload: str) -> dict[str, int]: + """Count in-repo basedpyright errors per rule from `--outputjson`. Warnings + and information are ignored; only `severity == "error"` is gated.""" + try: + data = json.loads(payload or "{}") + except json.JSONDecodeError as exc: + sys.stderr.write( + f"basedpyright did not emit valid JSON ({exc}); it likely crashed or " + f"printed text before the JSON. First 500 chars of its output:\n" + f"{payload[:500]}\n" + ) + raise SystemExit(1) from exc + counts: Counter[str] = Counter() + for diag in data.get("generalDiagnostics", []): + if diag.get("severity") != "error": + continue + if _to_repo_relative(diag.get("file", "")) is None: + continue + counts[diag.get("rule") or UNCODED] += 1 + return dict(counts) + + +def count_errors(stdin_text: str, tool: str) -> dict[str, int]: + if tool == "basedpyright": + return count_basedpyright(stdin_text) + return count_mypy(stdin_text.splitlines()) + + +def evaluate( + counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]] +) -> list[Breach]: + breaches = [] + for code, total in counts.items(): + spec = budget.get(code) + cap = spec["baseline"] + spec["slack"] if spec else DEFAULT_SLACK + if total > cap: + breaches.append(Breach(code, total, cap)) + return sorted(breaches) + + +def is_vacuous_run( + counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]] +) -> bool: + """True when nothing was parsed but the budget expects errors -- the + signature of a type checker that crashed or produced no output. The CI pipe + swallows the tool's exit code (`tool || true`), so without this guard an + empty run would clear every ceiling and pass silently.""" + return not counts and any(spec["baseline"] for spec in budget.values()) + + +def budget_path(tool: str) -> Path: + return REPO_ROOT / f"{tool}-code-budget.json" + + +def cmd_update(tool: str, counts: Mapping[str, int]) -> None: + path = budget_path(tool) + existing = json.loads(path.read_text()) if path.exists() else {} + budget = { + code: { + "baseline": count, + "slack": ( + existing[code]["slack"] if code in existing else _seed_slack(count) + ), + } + for code, count in sorted(counts.items()) + } + path.write_text(json.dumps(budget, indent=2, sort_keys=True) + "\n") + print( + f"Re-captured {tool} per-rule budget: {len(budget)} rules, {sum(counts.values())} errors total" + ) + + +def cmd_check(tool: str, counts: Mapping[str, int]) -> None: + budget = json.loads(budget_path(tool).read_text()) + if is_vacuous_run(counts, budget): + expected = sum(spec["baseline"] for spec in budget.values()) + print( + f"FAIL: {tool} produced no errors, but {budget_path(tool).name} expects " + f"~{expected}. The type checker almost certainly crashed or emitted " + f"nothing; refusing to certify a vacuous run." + ) + raise SystemExit(1) + breaches = evaluate(counts, budget) + if not breaches: + print( + f"OK: every rule is within its {tool} ceiling ({sum(counts.values())} errors total)" + ) + return + print(f"FAIL: {tool} errors exceed the per-rule ceiling:") + for breach in breaches: + print(f" {breach.code}: {breach.total} errors over cap {breach.cap}") + print( + f"Resolve the new errors, or run 'make lint-{tool}-budget-update' if the ceiling should move." + ) + raise SystemExit(1) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--tool", choices=("mypy", "basedpyright"), required=True) + parser.add_argument("--update", action="store_true") + args = parser.parse_args() + counts = count_errors(sys.stdin.read(), args.tool) + cmd_update(args.tool, counts) if args.update else cmd_check(args.tool, counts) + + +if __name__ == "__main__": + main() diff --git a/scripts/type_discipline_gate.py b/scripts/type_discipline_gate.py new file mode 100644 index 00000000000..c111486e56a --- /dev/null +++ b/scripts/type_discipline_gate.py @@ -0,0 +1,198 @@ +#!/usr/bin/env python3 +"""Total-count gate for the LIT* rules in scripts/check_type_discipline.py. + +Sibling of scripts/ruff_strict_gate.py. Each rule listed in +type-discipline-budget.json has a hard ceiling (baseline + slack). The gate counts +each rule across the whole `litellm` tree and fails when a rule is both over its +ceiling and higher than the base it merges into, so a change is blamed for the +violations it adds, never for drift that already exists in the base. + +Rules not present in the budget are ignored, but today every rule the checker +emits is gated: LIT001 (mutable collection in any annotation), LIT002 +(mutable-collection construction), LIT003/LIT004 (noqa / ignore without codes or +reason), LIT006 (cast), and LIT008 (`**kwargs`) carry slack-buffered ceilings to +ratchet down; LIT005 (`*-ok` suppression without a reason) is frozen at slack 0 +so any net-new reasonless suppression trips the gate; and LIT007 (TypeGuard/TypeIs) +is a hard zero. Re-baseline with `--update` to ratchet a ceiling down. +""" + +import argparse +import json +import re +import shutil +import subprocess +import sys +import tempfile +from collections import Counter +from pathlib import Path +from typing import NamedTuple + +REPO_ROOT = Path(__file__).resolve().parent.parent +CHECKER = REPO_ROOT / "scripts" / "check_type_discipline.py" +BUDGET_PATH = REPO_ROOT / "type-discipline-budget.json" +TARGET = "litellm" +DEFAULT_BASE = "origin/litellm_internal_staging" + +_HUNK = re.compile(r"^@@ -\d+(?:,\d+)? \+(\d+)(?:,(\d+))? @@") +_LINE = re.compile(r"^(?P.+?):(?P\d+): (?PLIT\d+) ") + + +class Violation(NamedTuple): + file: str + line: int + code: str + + +class Breach(NamedTuple): + rule: str + total: int + cap: int + added: int + + +def _run(cmd: list, cwd: Path = REPO_ROOT) -> str: + proc = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True) + if proc.returncode not in (0, 1): + sys.stderr.write(proc.stderr) + raise SystemExit(f"{cmd[0]} exited {proc.returncode}") + return proc.stdout + + +def _check(root: Path, checker: Path) -> list: + # Resolve root first: on macOS tempfile dirs (/var/...) resolve to /private/var/..., + # and the checker prints already-resolved absolute paths, so relative_to would fail. + root = root.resolve() + out = _run([sys.executable, str(checker), str(root / TARGET)], cwd=root) + found = [] + for line in out.splitlines(): + m = _LINE.match(line) + if m is None: + continue + name = Path(m.group("file")) + full = name if name.is_absolute() else root / name + rel = full.resolve().relative_to(root).as_posix() + found.append(Violation(rel, int(m.group("line")), m.group("code"))) + return found + + +def head_violations() -> list: + return _check(REPO_ROOT, CHECKER) + + +def count_by_rule(violations: list) -> dict: + return dict(Counter(v.code for v in violations)) + + +def base_counts(ref: str) -> dict: + parent = Path(tempfile.mkdtemp(prefix="lit_base_")) + worktree = parent / "wt" + try: + _run(["git", "worktree", "add", "--detach", str(worktree), ref]) + # Measure the base with the *current* rule logic, not whatever shipped at base. + (worktree / "scripts").mkdir(parents=True, exist_ok=True) + checker = worktree / "scripts" / "check_type_discipline.py" + shutil.copy(CHECKER, checker) + return count_by_rule(_check(worktree, checker)) + finally: + # Best-effort teardown: cleanup must never raise, or it masks the real error when + # the body (or the `worktree add` itself) failed. rmtree is already best-effort. + subprocess.run( + ["git", "worktree", "remove", "--force", str(worktree)], + cwd=REPO_ROOT, capture_output=True, text=True, + ) + shutil.rmtree(parent, ignore_errors=True) + + +def over_ceiling(head: dict, budget: dict) -> frozenset: + """Rules whose head count already exceeds baseline + slack. + + A rule can only breach when it is over its ceiling, so when none are the base + comparison cannot change the verdict and the base worktree scan can be skipped. + """ + return frozenset( + rule for rule, spec in budget.items() + if head.get(rule, 0) > spec["baseline"] + spec["slack"] + ) + + +def evaluate(head: dict, base: dict, budget: dict) -> list: + breaches = [] + for rule, spec in budget.items(): + cap = spec["baseline"] + spec["slack"] + total = head.get(rule, 0) + if total > cap and total > base.get(rule, 0): + breaches.append(Breach(rule, total, cap, total - base.get(rule, 0))) + return sorted(breaches) + + +def parse_changed_lines(diff_text: str) -> dict: + changed: dict = {} + path = None + for line in diff_text.splitlines(): + if line.startswith("+++ b/"): + path = line[6:] + elif path and (match := _HUNK.match(line)): + start = int(match.group(1)) + count = int(match.group(2)) if match.group(2) is not None else 1 + changed.setdefault(path, set()).update(range(start, start + count)) + return changed + + +def introduced(violations: list, changed: dict) -> list: + return [v for v in violations if v.line in changed.get(v.file, set())] + + +def cmd_check(base: str) -> None: + budget = json.loads(BUDGET_PATH.read_text()) + head = head_violations() + head_counts = count_by_rule(head) + if not over_ceiling(head_counts, budget): + print(f"OK: every LIT rule is within its codebase ceiling (base {base})") + return + base_point = _run(["git", "merge-base", base, "HEAD"]).strip() or base + breaches = evaluate(head_counts, base_counts(base_point), budget) + if not breaches: + print(f"OK: every LIT rule is within its codebase ceiling (base {base})") + return + new = introduced( + head, + parse_changed_lines( + _run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET]) + ), + ) + print(f"FAIL: LIT-rule totals exceed their ceiling (base {base}):") + for breach in breaches: + print( + f" {breach.rule}: total {breach.total} over cap {breach.cap} (this change added {breach.added})" + ) + for violation in sorted(v for v in new if v.code == breach.rule): + print(f" {violation.file}:{violation.line}") + print( + "Remove the new violations, give each a reason (`# noqa: XXX # `, " + "`# pyright: ignore[rule] # `, `# mutable-ok: `, " + "`# cast-ok: `, `# guard-ok: `, `# kwargs-ok: `), or " + "remove an equal number elsewhere; the ceiling is baseline + slack in " + "type-discipline-budget.json." + ) + raise SystemExit(1) + + +def cmd_update() -> None: + budget = json.loads(BUDGET_PATH.read_text()) + head = count_by_rule(head_violations()) + for rule in budget: + budget[rule]["baseline"] = head.get(rule, 0) + BUDGET_PATH.write_text(json.dumps(budget, indent=2, sort_keys=True) + "\n") + print("Re-captured per-rule baselines from the current tree") + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base", default=DEFAULT_BASE) + parser.add_argument("--update", action="store_true") + args = parser.parse_args() + cmd_update() if args.update else cmd_check(args.base) + + +if __name__ == "__main__": + main() diff --git a/tests/code_coverage_tests/enforce_llms_folder_style.py b/tests/code_coverage_tests/enforce_llms_folder_style.py index 43ab81b6c60..cbf5cd5266e 100644 --- a/tests/code_coverage_tests/enforce_llms_folder_style.py +++ b/tests/code_coverage_tests/enforce_llms_folder_style.py @@ -14,6 +14,7 @@ SEARCH_PROVIDERS = [ "exa_ai", "brave", "firecrawl", + "fastcrw", "searxng", "linkup", "duckduckgo", diff --git a/tests/litellm_utils_tests/test_logging_callback_manager.py b/tests/litellm_utils_tests/test_logging_callback_manager.py index d9540f8f850..d9bfca425e4 100644 --- a/tests/litellm_utils_tests/test_logging_callback_manager.py +++ b/tests/litellm_utils_tests/test_logging_callback_manager.py @@ -192,6 +192,29 @@ def test_remove_callback_from_list_by_object(): assert len(litellm._async_failure_callback) == 0 +def test_remove_callback_from_all_lists(): + manager = LoggingCallbackManager() + manager._reset_all_callbacks() + + class TestLogger(CustomLogger): + pass + + obj = TestLogger() + manager.add_litellm_callback(obj) + manager.add_litellm_success_callback(obj) + manager.add_litellm_failure_callback(obj) + manager.add_litellm_async_success_callback(obj) + manager.add_litellm_async_failure_callback(obj) + + manager.remove_callback_from_all_lists(obj) + + assert obj not in litellm.callbacks + assert obj not in litellm.success_callback + assert obj not in litellm.failure_callback + assert obj not in litellm._async_success_callback + assert obj not in litellm._async_failure_callback + + def test_reset_callbacks(callback_manager): # Add various callbacks callback_manager.add_litellm_callback("test") diff --git a/tests/test_anthropic_compaction_usage.py b/tests/test_anthropic_compaction_usage.py new file mode 100644 index 00000000000..1758a94fffc --- /dev/null +++ b/tests/test_anthropic_compaction_usage.py @@ -0,0 +1,96 @@ +from litellm.llms.anthropic.chat.transformation import AnthropicConfig + + +def test_anthropic_compaction_usage_calculation(): + """ + Test that calculate_usage correctly sums tokens from the iterations array + as requested in Issue #27060. + """ + anthropic_config = AnthropicConfig() + + # Mock usage object with compaction iterations + usage_object = { + "input_tokens": 100, # Top-level (excludes compaction) + "output_tokens": 50, # Top-level (excludes compaction) + "iterations": [ + { + "iteration": 1, + "type": "compaction", + "input_tokens": 1000, + "output_tokens": 500, + }, + { + "iteration": 2, + "type": "message", + "input_tokens": 100, + "output_tokens": 50, + }, + ], + } + + usage = anthropic_config.calculate_usage( + usage_object=usage_object, reasoning_content=None + ) + + # Assertions + # Total prompt tokens should be 1000 + 100 = 1100 + assert usage.prompt_tokens == 1100 + # Total completion tokens should be 500 + 50 = 550 + assert usage.completion_tokens == 550 + # Total tokens should be 1650 + assert usage.total_tokens == 1650 + + # Assert details + assert usage.prompt_tokens_details.text_tokens == 1100 + + # Assert iterations passthrough + assert usage.iterations is not None + assert len(usage.iterations) == 2 + assert usage.iterations[0]["type"] == "compaction" + + +def test_anthropic_compaction_usage_with_iteration_cache(): + """ + Test that calculate_usage correctly sums caching tokens FROM iterations. + This covers the specific case mentioned by JasonPan. + """ + anthropic_config = AnthropicConfig() + + usage_object = { + "input_tokens": 100, + "output_tokens": 50, + "iterations": [ + { + "type": "compaction", + "input_tokens": 500, + "output_tokens": 200, + "cache_creation_input_tokens": 50, + "cache_read_input_tokens": 17000, + }, + { + "type": "message", + "input_tokens": 100, + "output_tokens": 50, + "cache_creation_input_tokens": 10, + "cache_read_input_tokens": 20, + }, + ], + } + + usage = anthropic_config.calculate_usage( + usage_object=usage_object, reasoning_content=None + ) + + # input_tokens sum = 500 + 100 = 600 + # cache_creation sum = 50 + 10 = 60 + # cache_read sum = 17000 + 20 = 17020 + # Total prompt tokens = 600 + 60 + 17020 = 17680 + assert usage.prompt_tokens == 17680 + assert usage.completion_tokens == 250 + assert usage.prompt_tokens_details.cache_creation_tokens == 60 + assert usage.prompt_tokens_details.cached_tokens == 17020 + + +if __name__ == "__main__": + test_anthropic_compaction_usage_calculation() + test_anthropic_compaction_usage_with_iteration_cache() diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 06457dfebff..6a1de0586dd 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -2819,3 +2819,37 @@ def test_reasoning_items_streaming_emitted_on_response_completed(): ri["encrypted_content"] == encrypted ), "encrypted_content must be preserved in streaming" assert ri["summary"][0]["text"] == summary_text + + +def test_streaming_function_call_tool_id_for_degenerate_call_id(): + """In streaming, Bedrock Mantle's function_call event carries a unique ``id`` + (``fc_...``) and a non-unique, index-based ``call_id`` (``call_0``). For that + degenerate form the chat tool-call chunk must use the unique ``id`` so multi-turn + streaming agents don't collapse every tool call to the same id (which makes the + agent loop). A normal (unique) ``call_id`` must be preserved. Regression for the + bedrock-mantle gpt-5.5 streaming path.""" + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + OpenAiResponsesToChatCompletionStreamIterator, + ) + + def stream_tool_id(item_id, call_id): + chunk = { + "type": "response.output_item.added", + "output_index": 0, + "item": { + "type": "function_call", + "id": item_id, + "call_id": call_id, + "name": "get_weather", + "arguments": "", + }, + } + out = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream( + chunk + ) + tool_calls = out.model_dump()["choices"][0]["delta"]["tool_calls"] + assert tool_calls, "expected a tool_call chunk in the streaming delta" + return tool_calls[0]["id"] + + assert stream_tool_id("fc_unique_abc123", "call_0") == "fc_unique_abc123" + assert stream_tool_id("fc_2", "call_tokyo") == "call_tokyo" 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 19f4b0ff457..19eef284b91 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_components.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_components.py @@ -2,6 +2,8 @@ baggage helpers, metrics, the typed coercion helpers, mapper branches, span-name builders, and the registry validator's failure paths. Needs the OTel SDK.""" +import json + import pytest pytest.importorskip("opentelemetry") @@ -225,6 +227,47 @@ def test_genai_mapper_all_request_params(): assert attrs["server.port"] == 443 +def test_genai_mapper_stamps_input_output_messages(): + data = LLMCallSpanData( + operation=GenAIOperation.CHAT, + provider="openai", + request_model="gpt-4o", + response_model="gpt-4o-2024", + response_id="resp_1", + request_params=LLMRequestParams(), + usage=LLMUsage(), + finish_reasons=("stop",), + error=None, + response_cost=None, + server=None, + identity=RequestIdentity(call_id="c1"), + messages_in=( + {"role": "system", "content": "Be concise."}, + {"role": "user", "content": "What's the weather?"}, + ), + choices_out=( + { + "finish_reason": "stop", + "message": {"role": "assistant", "content": "Sunny."}, + }, + ), + ) + attrs = GenAIMapper().map(data) + assert json.loads(attrs[GenAI.INPUT_MESSAGES]) == [ + {"role": "system", "content": "Be concise."}, + {"role": "user", "content": "What's the weather?"}, + ] + assert json.loads(attrs[GenAI.OUTPUT_MESSAGES]) == [ + {"role": "assistant", "content": "Sunny."} + ] + + +def test_genai_mapper_omits_messages_when_content_not_captured(): + attrs = GenAIMapper().map(_full_llm_call()) + assert GenAI.INPUT_MESSAGES not in attrs + assert GenAI.OUTPUT_MESSAGES not in attrs + + def test_genai_mapper_cost_breakdown(): from litellm.integrations.otel.model.semconv import LiteLLM 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 8dffb71bbf0..77ee4d0a5a9 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py @@ -8,7 +8,7 @@ hooks, proxy SERVER span lifecycle (start + setters), parent-context resolution import asyncio import contextlib -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone import pytest @@ -1314,3 +1314,178 @@ def test_module_level_emit_guardrail_span_swallows_emit_errors(monkeypatch): monkeypatch.setattr(otel_logger, "_registered_v2_logger", lambda: _Boom()) otel_logger.emit_guardrail_span(_guardrail_entry(start=1.0, end=2.0)) + + +# --------------------------------------------------------------------------- # +# Metrics: invalid attribute-filter config is visible, not a silent no-op +# --------------------------------------------------------------------------- # + + +def _emitted_metric_names(reader) -> set: + data = reader.get_metrics_data() + if data is None: + return set() + return { + m.name + for rm in data.resource_metrics + for sm in rm.scope_metrics + for m in sm.metrics + if any(m.data.data_points) + } + + +def _metric_success_kwargs() -> dict: + return { + "model": "gpt-4o-mini", + "call_type": "acompletion", + "litellm_params": {"custom_llm_provider": "openai"}, + "optional_params": {}, + "response_cost": 0.001, + "standard_logging_object": {"metadata": {}, "hidden_params": {}}, + } + + +def test_invalid_metric_filter_logged_once_records_nothing(caplog, monkeypatch): + """An invalid ``callback_settings.otel.attributes`` (include_list + exclude_list + both set) must make the operator-fixable config error visible once at ERROR and + record no metrics — without raising out of the success path and without + per-request log spam. Mirrors the v1 fix against the silent-no-op failure mode. + """ + import logging + + from opentelemetry.sdk.metrics import MeterProvider + from opentelemetry.sdk.metrics.export import InMemoryMetricReader + + import litellm + + monkeypatch.setattr( + litellm, + "callback_settings", + { + "otel": { + "attributes": { + "include_list": ["gen_ai.system"], + "exclude_list": ["hidden_params"], + } + } + }, + raising=False, + ) + + cfg = OpenTelemetryV2Config(exporter="in_memory", enable_metrics=True) + reader = InMemoryMetricReader() + logger = OpenTelemetryV2( + config=cfg, + callback_name="otel", + tracer_provider=providers.build_tracer_provider(cfg), + meter_provider=MeterProvider(metric_readers=[reader]), + ) + + start = datetime.now(timezone.utc) + end = start + timedelta(seconds=1) + response_obj = {"usage": {"prompt_tokens": 1, "completion_tokens": 1}} + + with caplog.at_level(logging.ERROR, logger="LiteLLM"): + # Neither call may raise; the bad filter is caught in the logger. + asyncio.run( + logger.async_log_success_event( + _metric_success_kwargs(), response_obj, start, end + ) + ) + asyncio.run( + logger.async_log_success_event( + _metric_success_kwargs(), response_obj, start, end + ) + ) + + assert _emitted_metric_names(reader) == set() # nothing recorded + errors = [ + r + for r in caplog.records + if r.levelno == logging.ERROR and "metric filter" in r.getMessage() + ] + assert len(errors) == 1 # logged once, second bad record does not re-log + + +def test_valid_metric_filter_records_six_metrics(monkeypatch): + """The happy path: with no attribute filter, a successful LLM call records all + six GenAI histograms, and the token metric keeps its input/output split.""" + from opentelemetry.sdk.metrics import MeterProvider + from opentelemetry.sdk.metrics.export import InMemoryMetricReader + + import litellm + + monkeypatch.setattr(litellm, "callback_settings", {}, raising=False) + + cfg = OpenTelemetryV2Config(exporter="in_memory", enable_metrics=True) + reader = InMemoryMetricReader() + logger = OpenTelemetryV2( + config=cfg, + callback_name="otel", + tracer_provider=providers.build_tracer_provider(cfg), + meter_provider=MeterProvider(metric_readers=[reader]), + ) + + start = datetime.now(timezone.utc) + end = start + timedelta(seconds=2) + kwargs = _metric_success_kwargs() + kwargs["api_call_start_time"] = start.timestamp() + kwargs["completion_start_time"] = (start + timedelta(seconds=0.5)).timestamp() + kwargs["end_time"] = end.timestamp() + kwargs["optional_params"] = {"stream": True} + response_obj = {"usage": {"prompt_tokens": 5, "completion_tokens": 7}} + + asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) + + 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.client.response.duration", + } + + data = reader.get_metrics_data() + token_types = { + dp.attributes.get("gen_ai.token.type") + for rm in data.resource_metrics + for sm in rm.scope_metrics + for m in sm.metrics + if m.name == "gen_ai.client.token.usage" + for dp in m.data.data_points + } + assert token_types == {"input", "output"} + + +def test_metrics_disabled_by_default_records_nothing(monkeypatch): + """With ``enable_metrics`` off (the default), no meter is built and a success + event records nothing — the default behavior must stay unchanged.""" + from opentelemetry.sdk.metrics import MeterProvider + from opentelemetry.sdk.metrics.export import InMemoryMetricReader + + import litellm + + monkeypatch.setattr(litellm, "callback_settings", {}, raising=False) + + cfg = OpenTelemetryV2Config(exporter="in_memory") # enable_metrics defaults False + reader = InMemoryMetricReader() + logger = OpenTelemetryV2( + config=cfg, + callback_name="otel", + tracer_provider=providers.build_tracer_provider(cfg), + meter_provider=MeterProvider(metric_readers=[reader]), + ) + assert logger._metrics_recorder is None + + start = datetime.now(timezone.utc) + end = start + timedelta(seconds=1) + asyncio.run( + logger.async_log_success_event( + _metric_success_kwargs(), + {"usage": {"prompt_tokens": 1, "completion_tokens": 1}}, + start, + end, + ) + ) + assert _emitted_metric_names(reader) == set() diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py new file mode 100644 index 00000000000..29067f91b5a --- /dev/null +++ b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py @@ -0,0 +1,333 @@ +"""Tests for the V2 OTEL GenAI client metrics. + +Drives the real success path: the six ``gen_ai.client.*`` histograms are emitted +through ``OpenTelemetryV2.async_log_success_event`` into an injected +``InMemoryMetricReader``, and attributes/values are read straight off the +recorded data points (``resource_metrics`` -> ``scope_metrics`` -> ``metrics`` -> +``data.data_points``). The cardinality filter is resolved lazily from +``litellm.callback_settings['otel']['attributes']``, which the proxy populates +after the logger is built, so those tests set it AFTER construction. A +misconfigured filter (``gen_ai.token.type`` in a list, include+exclude together) +raises out of ``GenAIMetricRecorder.record`` -- asserted directly at the recorder +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. +""" + +import asyncio +from datetime import datetime, timedelta + +import pytest + +pytest.importorskip("opentelemetry") + +from opentelemetry.sdk.metrics import MeterProvider # noqa: E402 +from opentelemetry.sdk.metrics.export import InMemoryMetricReader # noqa: E402 + +import litellm # noqa: E402 +from litellm.integrations.otel.logger import OpenTelemetryV2 # noqa: E402 +from litellm.integrations.otel.model.config import ( # noqa: E402 + OpenTelemetryV2Config, +) +from litellm.integrations.otel.plumbing.metrics import ( # noqa: E402 + GenAIMetricRecorder, + create_genai_metrics, +) +from litellm.integrations.otel.plumbing.providers import ( # noqa: E402 + resolve_meter_provider, +) + +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" +RESPONSE_DURATION = "gen_ai.client.response.duration" + +ALL_METRICS = frozenset( + { + OPERATION_DURATION, + TOKEN_USAGE, + TOKEN_COST, + TIME_TO_FIRST_TOKEN, + TIME_PER_OUTPUT_TOKEN, + RESPONSE_DURATION, + } +) + +TOKEN_TYPE = "gen_ai.token.type" +MODEL_KEY = "gen_ai.request.model" + +# 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 = ( + "hidden_params", + "metadata.user_api_key_hash", + "metadata.requester_ip_address", + "metadata.requester_metadata", + "metadata.applied_guardrails", +) + +PROMPT_TOKENS = 137 +COMPLETION_TOKENS = 89 +RESPONSE_COST = 0.0023 + + +def _build_call(stream: bool = True): + """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.""" + 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"}, + "optional_params": {"stream": stream}, + "response_cost": RESPONSE_COST, + "api_call_start_time": api_call_start, + "completion_start_time": completion_start, + "end_time": end, + "standard_logging_object": { + "metadata": { + "user_api_key_hash": "hash-abc123", + "requester_ip_address": "10.0.0.7", + "requester_metadata": {"team": "alpha", "tier": "gold"}, + "applied_guardrails": ["pii", "toxicity"], + }, + "hidden_params": {"litellm_call_id": "abc", "model_id": "m-1"}, + }, + } + response_obj = { + "usage": { + "prompt_tokens": PROMPT_TOKENS, + "completion_tokens": COMPLETION_TOKENS, + } + } + return kwargs, response_obj, start, end + + +def _logger(reader, *, enable_metrics: bool): + return OpenTelemetryV2( + config=OpenTelemetryV2Config( + exporter="in_memory", enable_metrics=enable_metrics + ), + meter_provider=MeterProvider(metric_readers=[reader]), + ) + + +def _metrics_by_name(reader): + """{metric_name: [data_point, ...]} from everything the reader has collected.""" + data = reader.get_metrics_data() + out: dict = {} + if not data or not getattr(data, "resource_metrics", None): + return out + for rm in data.resource_metrics: + for sm in rm.scope_metrics: + for m in sm.metrics: + out.setdefault(m.name, []).extend(m.data.data_points) + return out + + +def _drive_success(reader, callback_settings_attributes=None): + """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) + previous = litellm.callback_settings + if callback_settings_attributes is not None: + litellm.callback_settings = { + "otel": {"attributes": callback_settings_attributes} + } + try: + kwargs, response_obj, start, end = _build_call() + asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) + finally: + litellm.callback_settings = previous + return _metrics_by_name(reader) + + +def test_all_six_metrics_emitted_when_enabled(): + """A successful streaming call with metrics on emits exactly the six + gen_ai.client.* histograms, and token.usage splits into an input and an + output point carrying the right token counts.""" + metrics = _drive_success(InMemoryMetricReader()) + + assert set(metrics.keys()) == set(ALL_METRICS) + + token_points = metrics[TOKEN_USAGE] + by_type = {dp.attributes[TOKEN_TYPE]: dp for dp in token_points} + assert set(by_type) == {"input", "output"} + assert by_type["input"].sum == PROMPT_TOKENS + assert by_type["output"].sum == COMPLETION_TOKENS + + cost_points = metrics[TOKEN_COST] + assert len(cost_points) == 1 + assert cost_points[0].sum == pytest.approx(RESPONSE_COST) + + +def test_time_to_first_token_is_streaming_only(): + """time_to_first_token is gated on streaming: a non-streaming call emits the + other five metrics but never that one.""" + reader = InMemoryMetricReader() + logger = _logger(reader, enable_metrics=True) + kwargs, response_obj, start, end = _build_call(stream=False) + asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) + + names = set(_metrics_by_name(reader).keys()) + assert TIME_TO_FIRST_TOKEN not in names + assert names == set(ALL_METRICS) - {TIME_TO_FIRST_TOKEN} + + +def test_metrics_disabled_records_nothing(): + """enable_metrics=False: the recorder is never built, so the injected reader + sees no gen_ai.client.* series even though the success hook runs.""" + reader = InMemoryMetricReader() + logger = _logger(reader, enable_metrics=False) + kwargs, response_obj, start, end = _build_call() + asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) + + assert set(_metrics_by_name(reader).keys()).isdisjoint(ALL_METRICS) + + +def test_metrics_off_by_default_records_nothing(): + """The default config has metrics off, so a default logger records nothing.""" + reader = InMemoryMetricReader() + logger = OpenTelemetryV2( + config=OpenTelemetryV2Config(exporter="in_memory"), + meter_provider=MeterProvider(metric_readers=[reader]), + ) + kwargs, response_obj, start, end = _build_call() + asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) + + assert set(_metrics_by_name(reader).keys()).isdisjoint(ALL_METRICS) + + +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.""" + metrics = _drive_success( + InMemoryMetricReader(), + callback_settings_attributes={"exclude_list": list(HIGH_CARDINALITY_KEYS)}, + ) + excluded = set(HIGH_CARDINALITY_KEYS) + + for name in (OPERATION_DURATION, TOKEN_USAGE): + points = metrics[name] + assert points, f"{name} was not recorded" + for dp in points: + keys = set(dp.attributes.keys()) + assert excluded.isdisjoint(keys), f"{name} leaked {excluded & keys}" + assert MODEL_KEY in keys + + +def test_include_list_allows_only_listed_attributes(): + """include_list caps emitted attributes to exactly the listed set; + gen_ai.token.type is the only key permitted beyond it, and only on the + token-usage metric.""" + include = [MODEL_KEY, "gen_ai.system"] + metrics = _drive_success( + InMemoryMetricReader(), + callback_settings_attributes={"include_list": include}, + ) + allowed = set(include) + + for dp in metrics[OPERATION_DURATION]: + assert set(dp.attributes.keys()) == allowed + + for dp in metrics[TOKEN_USAGE]: + 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.""" + metrics = _drive_success(InMemoryMetricReader()) + expected = set(HIGH_CARDINALITY_KEYS) + + for name in (OPERATION_DURATION, TOKEN_USAGE): + for dp in metrics[name]: + assert expected.issubset(set(dp.attributes.keys())) + + +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 + MeterProvider so its readers/exporters receive them. Before the fix the logger + built an isolated provider and the operator's reader saw nothing.""" + from opentelemetry import metrics + + reader = InMemoryMetricReader() + operator_provider = MeterProvider(metric_readers=[reader]) + monkeypatch.setattr(metrics, "get_meter_provider", lambda: operator_provider) + + logger = OpenTelemetryV2( + config=OpenTelemetryV2Config(exporter="in_memory", enable_metrics=True), + ) + kwargs, response_obj, start, end = _build_call() + asyncio.run(logger.async_log_success_event(kwargs, response_obj, start, end)) + + assert set(_metrics_by_name(reader).keys()) == set(ALL_METRICS) + operator_provider.shutdown() + + +def test_resolve_meter_provider_prefers_injected(): + """An injected provider is used verbatim, never replaced by the global.""" + injected = MeterProvider(metric_readers=[InMemoryMetricReader()]) + resolved = resolve_meter_provider( + OpenTelemetryV2Config(exporter="in_memory"), injected + ) + assert resolved is injected + injected.shutdown() + + +def test_resolve_meter_provider_honors_operator_noop(monkeypatch): + """An operator that disabled metrics with a NoOpMeterProvider is not silently + overridden by a freshly built provider.""" + from opentelemetry import metrics + from opentelemetry.metrics import NoOpMeterProvider + + noop = NoOpMeterProvider() + monkeypatch.setattr(metrics, "get_meter_provider", lambda: noop) + + resolved = resolve_meter_provider(OpenTelemetryV2Config(exporter="in_memory")) + assert resolved is noop + + +def _recorder(monkeypatch, attributes): + """A recorder wired to a fresh in-memory meter, with callback_settings carrying + `attributes`. record() resolves the filter lazily from there, so a misconfig + raises out of record() at this layer (the logger turns it into log-once).""" + monkeypatch.setattr( + litellm, + "callback_settings", + {"otel": {"attributes": attributes}}, + raising=False, + ) + meter = MeterProvider(metric_readers=[InMemoryMetricReader()]).get_meter("test") + return GenAIMetricRecorder(create_genai_metrics(meter), callback_name=None) + + +@pytest.mark.parametrize( + "attributes", + [ + {"exclude_list": [TOKEN_TYPE]}, + {"include_list": [TOKEN_TYPE]}, + ], +) +def test_token_type_rejected_from_either_list(attributes, monkeypatch): + """gen_ai.token.type is a structural discriminator stamped onto the + input/output series after filtering; it cannot itself be filtered without + collapsing the two series. Listing it in either list is rejected by the + recorder rather than silently ignored, so the misconfig is caught at all.""" + recorder = _recorder(monkeypatch, attributes) + kwargs, response_obj, start, end = _build_call() + with pytest.raises(ValueError) as exc_info: + recorder.record(kwargs, response_obj, start, end) + # The dedicated discriminator guard, not the generic unknown-name path: assert + # the specific reason so dropping that guard (and falling through to "unknown + # attribute name") is caught. + assert "discriminator" in str(exc_info.value) 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 20824ca09e6..3447f5bdb7e 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 @@ -528,6 +528,35 @@ def test_capture_span_content_resolves_modes(): ).capture_span_content is False ) + # V1 accepted UPPER_SNAKE_CASE; the env value is case-insensitive so an + # operator carrying ``SPAN_AND_EVENT`` forward still enables capture. + assert ( + OpenTelemetryV2Config( + capture_message_content="SPAN_AND_EVENT" + ).capture_span_content + is True + ) + assert ( + OpenTelemetryV2Config(capture_message_content="SPAN_ONLY").capture_span_content + is True + ) + assert ( + OpenTelemetryV2Config(capture_message_content="NO_CONTENT").capture_span_content + is False + ) + + +def test_capture_message_content_normalizer_only_touches_strings(): + """The casing normalizer lower-cases strings and leaves anything else + untouched, so a non-string value still fails the field's ``str`` validation + instead of being silently coerced into a bogus capture mode.""" + import pytest + from pydantic import ValidationError + + from litellm.integrations.otel.model.config import OpenTelemetryV2Config + + with pytest.raises(ValidationError): + OpenTelemetryV2Config(capture_message_content=123) def test_v2_flag_is_off_by_default(monkeypatch): diff --git a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py index 1a4d03528e7..6afe5efc54d 100644 --- a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py +++ b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py @@ -1087,3 +1087,357 @@ async def test_anthropic_cache_control_hook_string_negative_index(): f"Expected cachePoint in last message content, got: {last_message_content}. " "String index '-1' was not parsed correctly (str.isdigit() returns False for negative strings)." ) + + +def _count_cache_control(messages: List[AllMessageValues]) -> int: + """Count cache_control breakpoints across messages (message + content level).""" + count = 0 + for message in messages: + if message.get("cache_control") is not None: + count += 1 + content = message.get("content") + if isinstance(content, list): + for block in content: + if isinstance(block, dict) and block.get("cache_control") is not None: + count += 1 + return count + + +def _build_injection_points(): + return [ + { + "location": "message", + "role": "system", + "control": {"type": "ephemeral", "ttl": "1h"}, + }, + { + "location": "message", + "index": -1, + "control": {"type": "ephemeral", "ttl": "5m"}, + }, + ] + + +def test_cache_control_hook_caps_at_four_blocks_with_client_cache_control(): + """Regression for LIT-3667 / Anthropic 'A maximum of 4 blocks ... Found 5'. + + A Hermes-style request already carries 4 client cache_control breakpoints on + its system messages. With both auto-inject points configured the hook must + NOT add a 5th breakpoint, and must NOT overwrite the client's existing + breakpoints (TTL must be preserved). + """ + hook = AnthropicCacheControlHook() + + messages: List[AllMessageValues] = [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": f"System block {i}", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ], + } + for i in range(4) + ] + messages.append({"role": "user", "content": "hello"}) + + _, processed, _ = hook.get_chat_completion_prompt( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + non_default_params={ + "cache_control_injection_points": _build_injection_points() + }, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + + assert ( + _count_cache_control(processed) == 4 + ), "Hook must cap cache_control at Anthropic's limit of 4 blocks" + + # Client TTL on system blocks must be preserved (not overwritten by config). + for i in range(4): + assert processed[i]["content"][-1]["cache_control"] == { + "type": "ephemeral", + "ttl": "1h", + } + + # The last (user) message must not receive a 5th breakpoint. + user_message = processed[-1] + assert user_message.get("cache_control") is None + user_content = user_message.get("content") + if isinstance(user_content, list): + assert all( + block.get("cache_control") is None + for block in user_content + if isinstance(block, dict) + ) + + +def test_cache_control_hook_caps_at_four_blocks_without_client_cache_control(): + """Four plain system messages + role:system + index:-1 must stay at 4 blocks. + + role:system fills all four slots, so the index:-1 point is skipped. + """ + hook = AnthropicCacheControlHook() + + messages: List[AllMessageValues] = [ + {"role": "system", "content": f"System {i}"} for i in range(4) + ] + messages.append({"role": "user", "content": "hello"}) + + _, processed, _ = hook.get_chat_completion_prompt( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + non_default_params={ + "cache_control_injection_points": _build_injection_points() + }, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + + assert _count_cache_control(processed) == 4 + # All four system messages cached; user message skipped (limit reached). + assert all(processed[i].get("cache_control") is not None for i in range(4)) + assert processed[-1].get("cache_control") is None + + +def test_cache_control_hook_does_not_overwrite_existing_cache_control(): + """If a targeted message already has client cache_control, do not inject.""" + hook = AnthropicCacheControlHook() + + messages: List[AllMessageValues] = [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": "Cached by client", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ], + }, + {"role": "user", "content": "hello"}, + ] + + _, processed, _ = hook.get_chat_completion_prompt( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + # Target the already-cached system message with a different TTL. + non_default_params={ + "cache_control_injection_points": [ + { + "location": "message", + "index": 0, + "control": {"type": "ephemeral", "ttl": "5m"}, + } + ] + }, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + + # Client's 1h TTL must be preserved, not replaced by the config's 5m. + assert processed[0]["content"][-1]["cache_control"] == { + "type": "ephemeral", + "ttl": "1h", + } + assert _count_cache_control(processed) == 1 + + +@pytest.mark.asyncio +async def test_cache_control_hook_bedrock_payload_caps_cachepoints_at_four(): + """End-to-end: outgoing Bedrock payload must not exceed 4 cachePoint blocks. + + Reproduces the customer report where 4 client cache_control system blocks + plus auto-inject produced 5 cachePoint blocks and Bedrock returned 400. + """ + with patch.dict( + os.environ, + { + "AWS_ACCESS_KEY_ID": "fake_access_key_id", + "AWS_SECRET_ACCESS_KEY": "fake_secret_access_key", + "AWS_REGION_NAME": "us-east-1", + }, + ): + litellm.callbacks = [AnthropicCacheControlHook()] + + mock_response = MagicMock() + mock_response.json.return_value = { + "output": {"message": {"role": "assistant", "content": "ok"}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 100, "outputTokens": 4, "totalTokens": 104}, + } + mock_response.status_code = 200 + + client = AsyncHTTPHandler() + with patch.object(client, "post", return_value=mock_response) as mock_post: + messages = [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": f"System block {i}", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ], + } + for i in range(4) + ] + messages.append({"role": "user", "content": "hello"}) + + await litellm.acompletion( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + max_tokens=32, + cache_control_injection_points=_build_injection_points(), + client=client, + ) + + request_body = json.loads(mock_post.call_args.kwargs["data"]) + + cache_points = sum( + 1 + for block in request_body.get("system", []) + if isinstance(block, dict) and "cachePoint" in block + ) + for msg in request_body.get("messages", []): + content = msg.get("content", []) + if isinstance(content, list): + cache_points += sum( + 1 + for block in content + if isinstance(block, dict) and "cachePoint" in block + ) + + assert cache_points <= 4, ( + f"Bedrock payload exceeded Anthropic's 4 cache_control block limit: " + f"found {cache_points} cachePoint blocks" + ) + + +def test_cache_control_hook_reserves_slot_for_tool_config_point(): + """A tool_config injection point consumes one of the 4 slots downstream. + + With role:system targeting 4 system messages plus a tool_config point, the + hook must inject at most 3 message-level blocks so the tool_config cachePoint + appended by the Bedrock transform keeps the total at 4, not 5. + """ + hook = AnthropicCacheControlHook() + + messages: List[AllMessageValues] = [ + {"role": "system", "content": f"System {i}"} for i in range(4) + ] + messages.append({"role": "user", "content": "hello"}) + + _, processed, non_default_params = hook.get_chat_completion_prompt( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + non_default_params={ + "cache_control_injection_points": [ + { + "location": "message", + "role": "system", + "control": {"type": "ephemeral", "ttl": "1h"}, + }, + {"location": "tool_config"}, + ] + }, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + + assert _count_cache_control(processed) == 3 + # The tool_config point is passed through for the provider transform. + assert non_default_params["cache_control_injection_points"] == [ + {"location": "tool_config"} + ] + + +@pytest.mark.asyncio +async def test_cache_control_hook_bedrock_payload_caps_with_tool_config_point(): + """End-to-end: message + tool_config injection must not exceed 4 cachePoints.""" + with patch.dict( + os.environ, + { + "AWS_ACCESS_KEY_ID": "fake_access_key_id", + "AWS_SECRET_ACCESS_KEY": "fake_secret_access_key", + "AWS_REGION_NAME": "us-east-1", + }, + ): + litellm.callbacks = [AnthropicCacheControlHook()] + + mock_response = MagicMock() + mock_response.json.return_value = { + "output": {"message": {"role": "assistant", "content": "ok"}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 100, "outputTokens": 4, "totalTokens": 104}, + } + mock_response.status_code = 200 + + client = AsyncHTTPHandler() + with patch.object(client, "post", return_value=mock_response) as mock_post: + messages = [ + {"role": "system", "content": f"System block {i}"} for i in range(4) + ] + messages.append({"role": "user", "content": "What is the weather?"}) + + await litellm.acompletion( + model="bedrock/us.anthropic.claude-opus-4-6-v1:0", + messages=messages, + max_tokens=32, + tools=[ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather for a location", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + }, + }, + } + ], + cache_control_injection_points=[ + { + "location": "message", + "role": "system", + "control": {"type": "ephemeral", "ttl": "1h"}, + }, + {"location": "tool_config"}, + ], + client=client, + ) + + request_body = json.loads(mock_post.call_args.kwargs["data"]) + + cache_points = sum( + 1 + for block in request_body.get("system", []) + if isinstance(block, dict) and "cachePoint" in block + ) + for msg in request_body.get("messages", []): + content = msg.get("content", []) + if isinstance(content, list): + cache_points += sum( + 1 + for block in content + if isinstance(block, dict) and "cachePoint" in block + ) + for tool in request_body.get("toolConfig", {}).get("tools", []): + if isinstance(tool, dict) and "cachePoint" in tool: + cache_points += 1 + + assert cache_points <= 4, ( + f"Bedrock payload exceeded Anthropic's 4 cache_control block limit " + f"when mixing message and tool_config injection: found {cache_points}" + ) diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index f0bc7b8ebed..29e9f4529fc 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -84,6 +84,112 @@ class TestCustomGuardrailDeploymentHook: assert result["messages"] == mock_result["messages"] assert result["messages"] != original_messages + @pytest.mark.asyncio + async def test_deployment_hook_skips_when_pre_call_already_ran(self): + """The deployment hook must not re-run async_pre_call_hook once the proxy + pre-call loop has already run it for this request.""" + + class CountingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__(guardrail_name="g1", default_on=True) + self.pre_call_count = 0 + + async def async_pre_call_hook( + self, user_api_key_dict, cache, data, call_type + ): + self.pre_call_count += 1 + return data + + guardrail = CountingGuardrail() + kwargs = { + "messages": [{"role": "user", "content": "hi"}], + "model": "gpt-3.5-turbo", + "guardrails": ["g1"], + "metadata": {}, + } + + guardrail.mark_pre_call_hook_ran(kwargs) + await guardrail.async_pre_call_deployment_hook( + kwargs=kwargs, call_type=CallTypes.completion + ) + + assert guardrail.pre_call_count == 0 + + @pytest.mark.asyncio + async def test_deployment_hook_runs_when_not_marked(self): + """Without the proxy marker (direct-SDK usage) the deployment hook is the + only execution path and must still run the guardrail exactly once.""" + + class CountingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__(guardrail_name="g1", default_on=True) + self.pre_call_count = 0 + + async def async_pre_call_hook( + self, user_api_key_dict, cache, data, call_type + ): + self.pre_call_count += 1 + return data + + guardrail = CountingGuardrail() + kwargs = { + "messages": [{"role": "user", "content": "hi"}], + "model": "gpt-3.5-turbo", + "guardrails": ["g1"], + "metadata": {}, + } + + await guardrail.async_pre_call_deployment_hook( + kwargs=kwargs, call_type=CallTypes.completion + ) + + assert guardrail.pre_call_count == 1 + + def test_mark_pre_call_hook_ran_uses_litellm_metadata(self): + """The marker is recorded in litellm_metadata when that is the metadata + bucket in use, and is then visible to the skip check.""" + from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY + + guardrail = CustomGuardrail(guardrail_name="g1") + kwargs = {"litellm_metadata": {}} + + guardrail.mark_pre_call_hook_ran(kwargs) + + assert kwargs["litellm_metadata"][PRE_CALL_EXECUTED_GUARDRAILS_KEY] + assert guardrail._pre_call_hook_already_ran(kwargs) is True + + @pytest.mark.asyncio + async def test_deployment_hook_ignores_forged_caller_marker(self): + """A direct-SDK caller controls request metadata but cannot know the + per-process token, so a hand-crafted marker must not suppress a + requested guardrail in async_pre_call_deployment_hook.""" + from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY + + class CountingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__(guardrail_name="g1", default_on=True) + self.pre_call_count = 0 + + async def async_pre_call_hook( + self, user_api_key_dict, cache, data, call_type + ): + self.pre_call_count += 1 + return data + + guardrail = CountingGuardrail() + kwargs = { + "messages": [{"role": "user", "content": "hi"}], + "model": "gpt-3.5-turbo", + "guardrails": ["g1"], + "metadata": {PRE_CALL_EXECUTED_GUARDRAILS_KEY: ["g1"]}, + } + + await guardrail.async_pre_call_deployment_hook( + kwargs=kwargs, call_type=CallTypes.completion + ) + + assert guardrail.pre_call_count == 1 + class TestCustomGuardrailShouldRunGuardrail: @@ -257,6 +363,54 @@ class TestCustomGuardrailShouldRunGuardrail: result is False ), "Admin config in metadata must be respected when other metadata key is empty" + def test_should_run_guardrail_key_disable_global_not_overruled_by_team_guardrail_list( + self, + ): + """Key disable_global_guardrails must take precedence over the guardrail + appearing in the team's explicit guardrails list.""" + from litellm.types.guardrails import GuardrailEventHooks + + custom_guardrail = CustomGuardrail( + guardrail_name="global_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + ) + + # Key disabled globals; team added the same guardrail to its explicit list + # (simulates what _add_guardrails_from_key_or_team_metadata produces). + data_key_disabled_team_listed = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + "metadata": { + "user_api_key_metadata": {"disable_global_guardrails": True}, + "guardrails": ["global_guardrail"], + }, + } + assert ( + custom_guardrail.should_run_guardrail( + data=data_key_disabled_team_listed, + event_type=GuardrailEventHooks.pre_call, + ) + is False + ), "Key disable_global_guardrails must win over team's explicit guardrail list" + + # Complementary: key NOT disabled, team added guardrail → should run + data_key_enabled_team_listed = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + "metadata": { + "user_api_key_metadata": {}, + "guardrails": ["global_guardrail"], + }, + } + assert ( + custom_guardrail.should_run_guardrail( + data=data_key_enabled_team_listed, + event_type=GuardrailEventHooks.pre_call, + ) + is True + ), "Guardrail in team's explicit list should run when key has not disabled globals" + def test_should_run_guardrail_with_opted_out_global_guardrails(self): """Test that per-guardrail opt-out only works from admin metadata""" from litellm.types.guardrails import GuardrailEventHooks diff --git a/tests/test_litellm/litellm_core_utils/test_fallback_utils.py b/tests/test_litellm/litellm_core_utils/test_fallback_utils.py index 0c542ff6a1b..90a61696e9d 100644 --- a/tests/test_litellm/litellm_core_utils/test_fallback_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_fallback_utils.py @@ -1,7 +1,13 @@ +"""Tests for litellm.litellm_core_utils.fallback_utils.""" + import pytest +import httpx import litellm -from litellm.litellm_core_utils.fallback_utils import async_completion_with_fallbacks +from litellm.litellm_core_utils.core_helpers import process_response_headers +from litellm.litellm_core_utils.fallback_utils import ( + async_completion_with_fallbacks, +) @pytest.mark.asyncio @@ -41,3 +47,123 @@ async def test_fallback_dict_not_mutated(monkeypatch): "primary-model", "fallback-model", ] + + +@pytest.mark.asyncio +async def test_async_completion_with_fallbacks_sets_attempted_fallbacks_header(): + """ + When a fallback succeeds, the response must carry the + `x-litellm-attempted-fallbacks` header so the proxy and other callers can + detect that a fallback occurred. Without it, + `_override_openai_response_model` stamps the requested model back over the + fallback model used. See issue #28241. + """ + response = await async_completion_with_fallbacks( + model="openai/primary-llm", + messages=[{"role": "user", "content": "hi"}], + api_key="fake-key", + mock_response=Exception("forced failure"), + kwargs={ + "fallbacks": [ + { + "model": "openai/backup-llm", + "api_key": "fake-key", + "mock_response": "backup-resp", + } + ] + }, + ) + + hidden_params = getattr(response, "_hidden_params", None) + assert isinstance(hidden_params, dict) + headers = hidden_params.get("additional_headers") or {} + assert headers.get("x-litellm-attempted-fallbacks") == 1 + + +@pytest.mark.asyncio +async def test_async_completion_with_fallbacks_header_is_zero_when_primary_succeeds(): + """ + When the primary model succeeds on the first attempt, the header should be + `0` (no fallback was used). This mirrors the existing router-level + semantics in `async_function_with_fallbacks`. + """ + response = await async_completion_with_fallbacks( + model="openai/primary-llm", + messages=[{"role": "user", "content": "hi"}], + api_key="fake-key", + mock_response="primary-resp", + kwargs={ + "fallbacks": [ + { + "model": "openai/backup-llm", + "api_key": "fake-key", + "mock_response": "backup-resp", + } + ] + }, + ) + + hidden_params = getattr(response, "_hidden_params", None) + assert isinstance(hidden_params, dict) + headers = hidden_params.get("additional_headers") or {} + assert headers.get("x-litellm-attempted-fallbacks") == 0 + assert response.choices[0].message.content == "primary-resp" + + +def test_process_response_headers_preserves_x_litellm_headers_when_internal(): + """ + `process_response_headers` must not add the `llm_provider-` prefix to + LiteLLM's own internal headers (anything starting with `x-litellm-`) when + the caller has marked the input as LiteLLM-owned. These are markers set by + LiteLLM (e.g. fallback / retry headers); the proxy and other callers look + up the bare key. + """ + result = process_response_headers( + { + "x-litellm-attempted-fallbacks": 1, + "x-litellm-model-group": "gpt-4", + "x-stainless-arch": "arm64", + }, + preserve_litellm_internal_headers=True, + ) + assert result["x-litellm-attempted-fallbacks"] == 1 + assert result["x-litellm-model-group"] == "gpt-4" + assert result["llm_provider-x-stainless-arch"] == "arm64" + + +def test_process_response_headers_prefixes_x_litellm_from_raw_provider(): + """ + On raw upstream-provider headers (default `preserve_litellm_internal_headers=False`), + a header whose name starts with `x-litellm-` MUST still get the + `llm_provider-` prefix. Otherwise a malicious provider could return + `x-litellm-attempted-fallbacks` and spoof a LiteLLM-internal marker, + bypassing the proxy model-override guard. + """ + result = process_response_headers( + { + "x-litellm-attempted-fallbacks": 99, + "x-stainless-arch": "arm64", + } + ) + assert "x-litellm-attempted-fallbacks" not in result + assert result["llm_provider-x-litellm-attempted-fallbacks"] == 99 + assert result["llm_provider-x-stainless-arch"] == "arm64" + + +def test_process_response_headers_ignores_preserve_flag_for_httpx_headers(): + """ + Some providers store raw httpx.Headers directly in _hidden_params["additional_headers"] + without a prior normalization pass. If preserve_litellm_internal_headers=True were + honored for httpx.Headers inputs, a provider returning x-litellm-attempted-fallbacks + could spoof it as a bare LiteLLM-internal marker and make the proxy skip + stamping the correct response model. The flag must be ignored for httpx.Headers. + """ + raw = httpx.Headers( + { + "x-litellm-attempted-fallbacks": "1", + "content-type": "application/json", + } + ) + result = process_response_headers(raw, preserve_litellm_internal_headers=True) + assert "x-litellm-attempted-fallbacks" not in result + assert result["llm_provider-x-litellm-attempted-fallbacks"] == "1" diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py index c5794194528..b5eb7af88b3 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -8,7 +8,8 @@ sys.path.insert( 0, os.path.abspath("../../..") ) # Adds the parent directory to the system path -from litellm import stream_chunk_builder +from litellm import ChatCompletionUsageBlock, stream_chunk_builder +from litellm.types.utils import GenericStreamingChunk from litellm.litellm_core_utils.streaming_chunk_builder_utils import ChunkProcessor from litellm.types.utils import ( ChatCompletionDeltaToolCall, @@ -324,6 +325,42 @@ def test_cache_read_input_tokens_retained(): assert usage.cache_read_input_tokens == 11775 assert usage.prompt_tokens_details.cached_tokens == 11775 +def test_cache_read_input_tokens_retained_genericstreamingchunk(): + chunk1 = GenericStreamingChunk( + text="Test1", + is_finished=False, + finish_reason="", + usage=None, + index=1, + ) + + chunk2 = GenericStreamingChunk( + text="Test2", + is_finished=True, + finish_reason="stop", + usage=ChatCompletionUsageBlock( + completion_tokens=5, + prompt_tokens=1234, + total_tokens=1239, + completion_tokens_details=None, + prompt_tokens_details=PromptTokensDetails( + audio_tokens=None, cached_tokens=543 + ).model_dump(), + ), + index=2, + ) + + # Use dictionaries directly instead of ModelResponseStream + chunks = [chunk1, chunk2] + processor = ChunkProcessor(chunks=chunks) + + usage = processor.calculate_usage( + chunks=chunks, + model="gpt-5.5", + completion_output="", + ) + + assert usage.prompt_tokens_details.cached_tokens == 543 def test_stream_chunk_builder_litellm_usage_chunks(): """ diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index 8cd9461f39d..8b774ef4dd0 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -1509,6 +1509,44 @@ def test_raise_on_model_repetition( wrapper.raise_on_model_repetition() +@pytest.mark.parametrize( + "empty_chunk_index", + [-1, -2], + ids=["last_chunk_empty", "second_to_last_chunk_empty"], +) +def test_raise_on_model_repetition_tolerates_empty_choices( + initialized_custom_stream_wrapper: CustomStreamWrapper, + empty_chunk_index: int, +): + """ + Regression test for https://github.com/BerriAI/litellm/issues/28884 + + Vertex Gemini Flash / Flash Lite with web search streaming emits + metadata-only and usage-only chunks that carry no choices. These are + appended to self.chunks, and raise_on_model_repetition() previously + accessed choices[0] unconditionally, raising IndexError mid-stream + (surfaced to users as MidStreamFallbackError -> APIConnectionError). + """ + wrapper = initialized_custom_stream_wrapper + + chunks = [ + _make_chunk("hello world"), + ModelResponseStream( + id="usage-only", + created=1741037890, + model="vertex_ai/gemini-3.1-flash-lite", + choices=[], + usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10), + ), + ] + if empty_chunk_index == -2: + chunks.append(_make_chunk("hello world again")) + + for chunk in chunks: + wrapper.chunks.append(chunk) + wrapper.raise_on_model_repetition() + + def test_usage_chunk_after_finish_reason_updates_hidden_params(logging_obj): """ Test that provider-reported usage from a post-finish_reason chunk 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 a81261d5ffd..76aa3a9c6aa 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 @@ -1295,6 +1295,12 @@ CACHE_CONTROL_BEDROCK_CONVERSE_MODEL = ( "bedrock/converse/global.anthropic.claude-opus-4-5-20251101-v1:0" ) CACHE_CONTROL_NON_ANTHROPIC_MODEL = "gpt-4" +# Bedrock Application Inference Profile ARN: the string contains neither +# "anthropic" nor "claude", so the model can only be recognized via its ARN shape +CACHE_CONTROL_BEDROCK_ARN_MODEL = ( + "bedrock/converse/arn:aws:bedrock:us-east-1:123456789012:" + "application-inference-profile/abcdef123456" +) def test_should_add_cache_control_for_anthropic_model(): @@ -1411,6 +1417,68 @@ def test_cache_control_not_preserved_for_non_claude_model(): assert "cache_control" not in result[0]["content"][0] +@pytest.mark.parametrize( + "model, expected", + [ + (CACHE_CONTROL_BEDROCK_ARN_MODEL, True), + ( + "arn:aws-us-gov:bedrock:us-gov-west-1:123:application-inference-profile/x", + True, + ), + ("bedrock/amazon.titan-text-express-v1", False), + ("arn:aws:sagemaker:us-east-1:123:endpoint/my-endpoint", False), + ("arn:aws:sagemaker:us-east-1:123:endpoint/my-bedrock-transcriber", False), + (CACHE_CONTROL_NON_ANTHROPIC_MODEL, False), + ], +) +def test_is_bedrock_arn_model(model, expected): + """is_bedrock_arn_model requires an ARN with bedrock in the service field, not just anywhere.""" + assert LiteLLMAnthropicMessagesAdapter.is_bedrock_arn_model(model) is expected + + +def test_cache_control_preserved_for_bedrock_arn_inference_profile(): + """ + Regression for https://github.com/BerriAI/litellm/issues/26625 + + Bedrock Application Inference Profile ARNs hide the underlying Claude model + name, so cache_control must still be preserved through the /v1/messages adapter. + """ + anthropic_messages = [ + AnthropicMessagesUserMessageParam( + role="user", + content=[ + { + "type": "text", + "text": "This is cached content", + "cache_control": {"type": "ephemeral"}, + } + ], + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_messages_to_openai( + messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_ARN_MODEL + ) + + assert len(result) == 1 + assert result[0]["content"][0]["cache_control"] == {"type": "ephemeral"} + + +def test_cache_control_fix_does_not_broaden_claude_detection(): + """ + The cache_control fix is scoped to _add_cache_control_if_applicable; it must not + make is_anthropic_claude_model treat ARN profiles as Claude, which would route + thinking params through unmodified and break non-Claude Bedrock profiles. + """ + assert ( + LiteLLMAnthropicMessagesAdapter.is_anthropic_claude_model( + CACHE_CONTROL_BEDROCK_ARN_MODEL + ) + is False + ) + + 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/base_llm/test_base_model_iterator.py b/tests/test_litellm/llms/base_llm/test_base_model_iterator.py index b7f12a92ccc..f54e71c7c3c 100644 --- a/tests/test_litellm/llms/base_llm/test_base_model_iterator.py +++ b/tests/test_litellm/llms/base_llm/test_base_model_iterator.py @@ -221,3 +221,38 @@ async def test_pydantic_basemodel_chunk_passes_through_async(): assert len(chunks) == 1 assert "response.created" in chunks[0]["text"] + + +@pytest.mark.asyncio +async def test_aclose_closes_attached_http_response(): + """Regression for BerriAI/litellm#30244: CustomStreamWrapper.aclose() can + only release the upstream provider connection if the iterator exposes + aclose() and it reaches the underlying HTTP response. Without this, a + client disconnect leaves backends like vLLM generating into a dead pipe.""" + from unittest.mock import AsyncMock, MagicMock + + async def async_gen(): + yield "data: {}" + + iterator = BaseModelResponseIterator( + streaming_response=async_gen(), sync_stream=False + ) + http_response = MagicMock() + http_response.aclose = AsyncMock() + iterator.http_response = http_response + + await iterator.aclose() + + http_response.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_aclose_is_noop_without_http_response(): + async def async_gen(): + yield "data: {}" + + iterator = BaseModelResponseIterator( + streaming_response=async_gen(), sync_stream=False + ) + + await iterator.aclose() diff --git a/tests/test_litellm/llms/fastcrw/search/test_transformation.py b/tests/test_litellm/llms/fastcrw/search/test_transformation.py new file mode 100644 index 00000000000..adf8fec087c --- /dev/null +++ b/tests/test_litellm/llms/fastcrw/search/test_transformation.py @@ -0,0 +1,182 @@ +import os +from unittest.mock import Mock, patch + +import pytest + +import litellm +from litellm.llms.fastcrw.search.transformation import FastCRWSearchConfig + + +def _config() -> FastCRWSearchConfig: + return FastCRWSearchConfig() + + +def test_fastcrw_search_request_body(): + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "success": True, + "data": [ + { + "title": "Test Title", + "url": "https://example.com", + "description": "Test description", + "markdown": "Test content", + } + ], + } + + with ( + patch.dict(os.environ, {"CRW_API_KEY": "test-api-key"}), + patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=mock_response, + ) as mock_post, + ): + response = litellm.search( + query="test query", + search_provider="fastcrw", + max_results=10, + ) + + assert mock_post.called + call_kwargs = mock_post.call_args.kwargs + assert call_kwargs.get("url", "").endswith("/search") + + request_body = call_kwargs.get("json") + assert request_body is not None + assert request_body["query"] == "test query" + assert request_body["limit"] == 10 + + assert len(response.results) == 1 + result = response.results[0] + assert result.title == "Test Title" + assert result.url == "https://example.com" + assert result.snippet == "Test content" + + +def test_ui_friendly_name(): + assert _config().ui_friendly_name() == "fastCRW" + + +def test_validate_environment_with_explicit_key(): + headers = _config().validate_environment({}, api_key="explicit-key") + assert headers["Authorization"] == "Bearer explicit-key" + assert headers["Content-Type"] == "application/json" + + +def test_validate_environment_reads_env_key(): + with patch.dict(os.environ, {"CRW_API_KEY": "env-key"}, clear=False): + headers = _config().validate_environment({}) + assert headers["Authorization"] == "Bearer env-key" + + +def test_validate_environment_missing_key_raises(): + with patch.dict(os.environ, {}, clear=True): + with pytest.raises(ValueError, match="CRW_API_KEY"): + _config().validate_environment({}) + + +def test_get_complete_url_default_base(): + with patch.dict(os.environ, {}, clear=True): + assert _config().get_complete_url(None, {}) == "https://fastcrw.com/api/v1/search" + + +def test_get_complete_url_appends_search(): + assert ( + _config().get_complete_url("https://self-hosted.local/api/v1", {}) + == "https://self-hosted.local/api/v1/search" + ) + + +def test_get_complete_url_does_not_double_append(): + assert ( + _config().get_complete_url("https://self-hosted.local/api/v1/search", {}) + == "https://self-hosted.local/api/v1/search" + ) + + +def test_get_complete_url_reads_env_base(): + with patch.dict( + os.environ, {"CRW_API_BASE": "https://env-base.local/v1"}, clear=True + ): + assert _config().get_complete_url(None, {}) == "https://env-base.local/v1/search" + + +def test_transform_search_request_basic(): + data = _config().transform_search_request("hello", {"max_results": 5}) + assert data["query"] == "hello" + assert data["limit"] == 5 + assert data["scrapeOptions"]["formats"] == ["markdown"] + assert data["scrapeOptions"]["onlyMainContent"] is True + + +def test_transform_search_request_joins_list_query(): + assert _config().transform_search_request(["foo", "bar"], {})["query"] == "foo bar" + + +def test_transform_search_request_passes_through_extra_params(): + data = _config().transform_search_request("q", {"sources": ["web", "images"]}) + assert data["sources"] == ["web", "images"] + + +def test_transform_search_request_preserves_explicit_scrape_options(): + custom = {"formats": ["html"]} + data = _config().transform_search_request("q", {"scrapeOptions": custom}) + assert data["scrapeOptions"] == custom + + +def _resp(payload): + r = Mock() + r.json.return_value = payload + return r + + +def test_transform_search_response_prefers_markdown(): + resp = _config().transform_search_response( + _resp( + { + "success": True, + "data": [ + { + "title": "T", + "url": "https://e.com", + "description": "d", + "markdown": "md", + } + ], + } + ), + logging_obj=Mock(), + ) + assert len(resp.results) == 1 + assert resp.results[0].snippet == "md" + + +def test_transform_search_response_falls_back_to_description(): + resp = _config().transform_search_response( + _resp( + { + "success": True, + "data": [ + {"title": "T", "url": "https://e.com", "description": "only-desc"} + ], + } + ), + logging_obj=Mock(), + ) + assert resp.results[0].snippet == "only-desc" + + +def test_transform_search_response_empty_data(): + resp = _config().transform_search_response( + _resp({"success": True, "data": []}), logging_obj=Mock() + ) + assert resp.results == [] + + +def test_transform_search_response_non_list_data(): + resp = _config().transform_search_response( + _resp({"success": True, "data": {"unexpected": "shape"}}), logging_obj=Mock() + ) + assert resp.results == [] diff --git a/tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py b/tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py new file mode 100644 index 00000000000..2767deae176 --- /dev/null +++ b/tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py @@ -0,0 +1,394 @@ +""" +Unit tests for ModelScope configuration. + +These tests validate the ModelScopeChatConfig class which extends OpenAIGPTConfig. +ModelScope is an OpenAI-compatible provider with minor customizations. +""" + +import json +import os +import sys + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +from unittest.mock import patch + +import httpx +import pytest +import respx + +import litellm +from litellm import completion +from litellm.llms.modelscope.chat.transformation import ModelScopeChatConfig + +DEFAULT_MODEL = "Qwen/Qwen3.5-35B-A3B" + + +class TestModelScopeConfig: + """Test class for ModelScope functionality""" + + def test_default_api_base(self): + """Test that default API base is used when none is provided""" + config = ModelScopeChatConfig() + headers = {} + api_key = "fake-modelscope-key" + + result = config.validate_environment( + headers=headers, + model=DEFAULT_MODEL, + messages=[{"role": "user", "content": "Hey"}], + optional_params={}, + litellm_params={}, + api_key=api_key, + api_base=None, + ) + + assert result["Authorization"] == f"Bearer {api_key}" + assert result["Content-Type"] == "application/json" + + @pytest.mark.respx() + def test_modelscope_completion_mock(self, respx_mock): + """Mock test for basic ModelScope completion.""" + + litellm.disable_aiohttp_transport = True + + api_key = "fake-modelscope-key" + api_base = "https://api-inference.modelscope.cn/v1" + + respx_mock.post(f"{api_base}/chat/completions").respond( + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": DEFAULT_MODEL, + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": '```python\nprint("Hey from LiteLLM!")\n```', + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 9, + "completion_tokens": 12, + "total_tokens": 21, + }, + }, + status_code=200, + ) + + response = completion( + model=f"modelscope/{DEFAULT_MODEL}", + messages=[ + {"role": "user", "content": "write code for saying hey from LiteLLM"} + ], + api_key=api_key, + api_base=api_base, + ) + + assert response is not None + assert response.choices[0].message.content is not None + assert "```python" in response.choices[0].message.content + + # ── _transform_messages tests ────────────────────────────────────── + + def test_transform_messages_flattens_text_content_list(self): + """Content lists containing only text items should be flattened to a string.""" + config = ModelScopeChatConfig() + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Hello"}, + {"type": "text", "text": " world"}, + ], + } + ] + + result = config._transform_messages(messages=messages, model=DEFAULT_MODEL) + + assert result[0]["content"] == "Hello world" + + def test_transform_messages_preserves_multimodal_content_list(self): + """Content lists with image_url should be preserved as lists for vision models.""" + config = ModelScopeChatConfig() + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is this?"}, + {"type": "image_url", "image_url": {"url": "https://example.com/img.png"}}, + ], + } + ] + + result = config._transform_messages(messages=messages, model=DEFAULT_MODEL) + + assert isinstance(result[0]["content"], list) + assert len(result[0]["content"]) == 2 + assert result[0]["content"][0]["type"] == "text" + assert result[0]["content"][1]["type"] == "image_url" + + def test_transform_messages_string_content_unchanged(self): + """Messages with string content should pass through unchanged.""" + config = ModelScopeChatConfig() + messages = [{"role": "user", "content": "Hello"}] + + result = config._transform_messages(messages=messages, model=DEFAULT_MODEL) + + assert result[0]["content"] == "Hello" + + def test_transform_messages_multi_turn(self): + """Multi-turn conversations should be handled correctly.""" + config = ModelScopeChatConfig() + messages = [ + {"role": "user", "content": "Hi"}, + {"role": "assistant", "content": "Hello!"}, + { + "role": "user", + "content": [ + {"type": "text", "text": "Tell me more"}, + ], + }, + ] + + result = config._transform_messages(messages=messages, model=DEFAULT_MODEL) + + assert result[0]["content"] == "Hi" + assert result[1]["content"] == "Hello!" + assert result[2]["content"] == "Tell me more" + + def test_transform_messages_multimodal_multi_turn(self): + """Multi-turn with mixed text-only and multimodal messages.""" + config = ModelScopeChatConfig() + messages = [ + {"role": "user", "content": "Hi"}, + {"role": "assistant", "content": "Hello!"}, + { + "role": "user", + "content": [ + {"type": "text", "text": "Describe this image"}, + {"type": "image_url", "image_url": {"url": "https://example.com/photo.jpg"}}, + ], + }, + ] + + result = config._transform_messages(messages=messages, model=DEFAULT_MODEL) + + assert result[0]["content"] == "Hi" + assert result[1]["content"] == "Hello!" + # Multimodal message should keep list format + assert isinstance(result[2]["content"], list) + assert result[2]["content"][1]["type"] == "image_url" + + # ── get_complete_url tests ───────────────────────────────────────── + + def test_get_complete_url_default(self): + """Default api_base should append /chat/completions.""" + config = ModelScopeChatConfig() + + url = config.get_complete_url( + api_base=None, + api_key="fake-key", + model=DEFAULT_MODEL, + optional_params={}, + litellm_params={}, + ) + + assert url == "https://api-inference.modelscope.cn/v1/chat/completions" + + def test_get_complete_url_custom_base(self): + """Custom api_base should append /chat/completions.""" + config = ModelScopeChatConfig() + + url = config.get_complete_url( + api_base="https://custom.modelscope.cn/v1", + api_key="fake-key", + model=DEFAULT_MODEL, + optional_params={}, + litellm_params={}, + ) + + assert url == "https://custom.modelscope.cn/v1/chat/completions" + + def test_get_complete_url_already_has_endpoint(self): + """api_base already ending in /chat/completions should not be doubled.""" + config = ModelScopeChatConfig() + + url = config.get_complete_url( + api_base="https://api-inference.modelscope.cn/v1/chat/completions", + api_key="fake-key", + model=DEFAULT_MODEL, + optional_params={}, + litellm_params={}, + ) + + assert url == "https://api-inference.modelscope.cn/v1/chat/completions" + assert url.count("/chat/completions") == 1 + + # ── _get_openai_compatible_provider_info tests ───────────────────── + + def test_get_provider_info_with_explicit_api_base(self): + """Explicit api_base and api_key should be returned as-is.""" + config = ModelScopeChatConfig() + + api_base, api_key = config._get_openai_compatible_provider_info( + api_base="https://custom.example.com/v1", + api_key="my-key", + ) + + assert api_base == "https://custom.example.com/v1" + assert api_key == "my-key" + + def test_get_provider_info_default_fallback(self): + """When no api_base or env var is set, DEFAULT_BASE_URL should be used.""" + config = ModelScopeChatConfig() + + with patch.dict(os.environ, {}, clear=False): + os.environ.pop("MODELSCOPE_API_BASE", None) + os.environ.pop("MODELSCOPE_API_KEY", None) + + api_base, api_key = config._get_openai_compatible_provider_info( + api_base=None, + api_key=None, + ) + + assert api_base == "https://api-inference.modelscope.cn/v1" + assert api_key is None + + def test_get_provider_info_env_var_fallback(self): + """MODELSCOPE_API_BASE env var should be used when api_base is not provided.""" + config = ModelScopeChatConfig() + + with patch.dict( + os.environ, + {"MODELSCOPE_API_BASE": "https://env.modelscope.cn/v1"}, + ): + api_base, _ = config._get_openai_compatible_provider_info( + api_base=None, + api_key=None, + ) + + assert api_base == "https://env.modelscope.cn/v1" + + # ── Mock HTTP tests ──────────────────────────────────────────────── + + @pytest.mark.respx() + def test_completion_with_text_content_list(self, respx_mock): + """Verify that text-only content list messages are flattened before sending.""" + litellm.disable_aiohttp_transport = True + + api_key = "fake-modelscope-key" + api_base = "https://api-inference.modelscope.cn/v1" + captured_request = {} + + def capture_request(request): + captured_request["body"] = request.content + return httpx.Response( + 200, + json={ + "id": "chatcmpl-456", + "object": "chat.completion", + "created": 1677652288, + "model": DEFAULT_MODEL, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Sure!"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 1, "total_tokens": 6}, + }, + ) + + respx_mock.post(f"{api_base}/chat/completions").mock(side_effect=capture_request) + + response = completion( + model=f"modelscope/{DEFAULT_MODEL}", + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": "Hello"}, + {"type": "text", "text": " world"}, + ], + } + ], + api_key=api_key, + api_base=api_base, + ) + + assert response.choices[0].message.content == "Sure!" + + body = json.loads(captured_request["body"]) + assert isinstance(body["messages"][0]["content"], str) + assert body["messages"][0]["content"] == "Hello world" + + @pytest.mark.respx() + def test_completion_with_multimodal_messages(self, respx_mock): + """Verify that multimodal messages (text + image_url) are sent as content lists.""" + litellm.disable_aiohttp_transport = True + + api_key = "fake-modelscope-key" + api_base = "https://api-inference.modelscope.cn/v1" + captured_request = {} + + def capture_request(request): + captured_request["body"] = request.content + return httpx.Response( + 200, + json={ + "id": "chatcmpl-789", + "object": "chat.completion", + "created": 1677652288, + "model": DEFAULT_MODEL, + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "A cat sitting on a couch.", + }, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 100, "completion_tokens": 8, "total_tokens": 108}, + }, + ) + + respx_mock.post(f"{api_base}/chat/completions").mock(side_effect=capture_request) + + response = completion( + model=f"modelscope/{DEFAULT_MODEL}", + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this image?"}, + { + "type": "image_url", + "image_url": {"url": "https://example.com/cat.jpg"}, + }, + ], + } + ], + api_key=api_key, + api_base=api_base, + ) + + assert response.choices[0].message.content == "A cat sitting on a couch." + + body = json.loads(captured_request["body"]) + msg = body["messages"][0] + # Multimodal content should remain as a list + assert isinstance(msg["content"], list) + assert len(msg["content"]) == 2 + assert msg["content"][0] == {"type": "text", "text": "What is in this image?"} + assert msg["content"][1]["type"] == "image_url" + assert msg["content"][1]["image_url"]["url"] == "https://example.com/cat.jpg" diff --git a/tests/test_litellm/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py b/tests/test_litellm/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py new file mode 100644 index 00000000000..7f00f53c451 --- /dev/null +++ b/tests/test_litellm/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py @@ -0,0 +1,456 @@ +""" +Unit tests for ModelScope image generation configuration. + +These tests validate the ModelScopeImageGenerationConfig class which handles +transformation between OpenAI-compatible format and ModelScope API format. +""" + +import os +import sys +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +from litellm.llms.modelscope.image_generation.transformation import ( + ModelScopeImageGenerationConfig, +) +from litellm.types.utils import ImageResponse + + +class TestModelScopeImageGenerationTransformation: + def setup_method(self): + """Set up test fixtures before each test method.""" + self.config = ModelScopeImageGenerationConfig() + self.model = "modelscope/Qwen/Qwen-Image-Edit" + self.logging_obj = MagicMock() + + def test_get_supported_openai_params(self): + """Test that get_supported_openai_params returns correct parameters.""" + supported_params = self.config.get_supported_openai_params(self.model) + + assert "n" in supported_params + assert "size" in supported_params + assert "response_format" in supported_params + assert "user" in supported_params + + def test_map_openai_params(self): + """Test that map_openai_params correctly passes through parameters.""" + non_default_params = { + "n": 2, + "size": "1024x1024", + "response_format": "url", + } + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert result["n"] == 2 + assert result["size"] == "1024x1024" + assert result["response_format"] == "url" + + def test_map_openai_params_with_user(self): + """Test that map_openai_params correctly passes through user parameter.""" + non_default_params = {"user": "test-user-123"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert result["user"] == "test-user-123" + + def test_get_complete_url_default(self): + """Test that get_complete_url returns default ModelScope URL.""" + result = self.config.get_complete_url( + api_base=None, + api_key="test_key", + model=self.model, + optional_params={}, + litellm_params={}, + ) + + assert result == "https://api-inference.modelscope.cn/v1/images/generations" + + def test_get_complete_url_with_custom_base(self): + """Test that get_complete_url uses custom api_base.""" + custom_base = "https://custom.modelscope.cn/v1" + + result = self.config.get_complete_url( + api_base=custom_base, + api_key="test_key", + model=self.model, + optional_params={}, + litellm_params={}, + ) + + assert result == f"{custom_base}/images/generations" + + def test_get_complete_url_with_trailing_slash(self): + """Test that get_complete_url strips trailing slashes from base.""" + custom_base = "https://custom.modelscope.cn/v1/" + + result = self.config.get_complete_url( + api_base=custom_base, + api_key="test_key", + model=self.model, + optional_params={}, + litellm_params={}, + ) + + assert result == "https://custom.modelscope.cn/v1/images/generations" + + @patch("litellm.llms.modelscope.image_generation.transformation.get_secret_str") + def test_validate_environment_with_api_key(self, mock_get_secret): + """Test that validate_environment correctly sets authorization header.""" + headers = {} + api_key = "test_api_key" + + result = self.config.validate_environment( + headers=headers, + model=self.model, + messages=[], + optional_params={}, + litellm_params={}, + api_key=api_key, + ) + + assert result["Authorization"] == f"Bearer {api_key}" + assert result["Content-Type"] == "application/json" + mock_get_secret.assert_not_called() + + @patch("litellm.llms.modelscope.image_generation.transformation.get_secret_str") + def test_validate_environment_with_secret_key(self, mock_get_secret): + """Test that validate_environment uses secret API key when api_key is None.""" + mock_get_secret.return_value = "secret_api_key" + headers = {} + + result = self.config.validate_environment( + headers=headers, + model=self.model, + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + + assert result["Authorization"] == "Bearer secret_api_key" + mock_get_secret.assert_called_once_with("MODELSCOPE_API_KEY") + + @patch("litellm.llms.modelscope.image_generation.transformation.get_secret_str") + def test_validate_environment_no_api_key(self, mock_get_secret): + """Test that validate_environment raises error when no API key is available.""" + mock_get_secret.return_value = None + headers = {} + + with pytest.raises(ValueError) as exc_info: + self.config.validate_environment( + headers=headers, + model=self.model, + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + + assert "MODELSCOPE_API_KEY is not set" in str(exc_info.value) + + def test_transform_image_generation_request_basic(self): + """Test that transform_image_generation_request creates correct request body.""" + prompt = "A beautiful sunset over mountains" + optional_params = {} + + result = self.config.transform_image_generation_request( + model=self.model, + prompt=prompt, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert result["model"] == self.model + assert result["prompt"] == prompt + + def test_transform_image_generation_request_with_optional_params(self): + """Test that transform_image_generation_request includes optional params.""" + prompt = "A beautiful sunset" + optional_params = { + "n": 2, + "size": "1024x1024", + "response_format": "b64_json", + } + + result = self.config.transform_image_generation_request( + model=self.model, + prompt=prompt, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert result["model"] == self.model + assert result["prompt"] == prompt + assert result["n"] == 2 + assert result["size"] == "1024x1024" + assert result["response_format"] == "b64_json" + + def test_transform_image_generation_request_ignores_internal_params(self): + """Test that transform_image_generation_request ignores params starting with _.""" + prompt = "A beautiful sunset" + optional_params = { + "n": 2, + "_internal_param": "should_be_ignored", + } + + result = self.config.transform_image_generation_request( + model=self.model, + prompt=prompt, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert result["model"] == self.model + assert result["n"] == 2 + assert "_internal_param" not in result + + def test_transform_image_generation_response_with_url_images(self): + """Test that transform_image_generation_response correctly extracts URL images.""" + response_data = { + "created": 1234567890, + "data": [ + {"url": "https://example.com/image1.png"}, + {"url": "https://example.com/image2.png"}, + ], + } + + mock_response = MagicMock() + mock_response.json.return_value = response_data + mock_response.status_code = 200 + mock_response.headers = {} + + model_response = ImageResponse(data=[]) + + result = self.config.transform_image_generation_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert len(result.data) == 2 + assert result.data[0].url == "https://example.com/image1.png" + assert result.data[1].url == "https://example.com/image2.png" + + def test_transform_image_generation_response_with_b64_json(self): + """Test that transform_image_generation_response correctly extracts base64 images.""" + response_data = { + "created": 1234567890, + "data": [ + {"b64_json": "iVBORw0KGgoAAAANS"}, + ], + } + + mock_response = MagicMock() + mock_response.json.return_value = response_data + mock_response.status_code = 200 + mock_response.headers = {} + + model_response = ImageResponse(data=[]) + + result = self.config.transform_image_generation_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert len(result.data) == 1 + assert result.data[0].b64_json == "iVBORw0KGgoAAAANS" + assert result.data[0].url is None + + def test_transform_image_generation_response_with_revised_prompt(self): + """Test that transform_image_generation_response extracts revised_prompt.""" + response_data = { + "created": 1234567890, + "data": [ + { + "url": "https://example.com/image.png", + "revised_prompt": "A detailed description of a beautiful sunset", + }, + ], + } + + mock_response = MagicMock() + mock_response.json.return_value = response_data + mock_response.status_code = 200 + mock_response.headers = {} + + model_response = ImageResponse(data=[]) + + result = self.config.transform_image_generation_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert len(result.data) == 1 + assert ( + result.data[0].revised_prompt + == "A detailed description of a beautiful sunset" + ) + + def test_transform_image_generation_response_empty_data(self): + """Test that transform_image_generation_response handles empty data array.""" + response_data = { + "created": 1234567890, + "data": [], + } + + mock_response = MagicMock() + mock_response.json.return_value = response_data + mock_response.status_code = 200 + mock_response.headers = {} + + model_response = ImageResponse(data=[]) + + result = self.config.transform_image_generation_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert len(result.data) == 0 + + def test_transform_image_generation_response_error_handling(self): + """Test that transform_image_generation_response raises error on API error.""" + response_data = { + "error": { + "message": "Invalid prompt provided", + "type": "invalid_request_error", + } + } + + mock_response = MagicMock() + mock_response.json.return_value = response_data + mock_response.status_code = 400 + mock_response.headers = {} + + model_response = ImageResponse(data=[]) + + with pytest.raises(Exception) as exc_info: + self.config.transform_image_generation_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert "ModelScope error" in str(exc_info.value) + assert "Invalid prompt provided" in str(exc_info.value) + + def test_transform_image_generation_response_json_error(self): + """Test that transform_image_generation_response raises error on invalid JSON.""" + import json + + mock_response = MagicMock() + mock_response.json.side_effect = json.JSONDecodeError("Invalid JSON", "", 0) + mock_response.status_code = 500 + mock_response.headers = {} + + model_response = ImageResponse(data=[]) + + with pytest.raises(Exception) as exc_info: + self.config.transform_image_generation_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert "Error parsing ModelScope response" in str(exc_info.value) + + def test_get_error_class_bad_request(self): + """Test that get_error_class returns BadRequestError for 400 status.""" + from litellm.exceptions import BadRequestError + + error = self.config.get_error_class( + error_message="Bad request", + status_code=400, + headers={"Content-Type": "application/json"}, + ) + + assert isinstance(error, BadRequestError) + + def test_get_error_class_authentication_error(self): + """Test that get_error_class returns AuthenticationError for 401 status.""" + from litellm.exceptions import AuthenticationError + + error = self.config.get_error_class( + error_message="Invalid API key", + status_code=401, + headers={"Content-Type": "application/json"}, + ) + + assert isinstance(error, AuthenticationError) + + def test_get_error_class_internal_server_error(self): + """Test that get_error_class returns InternalServerError for 500+ status.""" + from litellm.exceptions import InternalServerError + + error = self.config.get_error_class( + error_message="Internal server error", + status_code=500, + headers={"Content-Type": "application/json"}, + ) + + assert isinstance(error, InternalServerError) + + def test_get_error_class_default(self): + """Test that get_error_class returns BadRequestError for other status codes.""" + from litellm.exceptions import BadRequestError + + error = self.config.get_error_class( + error_message="Some error", + status_code=404, + headers={"Content-Type": "application/json"}, + ) + + assert isinstance(error, BadRequestError) diff --git a/tests/test_litellm/llms/openai_like/test_libertai_provider.py b/tests/test_litellm/llms/openai_like/test_libertai_provider.py new file mode 100644 index 00000000000..fdbe3046e9b --- /dev/null +++ b/tests/test_litellm/llms/openai_like/test_libertai_provider.py @@ -0,0 +1,131 @@ +""" +Tests for LibertAI provider configuration and integration. +""" + +import litellm + + +class TestLibertAIProviderConfig: + """Test LibertAI provider configuration""" + + def test_libertai_in_provider_list(self): + """Test that libertai is in the provider list""" + from litellm import LlmProviders + + assert hasattr(LlmProviders, "LIBERTAI") + assert LlmProviders.LIBERTAI.value == "libertai" + assert "libertai" in litellm.provider_list + + def test_libertai_json_config_exists(self): + """Test that libertai is configured in providers.json""" + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + assert JSONProviderRegistry.exists("libertai") + + libertai = JSONProviderRegistry.get("libertai") + assert libertai is not None + assert libertai.base_url == "https://api.libertai.io/v1" + assert libertai.api_key_env == "LIBERTAI_API_KEY" + assert libertai.api_base_env == "LIBERTAI_API_BASE" + assert libertai.param_mappings.get("max_completion_tokens") == "max_tokens" + + def test_libertai_provider_resolution(self): + """Test that provider resolution finds libertai and the default base URL""" + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + model, provider, api_key, api_base = get_llm_provider( + model="libertai/qwen3.6-27b", + custom_llm_provider=None, + api_base=None, + api_key=None, + ) + + assert model == "qwen3.6-27b" + assert provider == "libertai" + assert api_base == "https://api.libertai.io/v1" + + def test_libertai_api_base_override(self): + """Test that an explicit api_base / api_key overrides the default""" + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + model, provider, api_key, api_base = get_llm_provider( + model="libertai/qwen3.6-27b", + custom_llm_provider=None, + api_base="https://custom.example.com/v1", + api_key="sk-test", + ) + + assert provider == "libertai" + assert api_base == "https://custom.example.com/v1" + assert api_key == "sk-test" + + def test_libertai_model_cost_map(self): + """Test that libertai models are present in the model cost map""" + model_cost = litellm.model_cost + + assert "libertai/qwen3.6-27b" in model_cost + info = model_cost["libertai/qwen3.6-27b"] + assert info["litellm_provider"] == "libertai" + assert info["mode"] == "chat" + assert info["max_input_tokens"] == 262144 + assert info["max_output_tokens"] == 262144 + + # thinking variants are marked as reasoning models + assert ( + model_cost["libertai/qwen3.6-27b-thinking"].get("supports_reasoning") + is True + ) + + def test_libertai_router_config(self): + """Test that libertai can be used in Router configuration""" + from litellm import Router + + router = Router( + model_list=[ + { + "model_name": "libertai-chat", + "litellm_params": { + "model": "libertai/qwen3.6-27b", + "api_key": "test-key", + }, + } + ] + ) + + assert len(router.model_list) == 1 + assert router.model_list[0]["model_name"] == "libertai-chat" + + def test_libertai_model_modes(self): + """Chat models carry mode 'chat'; the embedding model carries mode 'embedding'.""" + model_cost = litellm.model_cost + + # chat model + assert model_cost["libertai/qwen3.6-27b"]["mode"] == "chat" + + # embedding model (bge-m3) must be normalized to mode 'embedding' so + # /embeddings routing and the supported-endpoints matrix stay consistent + assert "libertai/bge-m3" in model_cost + bge = model_cost["libertai/bge-m3"] + assert bge["litellm_provider"] == "libertai" + assert bge["mode"] == "embedding" + + def test_libertai_supported_endpoints_matrix(self): + """The runtime-served backup matrix (GET /public/supported_endpoints) lists libertai.""" + import json + from pathlib import Path + + import litellm as _litellm + + backup_path = ( + Path(_litellm.__file__).parent / "provider_endpoints_support_backup.json" + ) + matrix = json.loads(backup_path.read_text()) + + assert "libertai" in matrix["providers"] + endpoints = matrix["providers"]["libertai"]["endpoints"] + assert endpoints["chat_completions"] is True + # embeddings is advertised false: the JSON-configured-provider path only + # wires chat routing (the OpenAILike embedding handler is reached solely + # for the literal openai_like/llamafile/lm_studio providers), matching + # the llamagate precedent. bge-m3 stays in the cost map for metadata. + assert endpoints["embeddings"] is False 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 cc8b14e5514..cf75964ddb7 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 @@ -1917,3 +1917,57 @@ class TestVertexAIGlobalLocation: assert "generativelanguage.googleapis.com" in url assert "cachedContents" in url + + +class TestContextCachingMultiRegionUrls: + """Regression coverage for #29571: multi-region vertex_location values + (`eu`, `us`) must resolve to the REP host (`aiplatform.{geo}.rep.googleapis.com`) + on the cachedContents endpoint, matching the inference path (already + fixed in #27293). Previously the URL was hardcoded to + `{location}-aiplatform.googleapis.com`, which doesn't exist for + multi-region locations and 404'd.""" + + def setup_method(self): + self.caching = ContextCachingEndpoints() + + @pytest.mark.parametrize("location", ["eu", "us"]) + def test_vertex_ai_multi_region_uses_rep_host(self, location): + _, url = self.caching._get_token_and_url_context_caching( + gemini_api_key=None, + custom_llm_provider="vertex_ai", + api_base=None, + vertex_project="my-project", + vertex_location=location, + vertex_auth_header="Bearer token", + ) + + assert url.startswith(f"https://aiplatform.{location}.rep.googleapis.com/") + assert f"/locations/{location}/cachedContents" in url + # Old broken host must no longer appear. + assert f"{location}-aiplatform.googleapis.com" not in url + + def test_vertex_ai_regional_still_uses_regional_host(self): + _, url = self.caching._get_token_and_url_context_caching( + gemini_api_key=None, + custom_llm_provider="vertex_ai", + api_base=None, + vertex_project="my-project", + vertex_location="us-central1", + vertex_auth_header="Bearer token", + ) + + assert url.startswith("https://us-central1-aiplatform.googleapis.com/") + assert "/locations/us-central1/cachedContents" in url + + def test_vertex_ai_global_still_uses_global_host(self): + _, url = self.caching._get_token_and_url_context_caching( + gemini_api_key=None, + custom_llm_provider="vertex_ai", + api_base=None, + vertex_project="my-project", + vertex_location="global", + vertex_auth_header="Bearer token", + ) + + assert url.startswith("https://aiplatform.googleapis.com/") + assert "/locations/global/cachedContents" in url diff --git a/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py b/tests/test_litellm/llms/vertex_ai/test_gemini_batch_embeddings.py similarity index 95% rename from tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py rename to tests/test_litellm/llms/vertex_ai/test_gemini_batch_embeddings.py index 54ea41a6450..98abf5459df 100644 --- a/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py +++ b/tests/test_litellm/llms/vertex_ai/test_gemini_batch_embeddings.py @@ -8,12 +8,8 @@ This test ensures that: """ import json -import os -import sys from unittest.mock import MagicMock, patch -sys.path.insert(0, os.path.abspath("../../../..")) - import pytest import litellm @@ -311,13 +307,16 @@ def test_gemini_multimodal_embedding_e2e(): ): mock_get_token.return_value = ( {"x-goog-api-key": "test-key"}, - "https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-2-preview:embedContent", + "https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-2-preview:batchEmbedContents", ) mock_response = MagicMock() mock_response.status_code = 200 mock_response.json.return_value = { - "embedding": {"values": [0.1, 0.2, 0.3, 0.4, 0.5]} + "embeddings": [ + {"values": [0.1, 0.2, 0.3, 0.4, 0.5]}, + {"values": [0.6, 0.7, 0.8, 0.9, 1.0]}, + ] } mock_post.return_value = mock_response @@ -338,17 +337,21 @@ def test_gemini_multimodal_embedding_e2e(): request_body = json.loads(kwargs.get("data", "{}")) - assert "content" in request_body - assert "parts" in request_body["content"] - parts = request_body["content"]["parts"] + assert "requests" in request_body + assert len(request_body["requests"]) == 2 - assert len(parts) == 2 - assert parts[0]["text"] == "The food was delicious" - assert "inline_data" in parts[1] - assert parts[1]["inline_data"]["mime_type"] == "image/png" + text_parts = request_body["requests"][0]["content"]["parts"] + image_parts = request_body["requests"][1]["content"]["parts"] - assert len(response.data) == 1 + assert len(text_parts) == 1 + assert text_parts[0]["text"] == "The food was delicious" + assert len(image_parts) == 1 + assert "inline_data" in image_parts[0] + assert image_parts[0]["inline_data"]["mime_type"] == "image/png" + + assert len(response.data) == 2 assert response.data[0].embedding == [0.1, 0.2, 0.3, 0.4, 0.5] + assert response.data[1].embedding == [0.6, 0.7, 0.8, 0.9, 1.0] def test_gemini_multimodal_embedding_with_audio(): @@ -581,17 +584,21 @@ def test_vertex_ai_text_only_embedding_uses_embed_content(): def test_filter_embed_params_drops_unsupported(): """Unsupported params like max_tokens should be filtered out.""" - result = _filter_embed_params({"dimensions": 768, "max_tokens": 256, "temperature": 0.5}) + result = _filter_embed_params( + {"dimensions": 768, "max_tokens": 256, "temperature": 0.5} + ) assert result == {"outputDimensionality": 768} def test_filter_embed_params_keeps_supported(): """All supported Gemini embedding params should pass through.""" - result = _filter_embed_params({ - "dimensions": 768, - "task_type": "RETRIEVAL_DOCUMENT", - "title": "My doc", - }) + result = _filter_embed_params( + { + "dimensions": 768, + "task_type": "RETRIEVAL_DOCUMENT", + "title": "My doc", + } + ) assert result == { "outputDimensionality": 768, "taskType": "RETRIEVAL_DOCUMENT", 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 7753378ab4f..ab42ee1e979 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 @@ -658,12 +658,11 @@ class TestMCPOAuth2AuthFlow: async def test_oauth2_token_in_authorization_header_fallback(self): """ - When only Authorization header is present with a non-LiteLLM OAuth2 token - AND the target server is operator-configured for ``auth_type=oauth2``, - auth should fall back to permissive mode (OAuth2 passthrough). + When only the Authorization header is present with a non-LiteLLM OAuth2 + token AND the target server delegates auth to upstream, LiteLLM skips its + own validation entirely (so the upstream token is never mistaken for a + virtual key) and forwards the bearer upstream. """ - from fastapi import HTTPException - from litellm.types.mcp import MCPAuth scope = { @@ -675,17 +674,16 @@ class TestMCPOAuth2AuthFlow: ], } - async def mock_user_api_key_auth_fails(api_key, request): - raise HTTPException(status_code=401, detail="Invalid API key") - oauth2_server = MagicMock() oauth2_server.auth_type = MCPAuth.oauth2 + oauth2_server.delegate_auth_to_upstream = True + oauth2_server.has_client_credentials = False with ( patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", - side_effect=mock_user_api_key_auth_fails, - ), + new_callable=AsyncMock, + ) as mock_auth, patch( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" ) as mock_mgr, @@ -700,10 +698,10 @@ class TestMCPOAuth2AuthFlow: raw_headers, ) = await MCPRequestHandler.process_mcp_request(scope) - # Should succeed with default UserAPIKeyAuth (OAuth2 fallback) - assert auth_result is not None assert isinstance(auth_result, UserAPIKeyAuth) - # OAuth2 headers should contain the token for upstream forwarding + # The upstream token is never validated as a LiteLLM key ... + mock_auth.assert_not_called() + # ... and is preserved for upstream forwarding. assert ( oauth2_headers.get("Authorization") == "Bearer atlassian-oauth2-access-token-xyz" @@ -813,11 +811,12 @@ class TestMCPOAuth2AuthFlow: await MCPRequestHandler.process_mcp_request(scope) assert exc_info.value.status_code == 500 - async def test_proxy_exception_oauth2_fallback(self): + async def test_proxy_exception_non_delegate_oauth2_propagates(self): """ - user_api_key_auth raises ProxyException (not HTTPException) in production. - The OAuth2 fallback must catch ProxyException with code 401/403 too, - but only when the target server is operator-configured for ``auth_type=oauth2``. + Production raises ProxyException (not HTTPException) on auth failure. For + a non-delegate oauth2 server the bearer is treated as a LiteLLM credential + and a 401 must propagate as a real auth error, not be exchanged for an + anonymous upstream-passthrough session. """ from litellm.proxy._types import ProxyException from litellm.types.mcp import MCPAuth @@ -841,6 +840,8 @@ class TestMCPOAuth2AuthFlow: oauth2_server = MagicMock() oauth2_server.auth_type = MCPAuth.oauth2 + oauth2_server.delegate_auth_to_upstream = False + oauth2_server.is_oauth_passthrough = False with ( patch( @@ -852,22 +853,9 @@ class TestMCPOAuth2AuthFlow: ) as mock_mgr, ): mock_mgr.get_mcp_server_by_name.return_value = oauth2_server - ( - auth_result, - mcp_auth_header, - mcp_servers, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - ) = await MCPRequestHandler.process_mcp_request(scope) - - # Should succeed with default UserAPIKeyAuth (OAuth2 fallback) - assert auth_result is not None - assert isinstance(auth_result, UserAPIKeyAuth) - assert ( - oauth2_headers.get("Authorization") - == "Bearer atlassian-oauth2-access-token-xyz" - ) + with pytest.raises(ProxyException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert str(exc_info.value.code) == "401" async def test_proxy_exception_non_auth_still_raises(self): """ @@ -1355,11 +1343,15 @@ class TestMCPOAuth2FallbackTargetGating: await MCPRequestHandler.process_mcp_request(scope) assert exc_info.value.status_code == 401 - async def test_fallback_allowed_when_target_is_oauth2_mode(self): + async def test_non_delegate_oauth2_does_not_fall_back_to_anonymous(self): """ - Operator-configured OAuth2 passthrough still works: target server has - ``auth_type=oauth2`` → failed LiteLLM auth falls back to anonymous so - the bearer can be forwarded to upstream. + An ``auth_type=oauth2`` server that has NOT opted into + ``delegate_auth_to_upstream`` must not exchange a failed LiteLLM auth for + an anonymous session: forwarding an arbitrary bearer upstream is only + allowed once the operator explicitly delegates auth. A failed validation + here is a genuine 401 and propagates (which is also what keeps the + success-path trace free of a phantom 401, since no doomed validation runs + for a delegated server). """ from fastapi import HTTPException @@ -1389,8 +1381,9 @@ class TestMCPOAuth2FallbackTargetGating: mock_mgr.get_mcp_server_by_name.return_value = ( TestMCPOAuth2FallbackTargetGating._make_server(MCPAuth.oauth2) ) - auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) - assert isinstance(auth_result, UserAPIKeyAuth) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 async def test_fallback_allowed_when_target_is_passthrough(self): """ @@ -1668,19 +1661,16 @@ class TestMCPDelegateAuthToUpstream: assert isinstance(auth_result, UserAPIKeyAuth) mock_auth.assert_not_called() - async def test_delegate_with_upstream_token_in_authorization_falls_back_to_anonymous( + async def test_delegate_with_upstream_token_in_authorization_skips_litellm_auth( self, ): """ oauth2 + delegate_auth_to_upstream=True with an upstream OAuth token in - ``Authorization`` (not a LiteLLM key): LiteLLM auth is attempted first - (and fails), then the existing oauth2 fallback returns anonymous so the - bearer is forwarded upstream untouched. The delegate branch itself does - not fire when Authorization is present — that is what protects spend - tracking for callers using Authorization-style LiteLLM keys. + ``Authorization``: the delegate gate fires before any LiteLLM validation, + so ``user_api_key_auth`` is never called and the bearer is forwarded + upstream untouched. Skipping the doomed validation is what keeps a tool + call that actually succeeds from carrying a phantom 401 auth span. """ - from fastapi import HTTPException - from litellm.types.mcp import MCPAuth scope = { @@ -1690,14 +1680,11 @@ class TestMCPDelegateAuthToUpstream: "headers": [(b"authorization", b"Bearer upstream-pkce-token")], } - async def mock_user_api_key_auth_fails(api_key, request): - raise HTTPException(status_code=401, detail="Invalid API key") - with ( patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", - side_effect=mock_user_api_key_auth_fails, - ), + new_callable=AsyncMock, + ) as mock_auth, patch( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" ) as mock_mgr, @@ -1718,6 +1705,7 @@ class TestMCPDelegateAuthToUpstream: ) = await MCPRequestHandler.process_mcp_request(scope) assert isinstance(auth_result, UserAPIKeyAuth) assert oauth2_headers.get("Authorization") == "Bearer upstream-pkce-token" + mock_auth.assert_not_called() async def test_delegate_off_still_requires_litellm_auth(self): """ @@ -1912,12 +1900,15 @@ class TestMCPDelegateAuthToUpstream: assert auth_result.user_id == "real-user" mock_auth.assert_called_once() - async def test_litellm_key_via_authorization_header_not_bypassed(self): + async def test_authorization_bearer_on_delegate_server_treated_as_upstream(self): """ - Regression: a LiteLLM key sent via the secondary ``Authorization`` header - (e.g. ``Authorization: Bearer sk-...``) must still trigger normal auth - and not be silently swallowed by the delegate bypass — otherwise spend - tracking and rate limiting are skipped for those callers. + On a delegate server the ``Authorization`` header is, by contract, an + upstream token rather than a LiteLLM key — even when it is sk-shaped. It + is forwarded upstream without LiteLLM validation, so ``user_api_key_auth`` + is not called and no LiteLLM identity is resolved. Callers who need + LiteLLM identity / spend tracking on a delegate server must supply + ``x-litellm-api-key`` (see + test_explicit_litellm_key_takes_precedence_over_delegate). """ from litellm.types.mcp import MCPAuth @@ -1944,10 +1935,18 @@ class TestMCPDelegateAuthToUpstream: delegate_auth_to_upstream=True, ) ) - auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) + ( + auth_result, + _, + _, + _, + oauth2_headers, + _, + ) = await MCPRequestHandler.process_mcp_request(scope) assert isinstance(auth_result, UserAPIKeyAuth) - assert auth_result.user_id == "real-user" - mock_auth.assert_called_once() + assert auth_result.user_id is None + assert oauth2_headers.get("Authorization") == "Bearer sk-1234" + mock_auth.assert_not_called() async def test_delegate_ignored_for_client_credentials_server(self): """ diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 63b61954cf6..07b04961205 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -1400,6 +1400,65 @@ def test_rag_routes_accessible_to_internal_user_viewer(): ) +@pytest.mark.parametrize( + "route", + [ + "/vector_stores/vs_123", + "/v1/vector_stores/vs_123", + "/vector_stores/vs_123/search", + "/v1/vector_stores/vs_123/search", + "/vector_stores/vs_123/files", + "/v1/vector_stores/vs_123/files", + ], +) +def test_vector_store_routes_are_llm_api_routes(route): + """Retrieve/update/delete on a single vector store must classify as LLM API routes. + + Regression for the missing bare `/v1/vector_stores/{vector_store_id}` entry in + `openai_routes` that left retrieve/update/delete blocked for internal roles + while `/search` and `/files` sub-routes worked. + """ + + assert RouteChecks.is_llm_api_route(route) is True + + +@pytest.mark.parametrize( + "user_role", + [ + LitellmUserRoles.INTERNAL_USER.value, + LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value, + ], +) +@pytest.mark.parametrize( + "method, route", + [ + ("GET", "/v1/vector_stores/vs_123"), + ("POST", "/v1/vector_stores/vs_123"), + ("DELETE", "/v1/vector_stores/vs_123"), + ], +) +def test_vector_store_crud_accessible_to_internal_roles(user_role, method, route): + """Internal user and internal viewer must reach vector store retrieve/update/delete. + + Object-level access is still gated by `assert_user_can_access_vector_store`; + this only verifies the route gate no longer 403s these roles. + """ + + valid_token = UserAPIKeyAuth(user_id="test_user", user_role=user_role) + request = MagicMock(spec=Request) + request.method = method + request.query_params = {} + + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=LiteLLM_UserTable(user_id="test_user", user_role=user_role), + _user_role=user_role, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + + def test_videos_route_accessible_to_internal_users(): """ Test that internal users can access the videos routes. diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 9f7173383b0..0ef9ad857f9 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -180,3 +180,172 @@ def test_sync_guardrail_from_db_marks_source_db_when_unchanged(): handler.sync_guardrail_from_db(g) assert handler.get_source("collide") == "db" + + +def _db_litellm_params() -> dict: + """ + Shape produced by GuardrailRegistry.get_all_guardrails_from_db: litellm_params + is a raw dict (not a LitellmParams), holding only the keys originally stored, + a non-schema extra key, and plain-string enum values. + """ + return { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "default_on": True, + "version": 2, + "blocked_words": [{"keyword": "secret", "action": "BLOCK"}], + } + + +def test_unchanged_db_params_do_not_register_as_changed(): + """ + A DB poll returns litellm_params as a raw dict while the in-memory copy is a + LitellmParams whose model_dump() fills every field default and coerces enums. + The two shapes must compare equal when the config is identical; otherwise + every poll cycle re-initializes the guardrail indefinitely. + """ + handler = InMemoryGuardrailHandler() + raw = _db_litellm_params() + gid = "11111111-1111-1111-1111-111111111111" + handler.IN_MEMORY_GUARDRAILS[gid] = Guardrail( + guardrail_id=gid, + guardrail_name="cf", + litellm_params=LitellmParams(**raw), + ) + + new = Guardrail(guardrail_id=gid, guardrail_name="cf", litellm_params=dict(raw)) + assert handler._has_guardrail_params_changed(gid, new) is False + + +def test_changed_db_params_register_as_changed(): + """Normalizing both sides must still surface a genuine config change.""" + handler = InMemoryGuardrailHandler() + raw = _db_litellm_params() + gid = "22222222-2222-2222-2222-222222222222" + handler.IN_MEMORY_GUARDRAILS[gid] = Guardrail( + guardrail_id=gid, + guardrail_name="cf", + litellm_params=LitellmParams(**raw), + ) + + changed = {**raw, "blocked_words": [{"keyword": "different", "action": "BLOCK"}]} + new = Guardrail(guardrail_id=gid, guardrail_name="cf", litellm_params=changed) + assert handler._has_guardrail_params_changed(gid, new) is True + + +def test_unnormalizable_db_params_register_as_changed_without_raising(): + """ + A DB row whose litellm_params fail LitellmParams validation must not crash the + poll loop. The comparison falls back to treating the guardrail as changed so it + re-initializes (and surfaces the bad row in logs) rather than propagating the + validation error up through the polling cycle. + """ + handler = InMemoryGuardrailHandler() + raw = _db_litellm_params() + gid = "55555555-5555-5555-5555-555555555555" + handler.IN_MEMORY_GUARDRAILS[gid] = Guardrail( + guardrail_id=gid, + guardrail_name="cf", + litellm_params=LitellmParams(**raw), + ) + + malformed = {**raw, "default_on": "not-a-bool-xyz"} + new = Guardrail(guardrail_id=gid, guardrail_name="cf", litellm_params=malformed) + assert handler._has_guardrail_params_changed(gid, new) is True + + +def _all_callback_lists(): + import litellm + + return [ + litellm.callbacks, + litellm.success_callback, + litellm.failure_callback, + litellm._async_success_callback, + litellm._async_failure_callback, + ] + + +def test_delete_in_memory_guardrail_removes_callback_from_all_lists(): + """ + Request handling promotes guardrail callbacks from litellm.callbacks into the + success/failure/async lists. delete_in_memory_guardrail must purge the callback + from every list, otherwise a re-initialized guardrail leaves its old instance + stranded in those lists and instances accumulate. + """ + handler = InMemoryGuardrailHandler() + callback = CustomGuardrail( + guardrail_name="cf-delete", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + ) + gid = "33333333-3333-3333-3333-333333333333" + handler.IN_MEMORY_GUARDRAILS[gid] = _make_guardrail(gid, "cf-delete") + handler._sources[gid] = "db" + handler.guardrail_id_to_custom_guardrail[gid] = callback + + lists = _all_callback_lists() + snapshots = [list(cb_list) for cb_list in lists] + try: + for cb_list in lists: + cb_list.append(callback) + + handler.delete_in_memory_guardrail(gid) + + for cb_list in lists: + assert callback not in cb_list + finally: + for cb_list, snapshot in zip(lists, snapshots): + cb_list[:] = snapshot + + +def test_repeated_db_sync_does_not_accumulate_runner_instances(): + """ + End-to-end regression for the OOM: across repeated DB polls (with the config + genuinely changing each cycle to force re-initialization), exactly one live + guardrail instance must exist across all callback lists. On the unfixed code + the stale instance lingers in the success/failure lists and the distinct count + climbs above one. + """ + import litellm + + handler = InMemoryGuardrailHandler() + gid = "44444444-4444-4444-4444-444444444444" + name = "cf-accum" + + def db_guardrail(word: str) -> Guardrail: + params = { + **_db_litellm_params(), + "blocked_words": [{"keyword": word, "action": "BLOCK"}], + } + return Guardrail(guardrail_id=gid, guardrail_name=name, litellm_params=params) + + def promote_into_request_lists() -> None: + manager = litellm.logging_callback_manager + for callback in list(litellm.callbacks): + manager.add_litellm_success_callback(callback) + manager.add_litellm_failure_callback(callback) + manager.add_litellm_async_success_callback(callback) + manager.add_litellm_async_failure_callback(callback) + + def distinct_runner_instances() -> int: + seen = set() + for callback in litellm.logging_callback_manager._get_all_callbacks(): + if ( + isinstance(callback, CustomGuardrail) + and getattr(callback, "guardrail_name", None) == name + ): + seen.add(id(callback)) + return len(seen) + + lists = _all_callback_lists() + snapshots = [list(cb_list) for cb_list in lists] + try: + for cycle in range(5): + handler.sync_guardrail_from_db(db_guardrail(f"word-{cycle}")) + promote_into_request_lists() + + assert distinct_runner_instances() == 1 + finally: + for cb_list, snapshot in zip(lists, snapshots): + cb_list[:] = snapshot diff --git a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py index d924d5ecdfe..3bdf9bafdc7 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py @@ -3,6 +3,7 @@ import os import sys import types +from datetime import datetime, timedelta, timezone import pytest from unittest.mock import AsyncMock, MagicMock from fastapi.testclient import TestClient @@ -11,7 +12,6 @@ import litellm.proxy.proxy_server as ps from litellm.proxy.proxy_server import app from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles, CommonProxyErrors - sys.path.insert( 0, os.path.abspath("../../../") ) # Adds the parent directory to the system path @@ -265,3 +265,98 @@ async def test_new_budget_invalid_model_max_budget(client_and_mocks, monkeypatch assert resp.status_code in (400, 422), resp.text detail = resp.json()["detail"] assert "model_max_budget" in str(detail) or "dictionary" in str(detail).lower() + + +def _capture_update_data(mock_table): + captured = {} + + async def capture(*, where, data): + captured.update(data) + return {**where, **data} + + mock_table.update = AsyncMock(side_effect=capture) + return captured + + +@pytest.mark.asyncio +async def test_update_budget_recomputes_reset_at_when_duration_changes( + client_and_mocks, +): + """ + Regression for LIT-3362: shortening budget_duration without an explicit + budget_reset_at must bring the reset forward instead of leaving it pinned + to the previous (longer) schedule. + """ + client, _, mock_table = client_and_mocks + captured = _capture_update_data(mock_table) + + before = datetime.now(timezone.utc) + resp = client.post( + "/budget/update", + json={"budget_id": "budget_reset_recompute", "budget_duration": "1d"}, + ) + assert resp.status_code == 200, resp.text + + assert ( + "budget_reset_at" in captured + ), "duration change must recompute budget_reset_at" + reset_at = captured["budget_reset_at"] + assert isinstance(reset_at, datetime) + assert reset_at > before, "recomputed reset must be in the future" + # "1d" resets at the next standardized day boundary, always within ~24h + assert reset_at <= before + timedelta(days=1, hours=1), reset_at + # and it must be far closer than a stale 30d schedule would have left it + assert reset_at < before + timedelta(days=29) + + +@pytest.mark.asyncio +async def test_update_budget_preserves_explicit_reset_at(client_and_mocks): + """An explicit budget_reset_at from the caller always wins over recompute.""" + client, _, mock_table = client_and_mocks + captured = _capture_update_data(mock_table) + + explicit = datetime(2027, 1, 1, tzinfo=timezone.utc) + resp = client.post( + "/budget/update", + json={ + "budget_id": "budget_explicit_reset", + "budget_duration": "1d", + "budget_reset_at": explicit.isoformat(), + }, + ) + assert resp.status_code == 200, resp.text + + assert captured["budget_reset_at"] == explicit + + +@pytest.mark.asyncio +async def test_update_budget_without_duration_leaves_reset_at_untouched( + client_and_mocks, +): + """Updates that do not touch budget_duration must not introduce budget_reset_at.""" + client, _, mock_table = client_and_mocks + captured = _capture_update_data(mock_table) + + resp = client.post( + "/budget/update", + json={"budget_id": "budget_other_field", "max_budget": 200.0}, + ) + assert resp.status_code == 200, resp.text + + assert "budget_reset_at" not in captured + + +@pytest.mark.asyncio +async def test_update_budget_duration_none_does_not_recompute(client_and_mocks): + """Clearing budget_duration (explicit null) must not recompute against a None duration.""" + client, _, mock_table = client_and_mocks + captured = _capture_update_data(mock_table) + + resp = client.post( + "/budget/update", + json={"budget_id": "budget_clear_duration", "budget_duration": None}, + ) + assert resp.status_code == 200, resp.text + + assert "budget_duration" in captured and captured["budget_duration"] is None + assert "budget_reset_at" not in captured diff --git a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py index d8a674c2681..6c5ccd3562f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py @@ -409,3 +409,101 @@ async def test_get_customer_daily_activity_with_end_user_aliases(monkeypatch): "end-user-1": {"alias": "Customer One"}, "end-user-2": {"alias": "Customer Two"}, } + + +@pytest.mark.asyncio +async def test_get_customer_daily_activity_non_admin_is_rejected(monkeypatch): + """ + Security regression: any non-admin caller must receive 401 from + /customer/daily/activity and /end_user/daily/activity. + + Before this fix, the endpoint performed no role check. A caller with + user_role=INTERNAL_USER could omit end_user_ids, causing entity_id=None + to flow into get_daily_activity where the SQL builder treats it as no + filter — returning every tenant's spend across the full + LiteLLM_DailyEndUserSpend table. + + LiteLLM_EndUserTable has no per-tenant ownership column, so non-admin + scoping is not possible. The correct fix is admin-only, matching the + existing /customer/list gate. + """ + from litellm.proxy.management_endpoints import customer_endpoints + from litellm.proxy.management_endpoints.customer_endpoints import ( + get_customer_daily_activity, + ) + + mock_prisma_client = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + get_daily_activity_mock = AsyncMock() + monkeypatch.setattr( + customer_endpoints, "get_daily_activity", get_daily_activity_mock + ) + + non_admin_key = UserAPIKeyAuth( + user_id="regular-user-abc", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + with pytest.raises(HTTPException) as exc_info: + await get_customer_daily_activity( + end_user_ids=None, + start_date="2025-01-01", + end_date="2025-01-31", + model=None, + api_key=None, + page=1, + page_size=10, + exclude_end_user_ids=None, + user_api_key_dict=non_admin_key, + ) + + assert exc_info.value.status_code == 401 + assert "Admin-only endpoint" in str(exc_info.value.detail) + get_daily_activity_mock.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_customer_daily_activity_service_account_key_is_rejected(monkeypatch): + """ + Security regression: service-account keys (user_id=None, role=INTERNAL_USER) + must be rejected at the admin gate before reaching get_daily_activity. + + A service-account key with end_user_ids omitted is the worst-case caller: + entity_id=None and no user identity to scope by — the SQL builder would + return the full LiteLLM_DailyEndUserSpend table with no WHERE clause. + """ + from litellm.proxy.management_endpoints import customer_endpoints + from litellm.proxy.management_endpoints.customer_endpoints import ( + get_customer_daily_activity, + ) + + mock_prisma_client = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + get_daily_activity_mock = AsyncMock() + monkeypatch.setattr( + customer_endpoints, "get_daily_activity", get_daily_activity_mock + ) + + service_account_key = UserAPIKeyAuth( + user_id=None, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + with pytest.raises(HTTPException) as exc_info: + await get_customer_daily_activity( + end_user_ids=None, + start_date="2025-01-01", + end_date="2025-01-31", + model=None, + api_key=None, + page=1, + page_size=10, + exclude_end_user_ids=None, + user_api_key_dict=service_account_key, + ) + + assert exc_info.value.status_code == 401 + assert "Admin-only endpoint" in str(exc_info.value.detail) + get_daily_activity_mock.assert_not_called() 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 f0198320f22..b81807ee19e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -3064,6 +3064,106 @@ async def test_list_team_v2_org_admin_sees_org_teams(): assert where["organization_id"] == {"in": ["org_A"]} +@pytest.mark.asyncio +async def test_list_team_v2_org_admin_own_user_id_sees_all_org_teams(): + """ + Test that an org admin whose own user_id is sent (as the UI does for + non-Admin roles) still sees all teams in their organization, not just + teams they are a direct member of. + + Regression test for https://github.com/BerriAI/litellm/issues/30215 + """ + from datetime import datetime + from unittest.mock import AsyncMock, Mock, patch + + from fastapi import Request + + from litellm.proxy._types import ( + LiteLLM_OrganizationMembershipTable, + LiteLLM_UserTable, + LitellmUserRoles, + UserAPIKeyAuth, + ) + from litellm.proxy.management_endpoints.team_endpoints import list_team_v2 + + mock_request = Mock(spec=Request) + mock_user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="org_admin_user", + ) + + mock_user = LiteLLM_UserTable( + user_id="org_admin_user", + teams=["team_1"], # direct member of only 1 team + organization_memberships=[ + LiteLLM_OrganizationMembershipTable( + user_id="org_admin_user", + organization_id="org_A", + user_role="org_admin", + spend=0.0, + created_at=datetime.now(), + updated_at=datetime.now(), + ), + ], + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache"), + patch("litellm.proxy.proxy_server.proxy_logging_obj"), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + new_callable=AsyncMock, + return_value=mock_user, + ), + ): + mock_db = Mock() + mock_prisma.db = mock_db + + mock_team_1 = Mock() + mock_team_1.model_dump.return_value = { + "team_id": "team_1", + "team_alias": "Team One", + "organization_id": "org_A", + "members_with_roles": [{"user_id": "org_admin_user", "role": "admin"}], + } + mock_team_2 = Mock() + mock_team_2.model_dump.return_value = { + "team_id": "team_2", + "team_alias": "Team Two", + "organization_id": "org_A", + "members_with_roles": [{"user_id": "other_user", "role": "user"}], + } + mock_db.litellm_teamtable.find_many = AsyncMock( + return_value=[mock_team_1, mock_team_2] + ) + mock_db.litellm_teamtable.count = AsyncMock(return_value=2) + mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[]) + + # UI sends the caller's own user_id for non-Admin roles + result = await list_team_v2( + http_request=mock_request, + user_id="org_admin_user", # same as caller — UI sends this + organization_id=None, + team_id=None, + team_alias=None, + user_api_key_dict=mock_user_api_key_dict, + page=1, + page_size=10, + sort_by=None, + sort_order="asc", + status=None, + ) + + assert result["total"] == 2 + assert len(result["teams"]) == 2 + + # Verify the where clause scopes by org only — no team_id filter + where = mock_db.litellm_teamtable.find_many.call_args.kwargs["where"] + assert where["organization_id"] == {"in": ["org_A"]} + assert "team_id" not in where + + @pytest.mark.asyncio async def test_list_team_v2_org_admin_cannot_view_other_orgs(): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 2efec3e0b34..acca357e641 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -2120,6 +2120,292 @@ class TestCLIKeyRegenerationFlow: assert exc_info.value.status_code == 429 mock_cache.set_cache.assert_not_called() + @pytest.mark.asyncio + async def test_cli_sso_start_returns_verification_uri_complete_when_enabled(self): + """Test CLI SSO start returns a verification_uri_complete that round-trips the user_code only when the operator opts in""" + from urllib.parse import parse_qs, urlparse + + from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER + from litellm.proxy.management_endpoints.ui_sso import cli_sso_start + + mock_request = MagicMock(spec=Request) + mock_request.client = SimpleNamespace(host="127.0.0.1") + mock_request.headers = {} + mock_request.base_url = "https://proxy.example.com/" + mock_cache = MagicMock() + mock_cache.increment_cache.return_value = 1 + + with ( + patch.dict( + os.environ, + {"PROXY_BASE_URL": "https://proxy.example.com", "SERVER_ROOT_PATH": ""}, + ), + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_cli_sso_verification_uri_complete": True}, + ), + ): + result = await cli_sso_start(request=mock_request) + + verification_uri_complete = result["verification_uri_complete"] + parsed = urlparse(verification_uri_complete) + query = parse_qs(parsed.query) + + assert parsed.path.endswith("/sso/key/generate") + assert query["source"] == [LITELLM_CLI_SOURCE_IDENTIFIER] + assert query["key"] == [result["login_id"]] + assert query["user_code"] == [result["user_code"]] + + @pytest.mark.asyncio + async def test_cli_sso_start_omits_verification_uri_complete_by_default(self): + """Test CLI SSO start does NOT advertise verification_uri_complete unless the operator enables it (default off)""" + from litellm.proxy.management_endpoints.ui_sso import cli_sso_start + + mock_request = MagicMock(spec=Request) + mock_request.client = SimpleNamespace(host="127.0.0.1") + mock_request.headers = {} + mock_request.base_url = "https://proxy.example.com/" + mock_cache = MagicMock() + mock_cache.increment_cache.return_value = 1 + + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + result = await cli_sso_start(request=mock_request) + + assert "verification_uri_complete" not in result + assert result["user_code"] + assert result["login_id"].startswith("cli-") + + def test_cli_sso_verification_uri_complete_enabled_reads_general_settings(self): + """Test the operator opt-in flag is read from general_settings and defaults off""" + from litellm.proxy.management_endpoints.ui_sso import ( + _cli_sso_verification_uri_complete_enabled, + ) + + with patch("litellm.proxy.proxy_server.general_settings", {}): + assert _cli_sso_verification_uri_complete_enabled() is False + with patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_cli_sso_verification_uri_complete": True}, + ): + assert _cli_sso_verification_uri_complete_enabled() is True + + @pytest.mark.asyncio + async def test_google_login_only_threads_user_code_when_enabled(self): + """Test google_login forwards user_code into the OAuth state only when the operator opt-in is on, dropping it otherwise""" + from litellm.proxy.management_endpoints.ui_sso import google_login + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://proxy.example.com/" + mock_cache = MagicMock() + mock_cache.get_cache.return_value = {"poll_secret_hash": "h"} + + async def drive(enabled: bool): + with ( + patch.dict(os.environ, {}, clear=True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch( + "litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", + None, + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_cli_sso_verification_uri_complete": enabled}, + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.show_missing_vars_in_env", + return_value=None, + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.get_redirect_url_for_sso", + return_value="https://proxy.example.com/sso/callback", + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler._get_cli_state", + return_value=None, + ) as mock_get_cli_state, + ): + try: + await google_login( + request=mock_request, + source="litellm-cli", + key="cli-validsessionkey123456", + user_code="WXYZ-2345", + ) + except Exception: + pass + return mock_get_cli_state.call_args.kwargs["user_code"] + + assert await drive(enabled=True) == "WXYZ-2345" + assert await drive(enabled=False) is None + + def test_get_cli_state_appends_user_code_for_prefill(self): + """Test the OAuth state carries the user_code only for the opt-in prefill flow""" + from litellm.constants import ( + LITELLM_CLI_SESSION_TOKEN_PREFIX, + LITELLM_CLI_SOURCE_IDENTIFIER, + ) + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + manual_state = SSOAuthenticationHandler._get_cli_state( + source=LITELLM_CLI_SOURCE_IDENTIFIER, key="cli-abc123" + ) + prefill_state = SSOAuthenticationHandler._get_cli_state( + source=LITELLM_CLI_SOURCE_IDENTIFIER, + key="cli-abc123", + user_code="WXYZ-2345", + ) + + assert manual_state == f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123" + assert ( + prefill_state == f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123:WXYZ-2345" + ) + assert ( + SSOAuthenticationHandler._get_cli_state( + source="not-cli", key="cli-abc123", user_code="WXYZ-2345" + ) + is None + ) + + def test_get_cli_state_drops_malformed_user_code(self): + """Test a user_code that is not a server-issued code is dropped before reaching the size-limited OAuth state""" + from litellm.constants import ( + LITELLM_CLI_SESSION_TOKEN_PREFIX, + LITELLM_CLI_SOURCE_IDENTIFIER, + ) + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + manual_only = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123" + for bad_user_code in ("A" * 4096, "not-a-code", "WXYZ2345", "WXYZ-234", ""): + assert ( + SSOAuthenticationHandler._get_cli_state( + source=LITELLM_CLI_SOURCE_IDENTIFIER, + key="cli-abc123", + user_code=bad_user_code, + ) + == manual_only + ) + + def test_is_valid_cli_sso_user_code_matches_generated_format(self): + """Test the user_code validator accepts a freshly generated code and rejects malformed input""" + from litellm.proxy.management_endpoints.ui_sso import ( + _generate_cli_sso_user_code, + _is_valid_cli_sso_user_code, + ) + + assert _is_valid_cli_sso_user_code(_generate_cli_sso_user_code()) + assert _is_valid_cli_sso_user_code("WXYZ-2345") + assert not _is_valid_cli_sso_user_code("WXYZ-2340") # 0 is not in the alphabet + assert not _is_valid_cli_sso_user_code("wxyz-2345") + assert not _is_valid_cli_sso_user_code("WXYZ2345") + assert not _is_valid_cli_sso_user_code("A" * 64) + assert not _is_valid_cli_sso_user_code(None) + + def test_cli_state_round_trips_user_code_to_callback_parser(self): + """Test the callback's state parser recovers login_id and user_code from the state _get_cli_state builds""" + from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + state = SSOAuthenticationHandler._get_cli_state( + source=LITELLM_CLI_SOURCE_IDENTIFIER, + key="cli-abc123", + user_code="WXYZ-2345", + ) + + state_parts = state.split(":", 2) + key_id = state_parts[1] if len(state_parts) > 1 else None + prefill_user_code = state_parts[2] if len(state_parts) > 2 else None + + assert key_id == "cli-abc123" + assert prefill_user_code == "WXYZ-2345" + + def test_render_cli_sso_verification_page_prefills_user_code(self): + """Test the verify page pre-fills the user_code input (HTML-escaped) when provided""" + from litellm.proxy.management_endpoints.ui_sso import ( + _render_cli_sso_verification_page, + ) + + html = _render_cli_sso_verification_page( + verify_url="https://proxy.example.com/sso/cli/complete/cli-abc123", + browser_complete_token="browser-token", + prefill_user_code='WXYZ-2345">