From cafc8c1455a7691b4cf2082bc809abeb3cc45af4 Mon Sep 17 00:00:00 2001 From: milan Date: Fri, 21 Aug 2026 03:01:26 +0000 Subject: [PATCH 01/93] fix(proxy): store the actual selected model in spend logs for Azure Model Router Co-authored-by: Filippo Mattia Menghi Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../spend_tracking/spend_tracking_utils.py | 4 +- .../test_spend_tracking_utils.py | 42 +++++++++++++++++++ 2 files changed, 45 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 0b56f0d8246..822f03873b3 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -411,7 +411,9 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs or None ) raw_model: Final = cast(str, kwargs.get("model") or "") - model_name: Final = reconstruct_model_name(raw_model, custom_llm_provider, metadata or {}) + model_name: Final = ( + standard_logging_payload.get("model") if standard_logging_payload is not None else None + ) or reconstruct_model_name(raw_model, custom_llm_provider, metadata or {}) try: payload: Final[SpendLogsPayload] = SpendLogsPayload( diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index 9710dc44e99..b1a45fb84a3 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -3241,3 +3241,45 @@ def test_get_logging_payload_failed_request_without_standard_logging_payload_lea assert payload["model_group"] == "" assert payload["api_base"] == "" assert payload["custom_llm_provider"] == "" + + +def _model_router_spend_log_kwargs(slp_model: str | None) -> dict[str, Any]: + standard_logging_payload: Final = cast( + StandardLoggingPayload, + { + "model": slp_model, + "metadata": {}, + "model_map_information": StandardLoggingModelInformation( + model_map_key="azure_ai/model_router", model_map_value=None + ), + }, + ) + return { + "model": "azure_ai/model_router/model-router", + "litellm_params": {"metadata": {"user_api_key": "sk-test-key"}}, + "standard_logging_object": standard_logging_payload, + } + + +@patch("litellm.proxy.proxy_server.master_key", None) +@patch("litellm.proxy.proxy_server.general_settings", {}) +def test_get_logging_payload_uses_standard_logging_payload_model(): + payload = get_logging_payload( + kwargs=_model_router_spend_log_kwargs(slp_model="azure_ai/gpt-5-mini"), + response_obj={}, + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), + ) + assert payload["model"] == "azure_ai/gpt-5-mini" + + +@patch("litellm.proxy.proxy_server.master_key", None) +@patch("litellm.proxy.proxy_server.general_settings", {}) +def test_get_logging_payload_falls_back_to_kwargs_model_when_slp_model_missing(): + payload = get_logging_payload( + kwargs=_model_router_spend_log_kwargs(slp_model=None), + response_obj={}, + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), + ) + assert payload["model"] == "azure_ai/model_router/model-router" From 57b367c78e6f691839a4c6dccf8ffe57bfb25478 Mon Sep 17 00:00:00 2001 From: milan Date: Fri, 21 Aug 2026 03:25:06 +0000 Subject: [PATCH 02/93] refactor(tests): type the model router spend log kwargs helper Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/spend_tracking/test_spend_tracking_utils.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index b1a45fb84a3..9c97b2683b2 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -6,6 +6,8 @@ import sys from datetime import timezone from typing import Any, Final, cast +from typing_extensions import ReadOnly, TypedDict + import pytest from fastapi.testclient import TestClient @@ -3243,7 +3245,13 @@ def test_get_logging_payload_failed_request_without_standard_logging_payload_lea assert payload["custom_llm_provider"] == "" -def _model_router_spend_log_kwargs(slp_model: str | None) -> dict[str, Any]: +class _ModelRouterSpendLogKwargs(TypedDict): + model: ReadOnly[str] + litellm_params: ReadOnly[dict[str, dict[str, str]]] + standard_logging_object: ReadOnly[StandardLoggingPayload] + + +def _model_router_spend_log_kwargs(slp_model: str | None) -> _ModelRouterSpendLogKwargs: standard_logging_payload: Final = cast( StandardLoggingPayload, { From 4e37425a782b7bf09d18e42ef9f072f6108155e4 Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Fri, 21 Aug 2026 15:26:42 -0700 Subject: [PATCH 03/93] ci: ban row-rewriting DML from prisma migrations --- .../check_migrations_no_data_rewrites.py | 318 ++++++++++++++++++ .../test_check_migrations_no_data_rewrites.py | 234 +++++++++++++ 2 files changed, 552 insertions(+) create mode 100644 tests/code_coverage_tests/check_migrations_no_data_rewrites.py create mode 100644 tests/test_litellm/test_check_migrations_no_data_rewrites.py diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py new file mode 100644 index 00000000000..fc0eafb3b16 --- /dev/null +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -0,0 +1,318 @@ +#!/usr/bin/env python3 +"""Ban row-rewriting DML from Prisma migrations. + +Migrations run synchronously at proxy boot, before the process serves traffic, so +anything whose cost scales with existing table size turns into downtime. A single +`UPDATE` with no batching over a spend-log-sized table is minutes of unavailability +plus a doubled heap that plain autovacuum will not give back. + +Flagged, per statement, by its leading keyword: + + UPDATE rewrites every matching row, and `WHERE` does not bound the scan + DELETE same scan, and the dead tuples outlive the migration + INSERT only when it draws rows from a `SELECT`; `INSERT ... VALUES` is bounded + by the literal row list and passes + WITH a CTE-led statement containing any of the above + +Referential actions (`ON DELETE CASCADE`, `ON UPDATE CASCADE`) are schema, never a +statement's leading keyword, so they pass. + +Statements inside dollar-quoted bodies are scanned too. `DO $$ ... $$` is this +repo's idiom for conditional DDL, so a body is where an `UPDATE` would otherwise +hide. + +Add a column and let the application populate it, or run the rewrite as an opt-in +batched job outside boot. When a rewrite is genuinely bounded and must ship inside +the migration, put `-- data-migration-ok: ` on the statement, naming what +bounds it. The reason is required. + +`GRANDFATHERED` freezes the violations that predate this check. Prisma records a +checksum for every applied migration and this repo treats applied files as +immutable, so those two cannot take an inline marker. The set is closed; a new +migration belongs nowhere in it. +""" + +from __future__ import annotations + +import re +import sys +from collections.abc import Iterator, Mapping +from dataclasses import dataclass +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[2] +MIGRATIONS_DIR = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations" + +GRANDFATHERED = frozenset( + { + "20260817000000_shadow_eval_multi_key", + "20260818224500_add_shadow_eval_stopped_by", + } +) + +MARKER = re.compile(r"--[ \t]*data-migration-ok:[ \t]*(\S.*?)[ \t]*$", re.MULTILINE) +DOLLAR_TAG = re.compile(r"\$(?:[A-Za-z_][A-Za-z0-9_]*)?\$") +FIRST_WORD = re.compile(r"[A-Za-z_][A-Za-z0-9_]*") +STATEMENT = re.compile(r"[^;]+") + +REWRITES_ROWS = frozenset({"UPDATE", "DELETE", "MERGE"}) + +STATEMENT_KEYWORDS = REWRITES_ROWS | frozenset( + { + "INSERT", + "SELECT", + "WITH", + "ALTER", + "CREATE", + "DROP", + "TRUNCATE", + "COMMENT", + "GRANT", + "REVOKE", + "COPY", + "SET", + "PERFORM", + "RAISE", + "RETURN", + "EXECUTE", + "CALL", + "REINDEX", + "REFRESH", + "VACUUM", + "ANALYZE", + } +) + +GUIDANCE = """ +Migrations apply at proxy boot, before it serves traffic, so a statement whose cost +scales with table size is downtime. Add the column and let the application backfill +it, or move the rewrite to a batched job outside boot. + +If the rewrite is genuinely bounded and has to ship in the migration, mark the +statement with the bound spelled out: + + -- data-migration-ok: + UPDATE ... +""" + + +@dataclass(frozen=True, slots=True) +class Violation: + migration: str + line: int + keyword: str + + def render(self) -> str: + location = f"{MIGRATIONS_DIR.relative_to(REPO_ROOT)}/{self.migration}/migration.sql" + return f"{location}:{self.line}: {self.keyword} rewrites existing rows at boot" + + +def blank(text: str) -> str: + return "".join(character if character == "\n" else " " for character in text) + + +def mask(sql: str) -> tuple[str, tuple[tuple[int, int], ...]]: + """Blank comments and quoted text, keeping offsets, and locate dollar-quoted bodies.""" + chunks: list[str] = [] + bodies: list[tuple[int, int]] = [] + index = 0 + length = len(sql) + + while index < length: + pair = sql[index : index + 2] + + if pair == "--": + stop = sql.find("\n", index) + stop = length if stop == -1 else stop + chunks.append(blank(sql[index:stop])) + index = stop + continue + + if pair == "/*": + stop = skip_block_comment(sql, index) + chunks.append(blank(sql[index:stop])) + index = stop + continue + + character = sql[index] + + if character in "'\"": + stop = skip_quoted(sql, index, character) + chunks.append(blank(sql[index:stop])) + index = stop + continue + + if character == "$": + tag = DOLLAR_TAG.match(sql, index) + if tag is not None: + closing = sql.find(tag.group(), tag.end()) + body_end = length if closing == -1 else closing + stop = length if closing == -1 else closing + len(tag.group()) + bodies.append((tag.end(), body_end)) + chunks.append(blank(sql[index:stop])) + index = stop + continue + + chunks.append(character) + index += 1 + + return "".join(chunks), tuple(bodies) + + +def skip_block_comment(sql: str, start: int) -> int: + depth = 1 + index = start + 2 + while index < len(sql) and depth > 0: + pair = sql[index : index + 2] + if pair == "/*": + depth += 1 + index += 2 + elif pair == "*/": + depth -= 1 + index += 2 + else: + index += 1 + return index + + +def skip_quoted(sql: str, start: int, quote: str) -> int: + index = start + 1 + while index < len(sql): + if sql[index] != quote: + index += 1 + elif sql[index + 1 : index + 2] == quote: + index += 2 + else: + return index + 1 + return len(sql) + + +def strip_parens(statement: str) -> str: + """Blank parenthesised groups in place, so an `IF EXISTS (SELECT ...)` guard does not + stand in for the statement it guards.""" + chunks: list[str] = [] + depth = 0 + + for character in statement: + if character == "(": + depth += 1 + chunks.append(" ") + elif character == ")": + depth = max(depth - 1, 0) + chunks.append(" ") + elif depth > 0 and character != "\n": + chunks.append(" ") + else: + chunks.append(character) + + return "".join(chunks) + + +def leading_keyword(statement: str) -> re.Match[str] | None: + """The statement's own keyword, looking past PL/pgSQL block syntax such as + `BEGIN`, `IF ... THEN` and `END`.""" + return next( + (word for word in FIRST_WORD.finditer(statement) if word.group().upper() in STATEMENT_KEYWORDS), + None, + ) + + +def offending_keyword(statement: str) -> str | None: + word = leading_keyword(strip_parens(statement)) + if word is None: + return None + + keyword = word.group().upper() + + if keyword in REWRITES_ROWS: + return keyword + + if keyword == "INSERT": + return "INSERT ... SELECT" if contains(statement, "SELECT") else None + + if keyword == "WITH": + nested = next((name for name in sorted(REWRITES_ROWS) if contains(statement, name)), None) + if nested is not None: + return f"WITH ... {nested}" + if contains(statement, "INSERT") and contains(statement, "SELECT"): + return "WITH ... INSERT ... SELECT" + + return None + + +def contains(statement: str, keyword: str) -> bool: + return re.search(rf"\b{keyword}\b", statement, re.IGNORECASE) is not None + + +def exempt_lines(sql: str) -> frozenset[int]: + return frozenset(sql.count("\n", 0, match.start()) + 1 for match in MARKER.finditer(sql)) + + +def scan(sql: str, migration: str, exempt: frozenset[int], offset: int = 0) -> Iterator[Violation]: + masked, bodies = mask(sql) + + for match in STATEMENT.finditer(masked): + keyword = offending_keyword(match.group()) + if keyword is None: + continue + first = line_of(sql, offset + keyword_start(match)) + last = line_of(sql, offset + match.end()) + if any(line in exempt for line in range(first - 1, last + 1)): + continue + yield Violation(migration, first, keyword) + + for start, end in bodies: + yield from scan(sql[start:end], migration, exempt, offset + start) + + +def keyword_start(statement: re.Match[str]) -> int: + word = leading_keyword(strip_parens(statement.group())) + return statement.start() + (0 if word is None else word.start()) + + +def line_of(sql: str, offset: int) -> int: + return sql.count("\n", 0, offset) + 1 + + +def scan_migration(directory: Path) -> tuple[Violation, ...]: + sql = (directory / "migration.sql").read_text(encoding="utf-8") + return tuple(scan(sql, directory.name, exempt_lines(sql))) + + +def stale_grandfathers(found: Mapping[str, tuple[Violation, ...]]) -> tuple[str, ...]: + clean = (name for name in GRANDFATHERED & found.keys() if not found[name]) + missing = GRANDFATHERED - found.keys() + return tuple(sorted((*clean, *missing))) + + +def main() -> int: + if not MIGRATIONS_DIR.is_dir(): + print(f"migrations directory not found: {MIGRATIONS_DIR}", file=sys.stderr) + return 2 + + directories = tuple(sorted(path for path in MIGRATIONS_DIR.iterdir() if (path / "migration.sql").is_file())) + found = {directory.name: scan_migration(directory) for directory in directories} + violations = tuple( + violation for name, results in found.items() if name not in GRANDFATHERED for violation in results + ) + + for violation in violations: + print(violation.render()) + + stale = stale_grandfathers(found) + for name in stale: + print(f"{name}: listed in GRANDFATHERED but no longer violates; remove it from the set") + + if violations: + print(GUIDANCE, file=sys.stderr) + print(f"{len(violations)} data-rewriting statement(s) in migrations.", file=sys.stderr) + + if violations or stale: + return 1 + + print(f"No data-rewriting statements in {len(directories)} migrations.") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py new file mode 100644 index 00000000000..40b55d741bc --- /dev/null +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -0,0 +1,234 @@ +"""Tests for tests/code_coverage_tests/check_migrations_no_data_rewrites.py. + +The checker reads migration.sql as SQL rather than as text, so the cases that matter +are the ones a grep would get wrong: `ON DELETE CASCADE` in a foreign key (60-odd +occurrences in the shipped migrations), an `UPDATE` inside a string literal or a +comment, and an `UPDATE` hidden in the `DO $$ ... $$` block this repo uses for +conditional DDL. +""" + +import importlib.util +import sys +from pathlib import Path + +_CHECKER_PATH = Path(__file__).resolve().parents[1] / "code_coverage_tests" / "check_migrations_no_data_rewrites.py" +_SPEC = importlib.util.spec_from_file_location("check_migrations_no_data_rewrites", _CHECKER_PATH) +assert _SPEC is not None and _SPEC.loader is not None +checker = importlib.util.module_from_spec(_SPEC) +sys.modules[_SPEC.name] = checker +_SPEC.loader.exec_module(checker) + + +def _scan(tmp_path: Path, sql: str) -> tuple: + directory = tmp_path / "20260101000000_fixture" + directory.mkdir(exist_ok=True) + (directory / "migration.sql").write_text(sql, encoding="utf-8") + return checker.scan_migration(directory) + + +def _keywords(tmp_path: Path, sql: str) -> tuple: + return tuple(violation.keyword for violation in _scan(tmp_path, sql)) + + +class TestRowRewritesAreFlagged: + def test_update_is_flagged(self, tmp_path): + assert _keywords(tmp_path, 'UPDATE "Foo" SET "a" = 1;') == ("UPDATE",) + + def test_delete_is_flagged(self, tmp_path): + assert _keywords(tmp_path, 'DELETE FROM "Foo" WHERE "a" IS NULL;') == ("DELETE",) + + def test_update_without_trailing_semicolon_is_flagged(self, tmp_path): + assert _keywords(tmp_path, 'UPDATE "Foo" SET "a" = 1') == ("UPDATE",) + + def test_lowercase_update_is_flagged(self, tmp_path): + assert _keywords(tmp_path, 'update "Foo" set "a" = 1;') == ("UPDATE",) + + def test_merge_is_flagged(self, tmp_path): + sql = 'MERGE INTO "Foo" t USING "Bar" s ON t."id" = s."id" WHEN MATCHED THEN UPDATE SET "a" = s."a";' + assert _keywords(tmp_path, sql) == ("MERGE",) + + def test_every_offending_statement_is_reported(self, tmp_path): + sql = 'UPDATE "Foo" SET "a" = 1;\nALTER TABLE "Foo" ADD COLUMN "b" TEXT;\nDELETE FROM "Bar";' + assert _keywords(tmp_path, sql) == ("UPDATE", "DELETE") + + def test_the_incident_migration_is_flagged(self, tmp_path): + sql = ( + 'UPDATE "LiteLLM_SpendLogs"\n' + ' SET "created_at" = "endTime",\n' + ' "updated_at" = "endTime"\n' + ' WHERE "created_at" > "endTime" + interval \'1 hour\';\n' + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + +class TestSchemaStatementsPass: + def test_on_delete_cascade_is_not_a_data_rewrite(self, tmp_path): + sql = ( + 'ALTER TABLE "A" ADD CONSTRAINT "A_b_fkey" FOREIGN KEY ("b") ' + 'REFERENCES "B"("id") ON DELETE CASCADE ON UPDATE CASCADE;' + ) + assert _keywords(tmp_path, sql) == () + + def test_on_delete_set_null_is_not_a_data_rewrite(self, tmp_path): + sql = ( + 'ALTER TABLE "A" ADD CONSTRAINT "A_b_fkey" FOREIGN KEY ("b") ' + 'REFERENCES "B"("id") ON DELETE SET NULL ON UPDATE CASCADE;' + ) + assert _keywords(tmp_path, sql) == () + + def test_add_column_with_default_passes(self, tmp_path): + sql = 'ALTER TABLE "Foo" ADD COLUMN IF NOT EXISTS "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP;' + assert _keywords(tmp_path, sql) == () + + def test_drop_table_passes(self, tmp_path): + assert _keywords(tmp_path, 'DROP TABLE IF EXISTS "Foo";') == () + + def test_empty_file_passes(self, tmp_path): + assert _keywords(tmp_path, "") == () + + def test_only_comments_passes(self, tmp_path): + assert _keywords(tmp_path, "-- nothing to do here\n") == () + + +class TestInsert: + def test_insert_values_is_bounded_and_passes(self, tmp_path): + assert _keywords(tmp_path, "INSERT INTO \"Foo\" (\"id\") VALUES ('a'), ('b');") == () + + def test_insert_select_scans_and_is_flagged(self, tmp_path): + assert _keywords(tmp_path, 'INSERT INTO "Foo" ("id") SELECT "id" FROM "Bar";') == ("INSERT ... SELECT",) + + +class TestCommonTableExpressions: + def test_cte_led_update_is_flagged(self, tmp_path): + sql = 'WITH batch AS (SELECT "id" FROM "Foo" LIMIT 100) UPDATE "Foo" SET "a" = 1 FROM batch;' + assert _keywords(tmp_path, sql) == ("WITH ... UPDATE",) + + def test_cte_led_delete_is_flagged(self, tmp_path): + sql = 'WITH batch AS (SELECT "id" FROM "Foo" LIMIT 100) DELETE FROM "Foo" USING batch;' + assert _keywords(tmp_path, sql) == ("WITH ... DELETE",) + + def test_cte_led_insert_select_is_flagged(self, tmp_path): + sql = 'WITH batch AS (SELECT "id" FROM "Bar") INSERT INTO "Foo" ("id") SELECT "id" FROM batch;' + assert _keywords(tmp_path, sql) == ("WITH ... INSERT ... SELECT",) + + def test_read_only_cte_passes(self, tmp_path): + sql = 'WITH batch AS (SELECT "id" FROM "Foo") SELECT count(*) FROM batch;' + assert _keywords(tmp_path, sql) == () + + +class TestDollarQuotedBlocks: + def test_update_inside_do_block_is_flagged(self, tmp_path): + sql = 'DO $$\nBEGIN\n UPDATE "Foo" SET "a" = 1;\nEND $$;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_conditional_ddl_do_block_passes(self, tmp_path): + sql = ( + "DO $$\n" + "BEGIN\n" + " IF EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'x') THEN\n" + ' ALTER TABLE "Foo" DROP CONSTRAINT "x";\n' + " END IF;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_tagged_dollar_quote_is_scanned(self, tmp_path): + sql = 'DO $body$\nBEGIN\n DELETE FROM "Foo";\nEND $body$;' + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_semicolons_inside_do_block_do_not_split_outer_statements(self, tmp_path): + sql = 'DO $$ BEGIN PERFORM 1; END $$;\nALTER TABLE "Foo" ADD COLUMN "b" TEXT;' + assert _keywords(tmp_path, sql) == () + + +class TestQuotingAndComments: + def test_update_inside_string_literal_passes(self, tmp_path): + sql = "ALTER TABLE \"Foo\" ADD COLUMN \"note\" TEXT NOT NULL DEFAULT 'UPDATE nothing';" + assert _keywords(tmp_path, sql) == () + + def test_escaped_quote_inside_string_does_not_leak(self, tmp_path): + sql = "ALTER TABLE \"Foo\" ADD COLUMN \"note\" TEXT NOT NULL DEFAULT 'it''s fine';\n" + assert _keywords(tmp_path, sql) == () + + def test_update_inside_line_comment_passes(self, tmp_path): + assert _keywords(tmp_path, '-- UPDATE "Foo" SET "a" = 1;\nDROP TABLE "Bar";') == () + + def test_update_inside_block_comment_passes(self, tmp_path): + assert _keywords(tmp_path, '/* UPDATE "Foo" SET "a" = 1; */\nDROP TABLE "Bar";') == () + + def test_nested_block_comment_passes(self, tmp_path): + sql = '/* outer /* UPDATE "Foo" SET "a" = 1; */ still comment */\nDROP TABLE "Bar";' + assert _keywords(tmp_path, sql) == () + + def test_update_inside_quoted_identifier_passes(self, tmp_path): + assert _keywords(tmp_path, 'ALTER TABLE "UPDATE Foo" ADD COLUMN "b" TEXT;') == () + + def test_positional_parameter_is_not_a_dollar_quote(self, tmp_path): + sql = 'ALTER TABLE "Foo" ADD COLUMN "b" TEXT;\nUPDATE "Foo" SET "b" = $1;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + + +class TestEscapeHatch: + def test_marker_with_reason_exempts_the_statement(self, tmp_path): + sql = '-- data-migration-ok: one row per tenant, at most a few hundred\nUPDATE "Foo" SET "a" = 1;' + assert _keywords(tmp_path, sql) == () + + def test_marker_without_reason_does_not_exempt(self, tmp_path): + assert _keywords(tmp_path, '-- data-migration-ok:\nUPDATE "Foo" SET "a" = 1;') == ("UPDATE",) + + def test_marker_exempts_only_its_own_statement(self, tmp_path): + sql = ( + "-- data-migration-ok: bounded to in-flight jobs\n" + 'UPDATE "Foo" SET "a" = 1;\n' + 'UPDATE "Bar" SET "b" = 2;\n' + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 3 + + def test_marker_works_inside_a_do_block(self, tmp_path): + sql = 'DO $$\nBEGIN\n -- data-migration-ok: single row\n UPDATE "Foo" SET "a" = 1;\nEND $$;' + assert _keywords(tmp_path, sql) == () + + def test_marker_below_the_statement_does_not_exempt_the_next_one(self, tmp_path): + sql = 'UPDATE "Foo" SET "a" = 1;\n-- data-migration-ok: bounded\nALTER TABLE "Bar" ADD COLUMN "b" TEXT;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + + +class TestReporting: + def test_line_number_points_at_the_statement_keyword(self, tmp_path): + sql = '-- CreateIndex\nCREATE INDEX "i" ON "Foo"("a");\n\nUPDATE "Foo" SET "a" = 1;' + assert _scan(tmp_path, sql)[0].line == 4 + + def test_render_names_the_migration_and_line(self, tmp_path): + violation = _scan(tmp_path, '\n\nDELETE FROM "Foo";')[0] + rendered = violation.render() + assert "20260101000000_fixture/migration.sql:3" in rendered + assert "DELETE" in rendered + + +class TestGrandfathering: + def test_every_grandfathered_migration_still_violates(self): + for name in sorted(checker.GRANDFATHERED): + directory = checker.MIGRATIONS_DIR / name + assert directory.is_dir(), f"{name} no longer exists; drop it from GRANDFATHERED" + assert checker.scan_migration(directory), f"{name} is clean; drop it from GRANDFATHERED" + + def test_stale_entry_is_reported_when_a_migration_stops_violating(self): + found = {name: () for name in checker.GRANDFATHERED} + assert checker.stale_grandfathers(found) == tuple(sorted(checker.GRANDFATHERED)) + + def test_missing_entry_is_reported(self): + assert checker.stale_grandfathers({}) == tuple(sorted(checker.GRANDFATHERED)) + + def test_no_stale_entries_against_the_real_tree(self): + found = { + path.name: checker.scan_migration(path) + for path in checker.MIGRATIONS_DIR.iterdir() + if (path / "migration.sql").is_file() + } + assert checker.stale_grandfathers(found) == () + + +class TestShippedMigrations: + def test_the_repo_is_clean(self): + assert checker.main() == 0 From 6c7dfbd2498a9b17d5f1afb57434cc9166fdbdcf Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Fri, 21 Aug 2026 15:27:56 -0700 Subject: [PATCH 04/93] ci: run the migration data-rewrite check in code quality --- .github/workflows/test-code-quality.yml | 3 +++ CLAUDE.md | 2 ++ 2 files changed, 5 insertions(+) diff --git a/.github/workflows/test-code-quality.yml b/.github/workflows/test-code-quality.yml index 8f62837d29a..d0ac0b6fdee 100644 --- a/.github/workflows/test-code-quality.yml +++ b/.github/workflows/test-code-quality.yml @@ -128,6 +128,9 @@ jobs: - name: check_e2e_no_raw_requests run: uv run --no-sync python ./tests/code_coverage_tests/check_e2e_no_raw_requests.py + - name: check_migrations_no_data_rewrites + run: uv run --no-sync python ./tests/code_coverage_tests/check_migrations_no_data_rewrites.py + - name: memory_test run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py diff --git a/CLAUDE.md b/CLAUDE.md index b3383b4a895..4a661f6effe 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -79,6 +79,8 @@ Do not put names of customers or customer company names in code, PR descriptions CI supply-chain safety: Never pipe a remote script into a shell (`curl ... | bash`, `wget ... | sh`); download the artifact to a file, verify its SHA-256 checksum, then install. Pin every external tool to a specific version with a full URL (not `latest` or `stable`). Verify checksums for all downloaded binaries, using the provider's official `.sha256` / `.sha256sum` sidecar when available. These rules apply to every download in CI +Prisma migrations apply synchronously at proxy boot, before it serves traffic, so a migration must only change schema, never rewrite rows. No `UPDATE`, `DELETE` or `MERGE`, and no `INSERT ... SELECT`: on a spend-log-sized table any of those is minutes of downtime plus a doubled heap that plain autovacuum won't give back. Add the column and let the application populate it, or run the rewrite as an opt-in batched job outside boot. `tests/code_coverage_tests/check_migrations_no_data_rewrites.py` enforces this; when a rewrite is genuinely bounded and has to ship inside the migration, mark the statement `-- data-migration-ok: ` + Follow these coding conventions for new/updated code (a three-line fix in a legacy file shouldn't trigger huge drive-by refactors): - Composition over inheritance From 777eb8af107bf03eba4120a510d15bb1180a6af8 Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Fri, 21 Aug 2026 15:37:05 -0700 Subject: [PATCH 05/93] test: close the sql-lexing gaps found by mutation testing Drops the doubled-quote branch in skip_quoted, which masked the same span either way and so could not be covered, and orders the failure report before the guidance text. --- .../check_migrations_no_data_rewrites.py | 18 +++----- .../test_check_migrations_no_data_rewrites.py | 45 ++++++++++++++++--- 2 files changed, 46 insertions(+), 17 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index fc0eafb3b16..61845ac521f 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -10,6 +10,7 @@ Flagged, per statement, by its leading keyword: UPDATE rewrites every matching row, and `WHERE` does not bound the scan DELETE same scan, and the dead tuples outlive the migration + MERGE both of the above in one statement INSERT only when it draws rows from a `SELECT`; `INSERT ... VALUES` is bounded by the literal row list and passes WITH a CTE-led statement containing any of the above @@ -176,15 +177,10 @@ def skip_block_comment(sql: str, start: int) -> int: def skip_quoted(sql: str, start: int, quote: str) -> int: - index = start + 1 - while index < len(sql): - if sql[index] != quote: - index += 1 - elif sql[index + 1 : index + 2] == quote: - index += 2 - else: - return index + 1 - return len(sql) + """One quoted run, up to and including its closing quote. A doubled quote needs no + special case: closing on the first and reopening on the second masks the same span.""" + stop = sql.find(quote, start + 1) + return len(sql) if stop == -1 else stop + 1 def strip_parens(statement: str) -> str: @@ -304,8 +300,8 @@ def main() -> int: print(f"{name}: listed in GRANDFATHERED but no longer violates; remove it from the set") if violations: - print(GUIDANCE, file=sys.stderr) - print(f"{len(violations)} data-rewriting statement(s) in migrations.", file=sys.stderr) + print(f"\n{len(violations)} data-rewriting statement(s) in migrations.") + print(GUIDANCE) if violations or stale: return 1 diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 40b55d741bc..a4841f7a699 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -132,10 +132,36 @@ class TestDollarQuotedBlocks: ) assert _keywords(tmp_path, sql) == () + def test_guarded_update_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "BEGIN\n" + ' IF EXISTS (SELECT 1 FROM "Foo") THEN\n' + ' UPDATE "Foo" SET "a" = 1;\n' + " END IF;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_guard_with_a_nested_call_still_flags_the_update(self, tmp_path): + sql = ( + "DO $$\n" + "BEGIN\n" + ' IF EXISTS (SELECT 1 FROM "Foo" WHERE lower("a") = \'x\' UNION SELECT 1) THEN\n' + ' UPDATE "Foo" SET "a" = 1;\n' + " END IF;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_tagged_dollar_quote_is_scanned(self, tmp_path): sql = 'DO $body$\nBEGIN\n DELETE FROM "Foo";\nEND $body$;' assert _keywords(tmp_path, sql) == ("DELETE",) + def test_tagged_dollar_quote_holds_an_apostrophe(self, tmp_path): + sql = 'INSERT INTO "Foo" ("t") VALUES ($body$don\'t$body$);\nUPDATE "Bar" SET "b" = 1;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_semicolons_inside_do_block_do_not_split_outer_statements(self, tmp_path): sql = 'DO $$ BEGIN PERFORM 1; END $$;\nALTER TABLE "Foo" ADD COLUMN "b" TEXT;' assert _keywords(tmp_path, sql) == () @@ -143,7 +169,7 @@ class TestDollarQuotedBlocks: class TestQuotingAndComments: def test_update_inside_string_literal_passes(self, tmp_path): - sql = "ALTER TABLE \"Foo\" ADD COLUMN \"note\" TEXT NOT NULL DEFAULT 'UPDATE nothing';" + sql = 'ALTER TABLE "Foo" ADD COLUMN "note" TEXT NOT NULL DEFAULT \'UPDATE nothing\';' assert _keywords(tmp_path, sql) == () def test_escaped_quote_inside_string_does_not_leak(self, tmp_path): @@ -160,9 +186,20 @@ class TestQuotingAndComments: sql = '/* outer /* UPDATE "Foo" SET "a" = 1; */ still comment */\nDROP TABLE "Bar";' assert _keywords(tmp_path, sql) == () + def test_nested_block_comment_masks_past_the_inner_close(self, tmp_path): + sql = '/* outer /* inner */ UPDATE "Foo" SET "a" = 1; */\nDROP TABLE "Bar";' + assert _keywords(tmp_path, sql) == () + def test_update_inside_quoted_identifier_passes(self, tmp_path): assert _keywords(tmp_path, 'ALTER TABLE "UPDATE Foo" ADD COLUMN "b" TEXT;') == () + def test_select_in_a_quoted_identifier_does_not_make_an_insert_a_rewrite(self, tmp_path): + assert _keywords(tmp_path, 'INSERT INTO "SELECT Foo" ("id") VALUES (\'a\');') == () + + def test_update_in_a_quoted_identifier_does_not_make_a_cte_a_rewrite(self, tmp_path): + sql = 'WITH batch AS (SELECT "id" FROM "UPDATE Foo") SELECT count(*) FROM batch;' + assert _keywords(tmp_path, sql) == () + def test_positional_parameter_is_not_a_dollar_quote(self, tmp_path): sql = 'ALTER TABLE "Foo" ADD COLUMN "b" TEXT;\nUPDATE "Foo" SET "b" = $1;' assert _keywords(tmp_path, sql) == ("UPDATE",) @@ -177,11 +214,7 @@ class TestEscapeHatch: assert _keywords(tmp_path, '-- data-migration-ok:\nUPDATE "Foo" SET "a" = 1;') == ("UPDATE",) def test_marker_exempts_only_its_own_statement(self, tmp_path): - sql = ( - "-- data-migration-ok: bounded to in-flight jobs\n" - 'UPDATE "Foo" SET "a" = 1;\n' - 'UPDATE "Bar" SET "b" = 2;\n' - ) + sql = '-- data-migration-ok: bounded to in-flight jobs\nUPDATE "Foo" SET "a" = 1;\nUPDATE "Bar" SET "b" = 2;\n' assert _keywords(tmp_path, sql) == ("UPDATE",) assert _scan(tmp_path, sql)[0].line == 3 From 7e6d303e382b386e5d2d562a069e91494196778d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 21 Aug 2026 17:13:47 -0700 Subject: [PATCH 06/93] fix: count migration lines against the whole file, scan EXECUTE'd sql, allow bounded inserts scan() recursed into a dollar-quoted body with the sliced text but kept absolute offsets, so line_of counted newlines in the slice against a position past its end. Any DO $$ block below the first line reported a wrong line, which also misaligned the -- data-migration-ok: markers: an unrelated marker earlier in the file could exempt a rewrite inside a block, and a marker sitting right above one failed to. Line numbers now always count against the whole migration text. EXECUTE was treated as harmless while its quoted SQL was masked, so a rewrite handed over as a string walked through the gate. The literal an EXECUTE runs is now scanned like a dollar-quoted body. INSERT was classified by searching the whole statement for SELECT, so a bounded INSERT ... VALUES holding a scalar subquery, or led by a helper CTE, was flagged as INSERT ... SELECT. A top-level VALUES now bounds the insert, and a VALUES buried in a subquery still does not. --- .../check_migrations_no_data_rewrites.py | 62 ++++++++-- .../test_check_migrations_no_data_rewrites.py | 108 ++++++++++++++++++ 2 files changed, 158 insertions(+), 12 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 61845ac521f..8ebdc43fac2 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -12,7 +12,8 @@ Flagged, per statement, by its leading keyword: DELETE same scan, and the dead tuples outlive the migration MERGE both of the above in one statement INSERT only when it draws rows from a `SELECT`; `INSERT ... VALUES` is bounded - by the literal row list and passes + by the literal row list and passes, scalar subqueries in that list + included WITH a CTE-led statement containing any of the above Referential actions (`ON DELETE CASCADE`, `ON UPDATE CASCADE`) are schema, never a @@ -20,7 +21,12 @@ statement's leading keyword, so they pass. Statements inside dollar-quoted bodies are scanned too. `DO $$ ... $$` is this repo's idiom for conditional DDL, so a body is where an `UPDATE` would otherwise -hide. +hide. The SQL an `EXECUTE` runs is scanned the same way, since a rewrite reads the +same to Postgres whether it is spelled out or handed over as a string. + +Line numbers always count against the whole migration file, however deeply the +statement is nested, so a reported line points at the statement and the markers +below line up with the statements they exempt. Add a column and let the application populate it, or run the rewrite as an opt-in batched job outside boot. When a rewrite is genuinely bounded and must ship inside @@ -112,10 +118,12 @@ def blank(text: str) -> str: return "".join(character if character == "\n" else " " for character in text) -def mask(sql: str) -> tuple[str, tuple[tuple[int, int], ...]]: - """Blank comments and quoted text, keeping offsets, and locate dollar-quoted bodies.""" +def mask(sql: str) -> tuple[str, tuple[tuple[int, int], ...], tuple[tuple[int, int], ...]]: + """Blank comments and quoted text, keeping offsets, and locate the spans that can still + hold SQL: dollar-quoted bodies, and the single-quoted literals `EXECUTE` runs.""" chunks: list[str] = [] bodies: list[tuple[int, int]] = [] + literals: list[tuple[int, int]] = [] index = 0 length = len(sql) @@ -139,6 +147,9 @@ def mask(sql: str) -> tuple[str, tuple[tuple[int, int], ...]]: if character in "'\"": stop = skip_quoted(sql, index, character) + if character == "'": + closed = sql[stop - 1 : stop] == character + literals.append((index + 1, max(index + 1, stop - 1 if closed else stop))) chunks.append(blank(sql[index:stop])) index = stop continue @@ -157,7 +168,7 @@ def mask(sql: str) -> tuple[str, tuple[tuple[int, int], ...]]: chunks.append(character) index += 1 - return "".join(chunks), tuple(bodies) + return "".join(chunks), tuple(bodies), tuple(literals) def skip_block_comment(sql: str, start: int) -> int: @@ -224,18 +235,31 @@ def offending_keyword(statement: str) -> str | None: return keyword if keyword == "INSERT": - return "INSERT ... SELECT" if contains(statement, "SELECT") else None + return "INSERT ... SELECT" if draws_rows_from_a_select(statement) else None if keyword == "WITH": nested = next((name for name in sorted(REWRITES_ROWS) if contains(statement, name)), None) if nested is not None: return f"WITH ... {nested}" - if contains(statement, "INSERT") and contains(statement, "SELECT"): + if contains(statement, "INSERT") and draws_rows_from_a_select(statement): return "WITH ... INSERT ... SELECT" return None +def draws_rows_from_a_select(statement: str) -> bool: + """Whether an `INSERT` takes its rows from a query rather than a literal list. A + top-level `VALUES` bounds the insert to the rows written out there, so the scalar + subqueries and helper CTEs that sit in parentheses around it do not make it a + rewrite.""" + return contains(statement, "SELECT") and not contains(strip_parens(statement), "VALUES") + + +def leads_with(statement: str, keyword: str) -> bool: + word = leading_keyword(strip_parens(statement)) + return word is not None and word.group().upper() == keyword + + def contains(statement: str, keyword: str) -> bool: return re.search(rf"\b{keyword}\b", statement, re.IGNORECASE) is not None @@ -244,21 +268,35 @@ def exempt_lines(sql: str) -> frozenset[int]: return frozenset(sql.count("\n", 0, match.start()) + 1 for match in MARKER.finditer(sql)) -def scan(sql: str, migration: str, exempt: frozenset[int], offset: int = 0) -> Iterator[Violation]: - masked, bodies = mask(sql) +def scan(sql: str, migration: str, exempt: frozenset[int]) -> Iterator[Violation]: + yield from scan_region(sql, sql, migration, exempt, 0) + + +def scan_region( + document: str, region: str, migration: str, exempt: frozenset[int], offset: int +) -> Iterator[Violation]: + """Violations in one region of `document`, whose text begins at `offset`. Lines are + always counted against the whole document, so a statement nested in a dollar-quoted + body reports its real file line and lines up with the markers read from that file.""" + masked, bodies, literals = mask(region) for match in STATEMENT.finditer(masked): + if leads_with(match.group(), "EXECUTE"): + for start, end in literals: + if match.start() <= start and end <= match.end(): + yield from scan_region(document, region[start:end], migration, exempt, offset + start) + continue keyword = offending_keyword(match.group()) if keyword is None: continue - first = line_of(sql, offset + keyword_start(match)) - last = line_of(sql, offset + match.end()) + first = line_of(document, offset + keyword_start(match)) + last = line_of(document, offset + match.end()) if any(line in exempt for line in range(first - 1, last + 1)): continue yield Violation(migration, first, keyword) for start, end in bodies: - yield from scan(sql[start:end], migration, exempt, offset + start) + yield from scan_region(document, region[start:end], migration, exempt, offset + start) def keyword_start(statement: re.Match[str]) -> int: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index a4841f7a699..9c6871cf557 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -97,6 +97,22 @@ class TestInsert: def test_insert_select_scans_and_is_flagged(self, tmp_path): assert _keywords(tmp_path, 'INSERT INTO "Foo" ("id") SELECT "id" FROM "Bar";') == ("INSERT ... SELECT",) + def test_insert_values_with_a_scalar_subquery_passes(self, tmp_path): + sql = 'INSERT INTO "Config" ("k", "v") VALUES (\'rev\', (SELECT max("id")::text FROM "Bar"));' + assert _keywords(tmp_path, sql) == () + + def test_insert_values_with_a_scalar_subquery_per_row_passes(self, tmp_path): + sql = ( + 'INSERT INTO "Config" ("k", "v") VALUES\n' + " ('a', (SELECT \"id\" FROM \"Bar\" WHERE \"n\" = 'a')),\n" + " ('b', (SELECT \"id\" FROM \"Bar\" WHERE \"n\" = 'b'));" + ) + assert _keywords(tmp_path, sql) == () + + def test_values_inside_a_subquery_does_not_bound_an_insert_select(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") SELECT "id" FROM (VALUES (1), (2)) AS "v"("id");' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + class TestCommonTableExpressions: def test_cte_led_update_is_flagged(self, tmp_path): @@ -115,6 +131,10 @@ class TestCommonTableExpressions: sql = 'WITH batch AS (SELECT "id" FROM "Foo") SELECT count(*) FROM batch;' assert _keywords(tmp_path, sql) == () + def test_cte_led_insert_values_is_bounded_and_passes(self, tmp_path): + sql = 'WITH latest AS (SELECT max("id") AS "id" FROM "Bar")\nINSERT INTO "Config" ("k", "v") VALUES (\'rev\', (SELECT "id"::text FROM latest));' + assert _keywords(tmp_path, sql) == () + class TestDollarQuotedBlocks: def test_update_inside_do_block_is_flagged(self, tmp_path): @@ -166,6 +186,32 @@ class TestDollarQuotedBlocks: sql = 'DO $$ BEGIN PERFORM 1; END $$;\nALTER TABLE "Foo" ADD COLUMN "b" TEXT;' assert _keywords(tmp_path, sql) == () + def test_line_number_inside_a_do_block_counts_from_the_top_of_the_file(self, tmp_path): + sql = ( + "-- AlterTable\n" + 'ALTER TABLE "Foo" ADD COLUMN "b" INT;\n' + "\n" + "DO $$\n" + "BEGIN\n" + ' UPDATE "Foo" SET "b" = 1;\n' + "END $$;" + ) + assert _scan(tmp_path, sql)[0].line == 6 + + def test_line_number_inside_a_nested_body_counts_from_the_top_of_the_file(self, tmp_path): + sql = ( + "-- CreateIndex\n" + 'CREATE INDEX "i" ON "Foo"("a");\n' + "\n" + "DO $outer$\n" + "BEGIN\n" + " EXECUTE $inner$\n" + ' UPDATE "Foo" SET "a" = 1\n' + " $inner$;\n" + "END $outer$;" + ) + assert _scan(tmp_path, sql)[0].line == 7 + class TestQuotingAndComments: def test_update_inside_string_literal_passes(self, tmp_path): @@ -226,6 +272,68 @@ class TestEscapeHatch: sql = 'UPDATE "Foo" SET "a" = 1;\n-- data-migration-ok: bounded\nALTER TABLE "Bar" ADD COLUMN "b" TEXT;' assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_marker_inside_a_do_block_below_the_first_line_exempts(self, tmp_path): + sql = ( + "-- AlterTable\n" + 'ALTER TABLE "Foo" ADD COLUMN "b" INT;\n' + "\n" + "DO $$\n" + "BEGIN\n" + " -- data-migration-ok: one config row\n" + ' UPDATE "Foo" SET "b" = 1;\n' + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_marker_above_a_do_block_does_not_exempt_a_rewrite_inside_it(self, tmp_path): + sql = ( + "-- data-migration-ok: bounded, this belongs to the insert below\n" + "INSERT INTO \"Config\" (\"k\") VALUES ('x');\n" + "\n" + 'DO $$ BEGIN UPDATE "Foo" SET "b" = 1; END $$;' + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 4 + + +class TestDynamicSql: + def test_execute_of_a_quoted_update_is_flagged(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'UPDATE \"Foo\" SET \"a\" = 1';\nEND $$;" + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 3 + + def test_execute_of_a_quoted_delete_is_flagged(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'DELETE FROM \"Foo\"';\nEND $$;" + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_execute_of_a_formatted_update_is_flagged(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE format('UPDATE %I SET \"a\" = 1', 'Foo');\nEND $$;" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_execute_of_a_dollar_quoted_update_is_flagged(self, tmp_path): + sql = 'DO $outer$\nBEGIN\n EXECUTE $q$UPDATE "Foo" SET "a" = 1$q$;\nEND $outer$;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_doubled_quote_inside_executed_sql_does_not_hide_the_rewrite(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'UPDATE \"Foo\" SET \"a\" = date_trunc(''day'', \"t\")';\nEND $$;" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_execute_of_ddl_passes(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'ALTER TABLE \"Foo\" ADD COLUMN \"b\" TEXT';\nEND $$;" + assert _keywords(tmp_path, sql) == () + + def test_execute_of_a_read_only_query_passes(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'SELECT 1';\nEND $$;" + assert _keywords(tmp_path, sql) == () + + def test_a_marker_exempts_an_executed_rewrite(self, tmp_path): + sql = "DO $$\nBEGIN\n -- data-migration-ok: one row\n EXECUTE 'UPDATE \"Foo\" SET \"a\" = 1';\nEND $$;" + assert _keywords(tmp_path, sql) == () + + def test_a_literal_that_is_not_executed_is_still_inert(self, tmp_path): + sql = "INSERT INTO \"Foo\" (\"note\") VALUES ('UPDATE \"Bar\" SET \"a\" = 1');" + assert _keywords(tmp_path, sql) == () + class TestReporting: def test_line_number_points_at_the_statement_keyword(self, tmp_path): From e7dea842c36be1585cc93c56473a087ed129e23d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 21 Aug 2026 17:29:26 -0700 Subject: [PATCH 07/93] fix: scan sql held in a variable, and bound inserts by their own row source --- .../check_migrations_no_data_rewrites.py | 30 ++++++---- .../test_check_migrations_no_data_rewrites.py | 58 +++++++++++++++++++ 2 files changed, 77 insertions(+), 11 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 8ebdc43fac2..1e67617cead 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -11,9 +11,10 @@ Flagged, per statement, by its leading keyword: UPDATE rewrites every matching row, and `WHERE` does not bound the scan DELETE same scan, and the dead tuples outlive the migration MERGE both of the above in one statement - INSERT only when it draws rows from a `SELECT`; `INSERT ... VALUES` is bounded - by the literal row list and passes, scalar subqueries in that list - included + INSERT only when it draws rows from a `SELECT`; an insert whose row source is + a leading `VALUES` is bounded by the rows spelled out there and passes, + scalar subqueries in that list included, while a `VALUES` reached + through a subquery or a set operation bounds nothing WITH a CTE-led statement containing any of the above Referential actions (`ON DELETE CASCADE`, `ON UPDATE CASCADE`) are schema, never a @@ -22,7 +23,9 @@ statement's leading keyword, so they pass. Statements inside dollar-quoted bodies are scanned too. `DO $$ ... $$` is this repo's idiom for conditional DDL, so a body is where an `UPDATE` would otherwise hide. The SQL an `EXECUTE` runs is scanned the same way, since a rewrite reads the -same to Postgres whether it is spelled out or handed over as a string. +same to Postgres whether it is spelled out or handed over as a string, and so is a +literal assigned to a variable with `:=`, which is where an `EXECUTE` further down +the body gets its statement from. Line numbers always count against the whole migration file, however deeply the statement is nested, so a reported line points at the statement and the markers @@ -248,11 +251,17 @@ def offending_keyword(statement: str) -> str | None: def draws_rows_from_a_select(statement: str) -> bool: - """Whether an `INSERT` takes its rows from a query rather than a literal list. A - top-level `VALUES` bounds the insert to the rows written out there, so the scalar - subqueries and helper CTEs that sit in parentheses around it do not make it a - rewrite.""" - return contains(statement, "SELECT") and not contains(strip_parens(statement), "VALUES") + """Whether an `INSERT` takes its rows from a query rather than a literal list. Only a + `SELECT` the insert is built on counts, so the scalar subqueries and helper CTEs that + sit in parentheses around a `VALUES` list do not make it a rewrite, while one reached + through a set operation does.""" + return contains(strip_parens(statement), "SELECT") + + +def hands_off_sql(statement: str) -> bool: + """Whether a statement gives the server a string literal to run as SQL. `EXECUTE` runs + one outright, and an assignment parks one in a variable for an `EXECUTE` further down.""" + return leads_with(statement, "EXECUTE") or ":=" in statement def leads_with(statement: str, keyword: str) -> bool: @@ -281,11 +290,10 @@ def scan_region( masked, bodies, literals = mask(region) for match in STATEMENT.finditer(masked): - if leads_with(match.group(), "EXECUTE"): + if hands_off_sql(match.group()): for start, end in literals: if match.start() <= start and end <= match.end(): yield from scan_region(document, region[start:end], migration, exempt, offset + start) - continue keyword = offending_keyword(match.group()) if keyword is None: continue diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 9c6871cf557..2aa61b50949 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -113,6 +113,18 @@ class TestInsert: sql = 'INSERT INTO "Foo" ("id") SELECT "id" FROM (VALUES (1), (2)) AS "v"("id");' assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + def test_values_after_a_set_operation_does_not_bound_an_insert_select(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") SELECT "id" FROM "Bar" UNION ALL VALUES (1);' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_values_after_an_except_does_not_bound_an_insert_select(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") SELECT "id" FROM "Bar" EXCEPT VALUES (1);' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_select_term_after_a_values_list_is_still_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") VALUES (1), (2) UNION ALL SELECT "id" FROM "Bar";' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + class TestCommonTableExpressions: def test_cte_led_update_is_flagged(self, tmp_path): @@ -272,6 +284,11 @@ class TestEscapeHatch: sql = 'UPDATE "Foo" SET "a" = 1;\n-- data-migration-ok: bounded\nALTER TABLE "Bar" ADD COLUMN "b" TEXT;' assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_a_marker_written_below_its_statement_leaves_that_statement_flagged(self, tmp_path): + sql = 'UPDATE "Foo" SET "a" = 1;\n-- data-migration-ok: bounded\nUPDATE "Bar" SET "b" = 2;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 1 + def test_marker_inside_a_do_block_below_the_first_line_exempts(self, tmp_path): sql = ( "-- AlterTable\n" @@ -334,6 +351,47 @@ class TestDynamicSql: sql = "INSERT INTO \"Foo\" (\"note\") VALUES ('UPDATE \"Bar\" SET \"a\" = 1');" assert _keywords(tmp_path, sql) == () + def test_a_rewrite_declared_into_a_variable_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text := 'UPDATE \"Foo\" SET \"a\" = 1';\n" + "BEGIN\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 3 + + def test_a_rewrite_assigned_in_the_body_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + "BEGIN\n" + " stmt := 'DELETE FROM \"Foo\" WHERE \"a\" = 1';\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("DELETE",) + assert _scan(tmp_path, sql)[0].line == 5 + + def test_ddl_assigned_to_a_variable_passes(self, tmp_path): + sql = "DO $$\nDECLARE\n stmt text := 'ALTER TABLE \"Foo\" ADD COLUMN \"b\" TEXT';\nBEGIN\n EXECUTE stmt;\nEND $$;" + assert _keywords(tmp_path, sql) == () + + def test_a_marker_exempts_a_rewrite_held_in_a_variable(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " -- data-migration-ok: one config row, keyed by its primary key\n" + " stmt text := 'UPDATE \"Config\" SET \"v\" = 1 WHERE \"k\" = ''rev''';\n" + "BEGIN\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + class TestReporting: def test_line_number_points_at_the_statement_keyword(self, tmp_path): From c4d9a1ac6c501da875636d72a7cf12dd676fa952 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 21 Aug 2026 17:40:44 -0700 Subject: [PATCH 08/93] fix: keep a marker trailing a statement from exempting the next one --- .../check_migrations_no_data_rewrites.py | 46 ++++++++++++++----- .../test_check_migrations_no_data_rewrites.py | 23 ++++++++++ 2 files changed, 58 insertions(+), 11 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 1e67617cead..558bce5ba6a 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -33,8 +33,10 @@ below line up with the statements they exempt. Add a column and let the application populate it, or run the rewrite as an opt-in batched job outside boot. When a rewrite is genuinely bounded and must ship inside -the migration, put `-- data-migration-ok: ` on the statement, naming what -bounds it. The reason is required. +the migration, put `-- data-migration-ok: ` on the statement or on the line +above it, naming what bounds it. The reason is required. A marker sharing a line +with the statement it follows exempts that statement alone, so the next statement +down is still checked rather than picking the marker up as its own. `GRANDFATHERED` freezes the violations that predate this check. Prisma records a checksum for every applied migration and this repo treats applied files as @@ -117,6 +119,19 @@ class Violation: return f"{location}:{self.line}: {self.keyword} rewrites existing rows at boot" +@dataclass(frozen=True, slots=True) +class Markers: + lines: frozenset[int] + standalone: frozenset[int] + + def exempt(self, first: int, last: int) -> bool: + """Whether a statement spanning `first` to `last` carries a marker. A marker alone on + its line speaks for the statement below it, which is how one written above a rewrite + exempts it. A marker sharing its line with the statement it follows speaks for that + statement only, so the next statement down does not inherit the exemption.""" + return any(line in self.lines for line in range(first, last + 1)) or first - 1 in self.standalone + + def blank(text: str) -> str: return "".join(character if character == "\n" else " " for character in text) @@ -273,16 +288,25 @@ def contains(statement: str, keyword: str) -> bool: return re.search(rf"\b{keyword}\b", statement, re.IGNORECASE) is not None -def exempt_lines(sql: str) -> frozenset[int]: - return frozenset(sql.count("\n", 0, match.start()) + 1 for match in MARKER.finditer(sql)) +def read_markers(sql: str) -> Markers: + lines: set[int] = set() + standalone: set[int] = set() + + for match in MARKER.finditer(sql): + line = sql.count("\n", 0, match.start()) + 1 + lines.add(line) + if not sql[sql.rfind("\n", 0, match.start()) + 1 : match.start()].strip(): + standalone.add(line) + + return Markers(frozenset(lines), frozenset(standalone)) -def scan(sql: str, migration: str, exempt: frozenset[int]) -> Iterator[Violation]: - yield from scan_region(sql, sql, migration, exempt, 0) +def scan(sql: str, migration: str, markers: Markers) -> Iterator[Violation]: + yield from scan_region(sql, sql, migration, markers, 0) def scan_region( - document: str, region: str, migration: str, exempt: frozenset[int], offset: int + document: str, region: str, migration: str, markers: Markers, offset: int ) -> Iterator[Violation]: """Violations in one region of `document`, whose text begins at `offset`. Lines are always counted against the whole document, so a statement nested in a dollar-quoted @@ -293,18 +317,18 @@ def scan_region( if hands_off_sql(match.group()): for start, end in literals: if match.start() <= start and end <= match.end(): - yield from scan_region(document, region[start:end], migration, exempt, offset + start) + yield from scan_region(document, region[start:end], migration, markers, offset + start) keyword = offending_keyword(match.group()) if keyword is None: continue first = line_of(document, offset + keyword_start(match)) last = line_of(document, offset + match.end()) - if any(line in exempt for line in range(first - 1, last + 1)): + if markers.exempt(first, last): continue yield Violation(migration, first, keyword) for start, end in bodies: - yield from scan_region(document, region[start:end], migration, exempt, offset + start) + yield from scan_region(document, region[start:end], migration, markers, offset + start) def keyword_start(statement: re.Match[str]) -> int: @@ -318,7 +342,7 @@ def line_of(sql: str, offset: int) -> int: def scan_migration(directory: Path) -> tuple[Violation, ...]: sql = (directory / "migration.sql").read_text(encoding="utf-8") - return tuple(scan(sql, directory.name, exempt_lines(sql))) + return tuple(scan(sql, directory.name, read_markers(sql))) def stale_grandfathers(found: Mapping[str, tuple[Violation, ...]]) -> tuple[str, ...]: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 2aa61b50949..6befc895b4b 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -302,6 +302,29 @@ class TestEscapeHatch: ) assert _keywords(tmp_path, sql) == () + def test_a_marker_trailing_a_statement_exempts_that_statement(self, tmp_path): + sql = 'UPDATE "Foo" SET "a" = 1; -- data-migration-ok: one row\nALTER TABLE "Bar" ADD COLUMN "b" TEXT;' + assert _keywords(tmp_path, sql) == () + + def test_a_marker_trailing_a_statement_does_not_exempt_the_next_one(self, tmp_path): + sql = 'UPDATE "Foo" SET "a" = 1; -- data-migration-ok: one row\nUPDATE "Bar" SET "b" = 2;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 2 + + def test_a_marker_trailing_a_multiline_statement_does_not_exempt_the_next_one(self, tmp_path): + sql = ( + 'UPDATE "Foo"\n' + ' SET "a" = 1; -- data-migration-ok: one row\n' + 'UPDATE "Bar" SET "b" = 2;' + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 3 + + def test_a_marker_alone_between_two_statements_belongs_to_the_one_below_it(self, tmp_path): + sql = 'UPDATE "Foo" SET "a" = 1;\n-- data-migration-ok: one row\nUPDATE "Bar" SET "b" = 2;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 1 + def test_marker_above_a_do_block_does_not_exempt_a_rewrite_inside_it(self, tmp_path): sql = ( "-- data-migration-ok: bounded, this belongs to the insert below\n" From 622d9c598d0dccc17cc97fefbe69f385fdba0ced Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 21 Aug 2026 17:56:02 -0700 Subject: [PATCH 09/93] fix: read dynamic SQL through the statement that hands it off A marker on an EXECUTE now covers the SQL that EXECUTE runs, so it goes where the migration reads rather than inside the string. A literal whose first line sat below its EXECUTE was missing the marker entirely, and the documented placement failed CI. A literal assigned with := counts as SQL only when an EXECUTE in the same body runs that variable by name. An error message naming a DELETE the application handles is text, and the only way to silence it before was a marker claiming a bounded data migration that was not there at all. --- .../check_migrations_no_data_rewrites.py | 49 +++++++++---- .../test_check_migrations_no_data_rewrites.py | 71 +++++++++++++++++++ 2 files changed, 107 insertions(+), 13 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 558bce5ba6a..2d5bb6ec726 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -24,8 +24,9 @@ Statements inside dollar-quoted bodies are scanned too. `DO $$ ... $$` is this repo's idiom for conditional DDL, so a body is where an `UPDATE` would otherwise hide. The SQL an `EXECUTE` runs is scanned the same way, since a rewrite reads the same to Postgres whether it is spelled out or handed over as a string, and so is a -literal assigned to a variable with `:=`, which is where an `EXECUTE` further down -the body gets its statement from. +literal assigned with `:=` to a variable some `EXECUTE` in the same body then runs +by name. A literal nothing runs is text, however much it reads like a statement, +so an error message naming a `DELETE` the application handles stays a message. Line numbers always count against the whole migration file, however deeply the statement is nested, so a reported line points at the statement and the markers @@ -36,7 +37,9 @@ batched job outside boot. When a rewrite is genuinely bounded and must ship insi the migration, put `-- data-migration-ok: ` on the statement or on the line above it, naming what bounds it. The reason is required. A marker sharing a line with the statement it follows exempts that statement alone, so the next statement -down is still checked rather than picking the marker up as its own. +down is still checked rather than picking the marker up as its own. A marker on an +`EXECUTE` or on the assignment feeding one covers the SQL that statement hands off, +so it goes where the migration reads rather than inside the string. `GRANDFATHERED` freezes the violations that predate this check. Prisma records a checksum for every applied migration and this repo treats applied files as @@ -66,6 +69,7 @@ MARKER = re.compile(r"--[ \t]*data-migration-ok:[ \t]*(\S.*?)[ \t]*$", re.MULTIL DOLLAR_TAG = re.compile(r"\$(?:[A-Za-z_][A-Za-z0-9_]*)?\$") FIRST_WORD = re.compile(r"[A-Za-z_][A-Za-z0-9_]*") STATEMENT = re.compile(r"[^;]+") +RUN_BY_NAME = re.compile(r"\bEXECUTE[ \t]+([A-Za-z_][A-Za-z0-9_]*)", re.IGNORECASE) REWRITES_ROWS = frozenset({"UPDATE", "DELETE", "MERGE"}) @@ -273,10 +277,27 @@ def draws_rows_from_a_select(statement: str) -> bool: return contains(strip_parens(statement), "SELECT") -def hands_off_sql(statement: str) -> bool: - """Whether a statement gives the server a string literal to run as SQL. `EXECUTE` runs - one outright, and an assignment parks one in a variable for an `EXECUTE` further down.""" - return leads_with(statement, "EXECUTE") or ":=" in statement +def hands_off_sql(statement: str, executed: frozenset[str]) -> bool: + """Whether a statement gives the server a string literal to run as SQL. `EXECUTE` runs one + outright. An assignment parks one in a variable, which counts only when something further + down runs that variable by name, since a string the migration never executes is text.""" + return leads_with(statement, "EXECUTE") or bool(assigned_names(statement) & executed) + + +def assigned_names(statement: str) -> frozenset[str]: + """The candidate variable names an assignment writes to, taken as every word ahead of the + `:=`. A declaration carries its type and sometimes a leading `DECLARE` alongside the name, + and none of that is worth parsing when the only question is which name is executed.""" + head, separator, _ = statement.partition(":=") + if not separator: + return frozenset() + return frozenset(word.group().lower() for word in FIRST_WORD.finditer(head)) + + +def executed_names(masked: str) -> frozenset[str]: + """The variables handed to an `EXECUTE` by name. Reading these off the masked text keeps + an `EXECUTE` written inside a comment or a string from counting.""" + return frozenset(match.group(1).lower() for match in RUN_BY_NAME.finditer(masked)) def leads_with(statement: str, keyword: str) -> bool: @@ -312,18 +333,20 @@ def scan_region( always counted against the whole document, so a statement nested in a dollar-quoted body reports its real file line and lines up with the markers read from that file.""" masked, bodies, literals = mask(region) + executed = executed_names(masked) for match in STATEMENT.finditer(masked): - if hands_off_sql(match.group()): + first = line_of(document, offset + keyword_start(match)) + last = line_of(document, offset + match.end()) + exempt = markers.exempt(first, last) + + if hands_off_sql(match.group(), executed) and not exempt: for start, end in literals: if match.start() <= start and end <= match.end(): yield from scan_region(document, region[start:end], migration, markers, offset + start) + keyword = offending_keyword(match.group()) - if keyword is None: - continue - first = line_of(document, offset + keyword_start(match)) - last = line_of(document, offset + match.end()) - if markers.exempt(first, last): + if keyword is None or exempt: continue yield Violation(migration, first, keyword) diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 6befc895b4b..926202f1d0e 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -415,6 +415,77 @@ class TestDynamicSql: ) assert _keywords(tmp_path, sql) == () + def test_a_marker_exempts_an_execute_whose_sql_starts_on_a_later_line(self, tmp_path): + sql = ( + "DO $$\n" + "BEGIN\n" + " -- data-migration-ok: one config row, keyed by its primary key\n" + " EXECUTE '\n" + " UPDATE \"Config\" SET \"v\" = 1 WHERE \"k\" = ''rev''';\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_marker_exempts_an_assignment_whose_sql_starts_on_a_later_line(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " -- data-migration-ok: one config row, keyed by its primary key\n" + " stmt text := '\n" + " UPDATE \"Config\" SET \"v\" = 1 WHERE \"k\" = ''rev''';\n" + "BEGIN\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_an_unmarked_execute_whose_sql_starts_on_a_later_line_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "BEGIN\n" + " EXECUTE '\n" + " UPDATE \"Foo\" SET \"a\" = 1';\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 4 + + def test_a_message_assigned_but_never_executed_passes(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " msg text := 'UPDATE of legacy rows skipped, the application backfills them';\n" + "BEGIN\n" + " RAISE NOTICE '%', msg;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_notice_naming_a_delete_it_never_runs_passes(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " note text := 'DELETE FROM legacy rows is handled by the application';\n" + "BEGIN\n" + " RAISE NOTICE '%', note;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_only_the_variable_that_is_executed_is_read_as_sql(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " msg text := 'UPDATE of legacy rows skipped';\n" + " stmt text := 'DELETE FROM \"Foo\" WHERE \"a\" = 1';\n" + "BEGIN\n" + " RAISE NOTICE '%', msg;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("DELETE",) + assert _scan(tmp_path, sql)[0].line == 4 + class TestReporting: def test_line_number_points_at_the_statement_keyword(self, tmp_path): From 3cca3f540286fd45514856f43eac95432c1dbd5a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 21 Aug 2026 18:04:16 -0700 Subject: [PATCH 10/93] fix: scan a DO body written in single quotes DO takes its body as a string literal, and dollar quoting is a convenience rather than a requirement. A migration spelling the body in single quotes got its rewrite through untouched, since nothing was reading that literal as SQL. It is ordinary syntax rather than an attempt to hide anything, so the miss was reachable by accident. The module docstring now also records where concatenated dynamic SQL stops being readable, which is a keyword split across fragments that do not hold it. Every fragment is scanned, so the shapes people actually write are all still caught. --- .../check_migrations_no_data_rewrites.py | 23 +++++++++++---- .../test_check_migrations_no_data_rewrites.py | 29 +++++++++++++++++++ 2 files changed, 47 insertions(+), 5 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 2d5bb6ec726..bf405a8aa31 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -25,8 +25,17 @@ repo's idiom for conditional DDL, so a body is where an `UPDATE` would otherwise hide. The SQL an `EXECUTE` runs is scanned the same way, since a rewrite reads the same to Postgres whether it is spelled out or handed over as a string, and so is a literal assigned with `:=` to a variable some `EXECUTE` in the same body then runs -by name. A literal nothing runs is text, however much it reads like a statement, -so an error message naming a `DELETE` the application handles stays a message. +by name, and so is the body of a `DO` written in single quotes rather than dollar +quotes. A literal nothing runs is text, however much it reads like a statement, so +an error message naming a `DELETE` the application handles stays a message. + +Each literal is read on its own, so a keyword built by concatenating fragments that +do not contain it (`'UPD' || 'ATE ...'`) is not caught. Every fragment is scanned, +so a concatenation is caught wherever the keyword survives whole in one of them, +which covers `'UPDATE ' || quote_ident(t)` and the rest of the readable shapes. The +gap needs a keyword deliberately split down the middle, and this check is a guard +against a rewrite reaching a boot unnoticed, not a defence against someone hiding +one on purpose. Line numbers always count against the whole migration file, however deeply the statement is nested, so a reported line points at the statement and the markers @@ -91,6 +100,7 @@ STATEMENT_KEYWORDS = REWRITES_ROWS | frozenset( "RAISE", "RETURN", "EXECUTE", + "DO", "CALL", "REINDEX", "REFRESH", @@ -279,9 +289,12 @@ def draws_rows_from_a_select(statement: str) -> bool: def hands_off_sql(statement: str, executed: frozenset[str]) -> bool: """Whether a statement gives the server a string literal to run as SQL. `EXECUTE` runs one - outright. An assignment parks one in a variable, which counts only when something further - down runs that variable by name, since a string the migration never executes is text.""" - return leads_with(statement, "EXECUTE") or bool(assigned_names(statement) & executed) + outright, and so does `DO`, whose body is a string wherever it is not dollar-quoted. An + assignment parks one in a variable, which counts only when something further down runs + that variable by name, since a string the migration never executes is text.""" + if leads_with(statement, "EXECUTE") or leads_with(statement, "DO"): + return True + return bool(assigned_names(statement) & executed) def assigned_names(statement: str) -> frozenset[str]: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 926202f1d0e..c394eb30c37 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -472,6 +472,35 @@ class TestDynamicSql: ) assert _keywords(tmp_path, sql) == () + def test_a_do_body_in_single_quotes_is_scanned(self, tmp_path): + sql = "DO 'BEGIN UPDATE \"Foo\" SET \"a\" = 1; END';" + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 1 + + def test_a_quoted_do_body_with_a_language_clause_is_scanned(self, tmp_path): + sql = "DO LANGUAGE plpgsql 'BEGIN DELETE FROM \"Foo\"; END';" + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_a_quoted_do_body_holding_only_ddl_passes(self, tmp_path): + sql = "DO 'BEGIN ALTER TABLE \"Foo\" ADD COLUMN \"b\" TEXT; END';" + assert _keywords(tmp_path, sql) == () + + def test_a_marker_exempts_a_quoted_do_body(self, tmp_path): + sql = ( + "-- data-migration-ok: one config row, keyed by its primary key\n" + "DO 'BEGIN UPDATE \"Config\" SET \"v\" = 1 WHERE \"k\" = ''rev''; END';" + ) + assert _keywords(tmp_path, sql) == () + + def test_concatenated_sql_is_flagged_when_the_keyword_leads_a_fragment(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'UPDATE ' || quote_ident('Foo') || ' SET \"a\" = 1';\nEND $$;" + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 3 + + def test_concatenated_sql_is_flagged_when_the_keyword_leads_a_later_fragment(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'WITH x AS (SELECT 1) ' || 'UPDATE \"Foo\" SET \"a\" = 1';\nEND $$;" + assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_only_the_variable_that_is_executed_is_read_as_sql(self, tmp_path): sql = ( "DO $$\n" From 3d16e327c6023a708f0ac8791ebe7400d0e46456 Mon Sep 17 00:00:00 2001 From: Mateo Date: Fri, 21 Aug 2026 18:17:03 -0700 Subject: [PATCH 11/93] fix: judge an EXPLAIN-wrapped statement on the statement itself EXPLAIN ANALYZE runs the statement it wraps rather than only planning it, but ANALYZE sits in the keyword set, so it stood in for the keyword underneath and a rewrite left under one reached boot unflagged. --- .../check_migrations_no_data_rewrites.py | 34 ++++++++++++++---- .../test_check_migrations_no_data_rewrites.py | 35 +++++++++++++++++++ 2 files changed, 63 insertions(+), 6 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index bf405a8aa31..572c20805a0 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -20,6 +20,12 @@ Flagged, per statement, by its leading keyword: Referential actions (`ON DELETE CASCADE`, `ON UPDATE CASCADE`) are schema, never a statement's leading keyword, so they pass. +A statement wrapped in `EXPLAIN` is judged on the statement itself, because the +`ANALYZE` form runs it rather than only planning it, and a rewrite left under one +rewrites the table on the way to printing its timings. Explaining a rewrite without +`ANALYZE` is flagged too: nothing here needs the plan of a statement it is being +told not to run at boot, and a marker is a cheap answer if one ever does. + Statements inside dollar-quoted bodies are scanned too. `DO $$ ... $$` is this repo's idiom for conditional DDL, so a body is where an `UPDATE` would otherwise hide. The SQL an `EXECUTE` runs is scanned the same way, since a rewrite reads the @@ -79,6 +85,8 @@ DOLLAR_TAG = re.compile(r"\$(?:[A-Za-z_][A-Za-z0-9_]*)?\$") FIRST_WORD = re.compile(r"[A-Za-z_][A-Za-z0-9_]*") STATEMENT = re.compile(r"[^;]+") RUN_BY_NAME = re.compile(r"\bEXECUTE[ \t]+([A-Za-z_][A-Za-z0-9_]*)", re.IGNORECASE) +EXPLAIN_OPTIONS = re.compile(r"\bEXPLAIN\b(?:\s+(?:ANALYZE|ANALYSE|VERBOSE)\b)+", re.IGNORECASE) +NOT_A_NEWLINE = re.compile(r"[^\n]") REWRITES_ROWS = frozenset({"UPDATE", "DELETE", "MERGE"}) @@ -247,17 +255,31 @@ def strip_parens(statement: str) -> str: return "".join(chunks) +def strip_explain(statement: str) -> str: + """Blank an `EXPLAIN` written with bare options, since the `ANALYZE` among them would + otherwise stand in for the keyword of the statement being explained. That statement is + the one worth reading: `EXPLAIN ANALYZE` runs it rather than only planning it, so a + rewrite underneath rewrites the table for real. The parenthesised option list needs + nothing here, already being blanked as a group.""" + return EXPLAIN_OPTIONS.sub(lambda match: NOT_A_NEWLINE.sub(" ", match.group()), statement) + + def leading_keyword(statement: str) -> re.Match[str] | None: - """The statement's own keyword, looking past PL/pgSQL block syntax such as - `BEGIN`, `IF ... THEN` and `END`.""" + """The statement's own keyword, looking past what wraps it: a parenthesised guard, + PL/pgSQL block syntax such as `BEGIN`, `IF ... THEN` and `END`, and an `EXPLAIN`. + Offsets survive both strips, so the match still points into `statement` itself.""" return next( - (word for word in FIRST_WORD.finditer(statement) if word.group().upper() in STATEMENT_KEYWORDS), + ( + word + for word in FIRST_WORD.finditer(strip_explain(strip_parens(statement))) + if word.group().upper() in STATEMENT_KEYWORDS + ), None, ) def offending_keyword(statement: str) -> str | None: - word = leading_keyword(strip_parens(statement)) + word = leading_keyword(statement) if word is None: return None @@ -314,7 +336,7 @@ def executed_names(masked: str) -> frozenset[str]: def leads_with(statement: str, keyword: str) -> bool: - word = leading_keyword(strip_parens(statement)) + word = leading_keyword(statement) return word is not None and word.group().upper() == keyword @@ -368,7 +390,7 @@ def scan_region( def keyword_start(statement: re.Match[str]) -> int: - word = leading_keyword(strip_parens(statement.group())) + word = leading_keyword(statement.group()) return statement.start() + (0 if word is None else word.start()) diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index c394eb30c37..fb56d032617 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -516,6 +516,41 @@ class TestDynamicSql: assert _scan(tmp_path, sql)[0].line == 4 +class TestExplain: + def test_explain_analyze_over_an_update_is_flagged(self, tmp_path): + sql = 'EXPLAIN ANALYZE UPDATE "Foo" SET "a" = 1;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 1 + + def test_explain_analyze_verbose_over_a_delete_is_flagged(self, tmp_path): + assert _keywords(tmp_path, 'EXPLAIN ANALYZE VERBOSE DELETE FROM "Foo";') == ("DELETE",) + + def test_explain_with_a_parenthesised_analyze_is_flagged(self, tmp_path): + assert _keywords(tmp_path, 'EXPLAIN (ANALYZE, BUFFERS) UPDATE "Foo" SET "a" = 1;') == ("UPDATE",) + + def test_explain_analyze_over_an_insert_select_is_flagged(self, tmp_path): + sql = 'EXPLAIN ANALYZE INSERT INTO "Foo" SELECT "a" FROM "Bar";' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_explain_analyze_over_a_select_passes(self, tmp_path): + assert _keywords(tmp_path, 'EXPLAIN ANALYZE SELECT * FROM "Foo";') == () + + def test_a_marker_exempts_an_explained_rewrite(self, tmp_path): + sql = '-- data-migration-ok: one config row\nEXPLAIN ANALYZE UPDATE "Foo" SET "a" = 1;' + assert _keywords(tmp_path, sql) == () + + def test_an_analyze_of_its_own_passes(self, tmp_path): + assert _keywords(tmp_path, 'ANALYZE "Foo";') == () + + def test_a_vacuum_analyze_passes(self, tmp_path): + assert _keywords(tmp_path, 'VACUUM ANALYZE "Foo";') == () + + def test_an_explained_rewrite_inside_a_block_reports_its_line(self, tmp_path): + sql = 'DO $$\nBEGIN\n EXPLAIN ANALYZE UPDATE "Foo" SET "a" = 1;\nEND $$;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 3 + + class TestReporting: def test_line_number_points_at_the_statement_keyword(self, tmp_path): sql = '-- CreateIndex\nCREATE INDEX "i" ON "Foo"("a");\n\nUPDATE "Foo" SET "a" = 1;' From aabaa5151b0f6de6201f50f044ef99ae057251b5 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 21 Aug 2026 18:43:04 -0700 Subject: [PATCH 12/93] docs: say where a marker goes for a dollar-quoted dynamic payload A dollar-quoted payload is read as its own region rather than as a handed-off string, so the marker belongs on the rewrite inside it. Pin that placement, and pin that a marker on a DO block header never covers the block's body. --- .../check_migrations_no_data_rewrites.py | 11 +-- .../test_check_migrations_no_data_rewrites.py | 71 +++++++++++++++++++ 2 files changed, 78 insertions(+), 4 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 572c20805a0..cc541a70102 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -53,8 +53,12 @@ the migration, put `-- data-migration-ok: ` on the statement or on the l above it, naming what bounds it. The reason is required. A marker sharing a line with the statement it follows exempts that statement alone, so the next statement down is still checked rather than picking the marker up as its own. A marker on an -`EXECUTE` or on the assignment feeding one covers the SQL that statement hands off, -so it goes where the migration reads rather than inside the string. +`EXECUTE` or on the assignment feeding one covers the single-quoted SQL that +statement hands off, so it goes where the migration reads rather than inside the +string. A dollar-quoted payload is not a string to this check but a region read like +any other body, so a rewrite inside one takes its marker on the rewrite itself. That +placement is deliberate rather than an oversight: a marker covering a whole body +would let one written for a `DO` block silence a rewrite added to that block later. `GRANDFATHERED` freezes the violations that predate this check. Prisma records a checksum for every applied migration and this repo treats applied files as @@ -86,7 +90,6 @@ FIRST_WORD = re.compile(r"[A-Za-z_][A-Za-z0-9_]*") STATEMENT = re.compile(r"[^;]+") RUN_BY_NAME = re.compile(r"\bEXECUTE[ \t]+([A-Za-z_][A-Za-z0-9_]*)", re.IGNORECASE) EXPLAIN_OPTIONS = re.compile(r"\bEXPLAIN\b(?:\s+(?:ANALYZE|ANALYSE|VERBOSE)\b)+", re.IGNORECASE) -NOT_A_NEWLINE = re.compile(r"[^\n]") REWRITES_ROWS = frozenset({"UPDATE", "DELETE", "MERGE"}) @@ -261,7 +264,7 @@ def strip_explain(statement: str) -> str: the one worth reading: `EXPLAIN ANALYZE` runs it rather than only planning it, so a rewrite underneath rewrites the table for real. The parenthesised option list needs nothing here, already being blanked as a group.""" - return EXPLAIN_OPTIONS.sub(lambda match: NOT_A_NEWLINE.sub(" ", match.group()), statement) + return EXPLAIN_OPTIONS.sub(lambda match: blank(match.group()), statement) def leading_keyword(statement: str) -> re.Match[str] | None: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index fb56d032617..86216f59174 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -335,6 +335,42 @@ class TestEscapeHatch: assert _keywords(tmp_path, sql) == ("UPDATE",) assert _scan(tmp_path, sql)[0].line == 4 + def test_marker_directly_above_a_do_block_does_not_exempt_its_body(self, tmp_path): + sql = ( + "-- data-migration-ok: seeding two default rows\n" + "DO $$\n" + "BEGIN\n" + ' INSERT INTO "Foo" ("a") VALUES (1);\n' + ' UPDATE "Foo" SET "a" = 1;\n' + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 5 + + def test_marker_on_the_do_line_does_not_exempt_its_body(self, tmp_path): + sql = ( + "DO $$ -- data-migration-ok: bounded to one row\n" + "BEGIN\n" + ' IF EXISTS (SELECT 1 FROM "Foo") THEN\n' + ' UPDATE "Foo" SET "a" = 1;\n' + " END IF;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 4 + + def test_a_marked_rewrite_does_not_exempt_a_later_one_in_the_same_block(self, tmp_path): + sql = ( + "DO $$\n" + "BEGIN\n" + " -- data-migration-ok: bounded to one row\n" + ' UPDATE "Foo" SET "a" = 1;\n' + ' UPDATE "Bar" SET "b" = 2;\n' + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 5 + class TestDynamicSql: def test_execute_of_a_quoted_update_is_flagged(self, tmp_path): @@ -515,6 +551,41 @@ class TestDynamicSql: assert _keywords(tmp_path, sql) == ("DELETE",) assert _scan(tmp_path, sql)[0].line == 4 + def test_a_marker_on_an_execute_covers_its_single_quoted_payload(self, tmp_path): + sql = ( + "DO $$\n" + "BEGIN\n" + " EXECUTE ' -- data-migration-ok: bounded to one row\n" + ' UPDATE "Foo" SET "a" = 1;\n' + " ';\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_marker_on_an_execute_does_not_reach_into_a_dollar_quoted_payload(self, tmp_path): + sql = ( + "DO $$\n" + "BEGIN\n" + " EXECUTE $x$ -- data-migration-ok: bounded to one row\n" + ' UPDATE "Foo" SET "a" = 1;\n' + " $x$;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 4 + + def test_a_marker_inside_a_dollar_quoted_payload_exempts_its_rewrite(self, tmp_path): + sql = ( + "DO $$\n" + "BEGIN\n" + " EXECUTE $x$\n" + " -- data-migration-ok: bounded to one row\n" + ' UPDATE "Foo" SET "a" = 1;\n' + " $x$;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + class TestExplain: def test_explain_analyze_over_an_update_is_flagged(self, tmp_path): From 64267ebd28ec35fa297a164b3c2405d6b5abc2bd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 21 Aug 2026 18:58:34 -0700 Subject: [PATCH 13/93] fix: flag an INSERT whose rows come from a parenthesised query Postgres takes the row source parenthesised, so `INSERT INTO "t" ("a") (SELECT ...)` copies a whole table at boot. Reading only the unparenthesised text let it through: 777eb8af10 caught it, then e7dea842c3 traded it away to stop a VALUES list joined to a query by a set operation from bounding nothing. Read the top level first so set operations still count, then fall back to the whole statement when no top-level VALUES bounds the insert. `TABLE t` is a row source as much as a `SELECT` is, and it was passing too --- .../check_migrations_no_data_rewrites.py | 42 +++++++++++++------ .../test_check_migrations_no_data_rewrites.py | 34 +++++++++++++++ 2 files changed, 63 insertions(+), 13 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index cc541a70102..a8975e2529e 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -11,10 +11,12 @@ Flagged, per statement, by its leading keyword: UPDATE rewrites every matching row, and `WHERE` does not bound the scan DELETE same scan, and the dead tuples outlive the migration MERGE both of the above in one statement - INSERT only when it draws rows from a `SELECT`; an insert whose row source is - a leading `VALUES` is bounded by the rows spelled out there and passes, - scalar subqueries in that list included, while a `VALUES` reached - through a subquery or a set operation bounds nothing + INSERT only when its rows come from a query rather than a literal `VALUES` + list. The query counts wherever it sits, since Postgres takes it + parenthesised, and `TABLE t` is one as much as a `SELECT` is. An + insert bounded by a leading `VALUES` passes, scalar subqueries in that + list included, while a `VALUES` reached through a subquery or joined + to a query by a set operation bounds nothing WITH a CTE-led statement containing any of the above Referential actions (`ON DELETE CASCADE`, `ON UPDATE CASCADE`) are schema, never a @@ -292,24 +294,38 @@ def offending_keyword(statement: str) -> str | None: return keyword if keyword == "INSERT": - return "INSERT ... SELECT" if draws_rows_from_a_select(statement) else None + source = row_source_keyword(statement) + return None if source is None else f"INSERT ... {source}" if keyword == "WITH": nested = next((name for name in sorted(REWRITES_ROWS) if contains(statement, name)), None) if nested is not None: return f"WITH ... {nested}" - if contains(statement, "INSERT") and draws_rows_from_a_select(statement): - return "WITH ... INSERT ... SELECT" + if contains(statement, "INSERT"): + source = row_source_keyword(statement) + if source is not None: + return f"WITH ... INSERT ... {source}" return None -def draws_rows_from_a_select(statement: str) -> bool: - """Whether an `INSERT` takes its rows from a query rather than a literal list. Only a - `SELECT` the insert is built on counts, so the scalar subqueries and helper CTEs that - sit in parentheses around a `VALUES` list do not make it a rewrite, while one reached - through a set operation does.""" - return contains(strip_parens(statement), "SELECT") +def row_source_keyword(statement: str) -> str | None: + """Which keyword supplies an `INSERT` its rows, or `None` when a literal `VALUES` list + does. A query outside every parenthesis is the row source outright, including one a set + operation joins to a `VALUES` list. Failing that, a `VALUES` outside every parenthesis + is itself the row source, so the scalar subqueries and helper CTEs nested within that + list do not make the insert a rewrite. Failing both, the rows come from a parenthesised + query, which Postgres accepts and which reading only the unparenthesised text would let + through: `INSERT INTO "t" ("a") (SELECT ...)` copies a whole table.""" + outer = strip_parens(statement) + joined = row_source_in(outer) + if joined is not None: + return joined + return None if contains(outer, "VALUES") else row_source_in(statement) + + +def row_source_in(text: str) -> str | None: + return next((word for word in ("SELECT", "TABLE") if contains(text, word)), None) def hands_off_sql(statement: str, executed: frozenset[str]) -> bool: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 86216f59174..15603218a27 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -125,6 +125,32 @@ class TestInsert: sql = 'INSERT INTO "Foo" ("id") VALUES (1), (2) UNION ALL SELECT "id" FROM "Bar";' assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + def test_a_parenthesised_select_row_source_is_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (SELECT "id" FROM "Bar");' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_parenthesised_select_row_source_without_a_column_list_is_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" (SELECT "id" FROM "Bar");' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_parenthesised_select_row_source_spanning_lines_is_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id")\n(\n SELECT "id" FROM "Bar"\n);' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_parenthesised_select_over_a_values_list_is_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (SELECT * FROM (VALUES (1), (2)) AS "v"("id"));' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_set_operation_over_parenthesised_selects_is_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (SELECT 1) UNION (SELECT 2);' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_table_row_source_is_flagged(self, tmp_path): + assert _keywords(tmp_path, 'INSERT INTO "Foo" TABLE "Bar";') == ("INSERT ... TABLE",) + + def test_a_table_named_in_the_insert_target_does_not_flag_it(self, tmp_path): + assert _keywords(tmp_path, 'INSERT INTO "audit table" ("id") VALUES (1);') == () + class TestCommonTableExpressions: def test_cte_led_update_is_flagged(self, tmp_path): @@ -139,6 +165,14 @@ class TestCommonTableExpressions: sql = 'WITH batch AS (SELECT "id" FROM "Bar") INSERT INTO "Foo" ("id") SELECT "id" FROM batch;' assert _keywords(tmp_path, sql) == ("WITH ... INSERT ... SELECT",) + def test_cte_led_insert_from_a_parenthesised_select_is_flagged(self, tmp_path): + sql = 'WITH batch AS (SELECT "id" FROM "Bar") INSERT INTO "Foo" ("id") (SELECT "id" FROM batch);' + assert _keywords(tmp_path, sql) == ("WITH ... INSERT ... SELECT",) + + def test_cte_led_insert_into_a_values_list_passes(self, tmp_path): + sql = 'WITH batch AS (SELECT max("id") FROM "Bar") INSERT INTO "Foo" ("id") VALUES (1);' + assert _keywords(tmp_path, sql) == () + def test_read_only_cte_passes(self, tmp_path): sql = 'WITH batch AS (SELECT "id" FROM "Foo") SELECT count(*) FROM batch;' assert _keywords(tmp_path, sql) == () From 73af0e96921d6e0b3c855c2618facb7f2ea01cb8 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 21 Aug 2026 19:17:27 -0700 Subject: [PATCH 14/93] fix: catch two more row sources the gate let through A parenthesised query term joined to a top-level VALUES list sat behind strip_parens, so an insert reading `VALUES (1) UNION ALL (SELECT ...)` copied a whole table past the gate. A VALUES list now bounds an insert only while no set operation sits beside it at that same level. PL/pgSQL also parks dynamic SQL in a variable through a query's INTO and through the bare `=` it takes as the assignment operator, and assigned_names read neither, so a rewrite handed to a later EXECUTE went unseen. A bare `=` counts only where the words ahead of it make it an assignment rather than a test. --- .../check_migrations_no_data_rewrites.py | 77 ++++++++++--- .../test_check_migrations_no_data_rewrites.py | 102 ++++++++++++++++++ 2 files changed, 162 insertions(+), 17 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index a8975e2529e..912943265bd 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -32,9 +32,10 @@ Statements inside dollar-quoted bodies are scanned too. `DO $$ ... $$` is this repo's idiom for conditional DDL, so a body is where an `UPDATE` would otherwise hide. The SQL an `EXECUTE` runs is scanned the same way, since a rewrite reads the same to Postgres whether it is spelled out or handed over as a string, and so is a -literal assigned with `:=` to a variable some `EXECUTE` in the same body then runs -by name, and so is the body of a `DO` written in single quotes rather than dollar -quotes. A literal nothing runs is text, however much it reads like a statement, so +literal parked in a variable some `EXECUTE` in the same body then runs by name, +however it got there: an assignment with `:=`, the bare `=` PL/pgSQL takes as the +same operator, or a query returning it through `INTO`. So is the body of a `DO` +written in single quotes rather than dollar quotes. A literal nothing runs is text, however much it reads like a statement, so an error message naming a `DELETE` the application handles stays a message. Each literal is read on its own, so a keyword built by concatenating fragments that @@ -91,10 +92,17 @@ DOLLAR_TAG = re.compile(r"\$(?:[A-Za-z_][A-Za-z0-9_]*)?\$") FIRST_WORD = re.compile(r"[A-Za-z_][A-Za-z0-9_]*") STATEMENT = re.compile(r"[^;]+") RUN_BY_NAME = re.compile(r"\bEXECUTE[ \t]+([A-Za-z_][A-Za-z0-9_]*)", re.IGNORECASE) +INTO_TARGETS = re.compile( + r"\bINTO[ \t]+(?:STRICT[ \t]+)?" + r"([A-Za-z_][A-Za-z0-9_]*(?:[ \t]*,[ \t]*[A-Za-z_][A-Za-z0-9_]*)*)", + re.IGNORECASE, +) EXPLAIN_OPTIONS = re.compile(r"\bEXPLAIN\b(?:\s+(?:ANALYZE|ANALYSE|VERBOSE)\b)+", re.IGNORECASE) REWRITES_ROWS = frozenset({"UPDATE", "DELETE", "MERGE"}) +JOINS_QUERIES = ("UNION", "INTERSECT", "EXCEPT") + STATEMENT_KEYWORDS = REWRITES_ROWS | frozenset( { "INSERT", @@ -122,6 +130,10 @@ STATEMENT_KEYWORDS = REWRITES_ROWS | frozenset( } ) +GUARDS_A_CONDITION = frozenset({"IF", "ELSIF", "ELSEIF", "CASE", "WHEN", "WHILE", "EXIT", "ASSERT"}) + +OPENS_A_BLOCK = frozenset({"BEGIN", "THEN", "ELSE", "LOOP"}) + GUIDANCE = """ Migrations apply at proxy boot, before it serves traffic, so a statement whose cost scales with table size is downtime. Add the column and let the application backfill @@ -311,17 +323,21 @@ def offending_keyword(statement: str) -> str | None: def row_source_keyword(statement: str) -> str | None: """Which keyword supplies an `INSERT` its rows, or `None` when a literal `VALUES` list - does. A query outside every parenthesis is the row source outright, including one a set - operation joins to a `VALUES` list. Failing that, a `VALUES` outside every parenthesis - is itself the row source, so the scalar subqueries and helper CTEs nested within that - list do not make the insert a rewrite. Failing both, the rows come from a parenthesised - query, which Postgres accepts and which reading only the unparenthesised text would let - through: `INSERT INTO "t" ("a") (SELECT ...)` copies a whole table.""" + does. A query outside every parenthesis is the row source outright. Failing that, a + `VALUES` outside every parenthesis is itself the row source, so the scalar subqueries + and helper CTEs nested within that list do not make the insert a rewrite, though only + while no set operation sits beside it at that same level: one that does joins the list + to a second query term, and that term is the row source however deeply it is + parenthesised. Failing both, the rows come from a parenthesised query, which Postgres + accepts and which reading only the unparenthesised text would let through: + `INSERT INTO "t" ("a") (SELECT ...)` copies a whole table.""" outer = strip_parens(statement) joined = row_source_in(outer) if joined is not None: return joined - return None if contains(outer, "VALUES") else row_source_in(statement) + if contains(outer, "VALUES") and not any(contains(outer, word) for word in JOINS_QUERIES): + return None + return row_source_in(statement) def row_source_in(text: str) -> str | None: @@ -339,13 +355,40 @@ def hands_off_sql(statement: str, executed: frozenset[str]) -> bool: def assigned_names(statement: str) -> frozenset[str]: - """The candidate variable names an assignment writes to, taken as every word ahead of the - `:=`. A declaration carries its type and sometimes a leading `DECLARE` alongside the name, - and none of that is worth parsing when the only question is which name is executed.""" - head, separator, _ = statement.partition(":=") - if not separator: - return frozenset() - return frozenset(word.group().lower() for word in FIRST_WORD.finditer(head)) + """The candidate variable names a statement writes to, taken as every word ahead of the + assignment operator. A declaration carries its type and sometimes a leading `DECLARE` + alongside the name, and none of that is worth parsing when the only question is which + name is executed. PL/pgSQL spells that operator `:=` and takes a bare `=` as the same + thing, so both count, the second only where `assigns_rather_than_compares` reads it as + an assignment. A query assigns through the target list after its `INTO` instead, which + is how a rewrite reaches a variable with neither operator appearing at all.""" + names: set[str] = set() + + head, operator, _ = statement.partition(":=") + assigns = bool(operator) + if not assigns: + head, operator, _ = statement.partition("=") + assigns = bool(operator) and assigns_rather_than_compares(head) + if assigns: + names.update(word.group().lower() for word in FIRST_WORD.finditer(head)) + + for targets in INTO_TARGETS.finditer(statement): + names.update(word.group().lower() for word in FIRST_WORD.finditer(targets.group(1))) + + return frozenset(names) + + +def assigns_rather_than_compares(head: str) -> bool: + """Whether the bare `=` this text runs up to writes a variable or tests one. Only the + words ahead of it tell the two apart: an assignment is reached with a name and perhaps a + type, while a comparison is reached either through a statement carrying its own keyword + or through a word that guards a condition. Those words stop counting once something + opens a block after them, since a `THEN` ends the condition its `IF` began and the + assignment that follows on the same line is an assignment like any other.""" + words = [word.group().upper() for word in FIRST_WORD.finditer(head)] + opened = max((index for index, word in enumerate(words) if word in OPENS_A_BLOCK), default=-1) + reached = set(words[opened + 1 :]) + return not (reached & STATEMENT_KEYWORDS) and not (reached & GUARDS_A_CONDITION) def executed_names(masked: str) -> frozenset[str]: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 15603218a27..f70d49ddb83 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -145,6 +145,22 @@ class TestInsert: sql = 'INSERT INTO "Foo" ("id") (SELECT 1) UNION (SELECT 2);' assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + def test_a_values_list_joined_to_a_parenthesised_select_is_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") VALUES (1) UNION ALL (SELECT "id" FROM "Bar");' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_values_list_excepting_a_parenthesised_select_is_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") VALUES (1) EXCEPT (SELECT "id" FROM "Bar");' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_values_list_joined_to_a_parenthesised_table_is_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") VALUES (1) UNION ALL (TABLE "Bar");' + assert _keywords(tmp_path, sql) == ("INSERT ... TABLE",) + + def test_a_set_operation_inside_a_values_list_does_not_flag_it(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") VALUES ((SELECT 1 UNION SELECT 2 LIMIT 1));' + assert _keywords(tmp_path, sql) == () + def test_a_table_row_source_is_flagged(self, tmp_path): assert _keywords(tmp_path, 'INSERT INTO "Foo" TABLE "Bar";') == ("INSERT ... TABLE",) @@ -469,6 +485,92 @@ class TestDynamicSql: assert _keywords(tmp_path, sql) == ("DELETE",) assert _scan(tmp_path, sql)[0].line == 5 + def test_a_rewrite_selected_into_a_variable_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + "BEGIN\n" + " SELECT 'UPDATE \"Foo\" SET \"a\" = 1' INTO stmt;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 5 + + def test_a_rewrite_selected_into_a_strict_target_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + "BEGIN\n" + " SELECT 'DELETE FROM \"Foo\"' INTO STRICT stmt;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_a_rewrite_assigned_with_a_bare_equals_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + "BEGIN\n" + " stmt = 'UPDATE \"Foo\" SET \"a\" = 1';\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 5 + + def test_a_rewrite_assigned_with_a_bare_equals_after_then_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + "BEGIN\n" + " IF true THEN stmt = 'DELETE FROM \"Foo\"'; END IF;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("DELETE",) + assert _scan(tmp_path, sql)[0].line == 5 + + def test_a_rewrite_declared_with_a_bare_equals_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text = 'UPDATE \"Foo\" SET \"a\" = 1';\n" + "BEGIN\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_literal_selected_into_a_variable_nothing_runs_is_inert(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " msg text;\n" + "BEGIN\n" + " SELECT 'UPDATE of legacy rows is skipped' INTO msg;\n" + " RAISE NOTICE '%', msg;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_comparing_an_executed_variable_does_not_flag_the_comparison(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text := 'ALTER TABLE \"Foo\" ADD COLUMN \"b\" TEXT';\n" + "BEGIN\n" + " IF stmt = 'DELETE FROM \"Foo\"' THEN RAISE NOTICE 'never'; END IF;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + def test_ddl_assigned_to_a_variable_passes(self, tmp_path): sql = "DO $$\nDECLARE\n stmt text := 'ALTER TABLE \"Foo\" ADD COLUMN \"b\" TEXT';\nBEGIN\n EXECUTE stmt;\nEND $$;" assert _keywords(tmp_path, sql) == () From 1f57e7ea19f7889907956683dd523a88cdece74e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 08:49:02 -0700 Subject: [PATCH 15/93] fix: stop reading an INSERT target table as an assignment target Scanning a statement's own literals for SQL only makes sense when the name before INTO is a variable the body later executes. INSERT INTO names a table there, so an insert into a table sharing a variable's name was flagged for whatever its column values happened to spell. An INSERT that really does assign reaches INTO through RETURNING, which the preceding word separates. Also names the scope boundary in the module docstring: the ban is on row-rewriting DML, not on everything whose cost scales with table size. --- .../check_migrations_no_data_rewrites.py | 20 +++++++++++++++ .../test_check_migrations_no_data_rewrites.py | 25 +++++++++++++++++++ 2 files changed, 45 insertions(+) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 912943265bd..cbc24fe3660 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -6,6 +6,13 @@ anything whose cost scales with existing table size turns into downtime. A singl `UPDATE` with no batching over a spend-log-sized table is minutes of unavailability plus a doubled heap that plain autovacuum will not give back. +What is banned is the row-rewriting DML behind that, not everything whose cost +scales that way. A non-concurrent `CREATE INDEX`, an `ALTER COLUMN ... TYPE` that is +not binary coercible, and a volatile `DEFAULT` on a new column all read the whole +table and all pass. That is deliberate: a rule wide enough to reach them fires on +most ordinary migrations, and a marker everyone adds by reflex stops carrying +information. The outage this was written for was a backfill. + Flagged, per statement, by its leading keyword: UPDATE rewrites every matching row, and `WHERE` does not bound the scan @@ -97,6 +104,7 @@ INTO_TARGETS = re.compile( r"([A-Za-z_][A-Za-z0-9_]*(?:[ \t]*,[ \t]*[A-Za-z_][A-Za-z0-9_]*)*)", re.IGNORECASE, ) +PRECEDING_WORD = re.compile(r"([A-Za-z_][A-Za-z0-9_]*)[^A-Za-z0-9_]*$") EXPLAIN_OPTIONS = re.compile(r"\bEXPLAIN\b(?:\s+(?:ANALYZE|ANALYSE|VERBOSE)\b)+", re.IGNORECASE) REWRITES_ROWS = frozenset({"UPDATE", "DELETE", "MERGE"}) @@ -373,11 +381,23 @@ def assigned_names(statement: str) -> frozenset[str]: names.update(word.group().lower() for word in FIRST_WORD.finditer(head)) for targets in INTO_TARGETS.finditer(statement): + if names_a_table(statement[: targets.start()]): + continue names.update(word.group().lower() for word in FIRST_WORD.finditer(targets.group(1))) return frozenset(names) +def names_a_table(before: str) -> bool: + """Whether the `INTO` this text runs up to introduces a table rather than a query's + target list. `INSERT INTO` is the one that does, and reading its table as somewhere a + string was parked would have an insert scanned for the SQL its own literals spell out. + An `INSERT` that really does assign reaches its `INTO` through a `RETURNING` list, so + the word immediately before is what separates the two.""" + word = PRECEDING_WORD.search(before) + return word is not None and word.group(1).upper() == "INSERT" + + def assigns_rather_than_compares(head: str) -> bool: """Whether the bare `=` this text runs up to writes a variable or tests one. Only the words ahead of it tell the two apart: an assignment is reached with a name and perhaps a diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index f70d49ddb83..7ed8b9978a3 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -547,6 +547,31 @@ class TestDynamicSql: ) assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_an_insert_target_table_is_not_read_as_an_assignment(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " audit text := 'ALTER TABLE \"Foo\" ADD COLUMN \"b\" TEXT';\n" + "BEGIN\n" + " INSERT INTO audit (note) VALUES ('DELETE FROM \"Foo\" is left to the app');\n" + " EXECUTE audit;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_rewrite_returned_into_a_variable_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + "BEGIN\n" + " INSERT INTO \"Log\" (\"sql\") VALUES ('UPDATE \"Foo\" SET \"a\" = 1')\n" + " RETURNING \"sql\" INTO stmt;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_a_literal_selected_into_a_variable_nothing_runs_is_inert(self, tmp_path): sql = ( "DO $$\n" From 813d3a991fabb7d75a314733e9fb88b0d5b94515 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 09:55:18 -0700 Subject: [PATCH 16/93] fix: read every assignment operator, not only the statement's first A bare = was found with partition, so a comparison earlier on the line took the one slot and the assignment after it went unread. INTO targets and the name an EXECUTE runs are also allowed to sit on the next line now. --- .../check_migrations_no_data_rewrites.py | 55 ++++++----- .../test_check_migrations_no_data_rewrites.py | 91 +++++++++++++++++++ 2 files changed, 124 insertions(+), 22 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index cbc24fe3660..9e6df46899b 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -98,12 +98,13 @@ MARKER = re.compile(r"--[ \t]*data-migration-ok:[ \t]*(\S.*?)[ \t]*$", re.MULTIL DOLLAR_TAG = re.compile(r"\$(?:[A-Za-z_][A-Za-z0-9_]*)?\$") FIRST_WORD = re.compile(r"[A-Za-z_][A-Za-z0-9_]*") STATEMENT = re.compile(r"[^;]+") -RUN_BY_NAME = re.compile(r"\bEXECUTE[ \t]+([A-Za-z_][A-Za-z0-9_]*)", re.IGNORECASE) +RUN_BY_NAME = re.compile(r"\bEXECUTE\s+([A-Za-z_][A-Za-z0-9_]*)", re.IGNORECASE) INTO_TARGETS = re.compile( - r"\bINTO[ \t]+(?:STRICT[ \t]+)?" - r"([A-Za-z_][A-Za-z0-9_]*(?:[ \t]*,[ \t]*[A-Za-z_][A-Za-z0-9_]*)*)", + r"\bINTO\s+(?:STRICT\s+)?" + r"([A-Za-z_][A-Za-z0-9_]*(?:\s*,\s*[A-Za-z_][A-Za-z0-9_]*)*)", re.IGNORECASE, ) +ASSIGNS = re.compile(r":=|(?!:=])=(?!=)") PRECEDING_WORD = re.compile(r"([A-Za-z_][A-Za-z0-9_]*)[^A-Za-z0-9_]*$") EXPLAIN_OPTIONS = re.compile(r"\bEXPLAIN\b(?:\s+(?:ANALYZE|ANALYSE|VERBOSE)\b)+", re.IGNORECASE) @@ -368,17 +369,17 @@ def assigned_names(statement: str) -> frozenset[str]: alongside the name, and none of that is worth parsing when the only question is which name is executed. PL/pgSQL spells that operator `:=` and takes a bare `=` as the same thing, so both count, the second only where `assigns_rather_than_compares` reads it as - an assignment. A query assigns through the target list after its `INTO` instead, which - is how a rewrite reaches a variable with neither operator appearing at all.""" + an assignment. Every operator in the statement is read rather than only the first, since + a comparison earlier on the line would otherwise claim the one slot and hide the + assignment after it: `IF n = 1 THEN stmt = '...'` writes `stmt` at its second `=`. A + query assigns through the target list after its `INTO` instead, which is how a rewrite + reaches a variable with neither operator appearing at all.""" names: set[str] = set() - head, operator, _ = statement.partition(":=") - assigns = bool(operator) - if not assigns: - head, operator, _ = statement.partition("=") - assigns = bool(operator) and assigns_rather_than_compares(head) - if assigns: - names.update(word.group().lower() for word in FIRST_WORD.finditer(head)) + for operator in ASSIGNS.finditer(statement): + reached = reached_words(statement[: operator.start()]) + if operator.group() == ":=" or assigns_rather_than_compares(reached): + names.update(word.lower() for word in reached) for targets in INTO_TARGETS.finditer(statement): if names_a_table(statement[: targets.start()]): @@ -398,22 +399,32 @@ def names_a_table(before: str) -> bool: return word is not None and word.group(1).upper() == "INSERT" -def assigns_rather_than_compares(head: str) -> bool: - """Whether the bare `=` this text runs up to writes a variable or tests one. Only the - words ahead of it tell the two apart: an assignment is reached with a name and perhaps a - type, while a comparison is reached either through a statement carrying its own keyword - or through a word that guards a condition. Those words stop counting once something - opens a block after them, since a `THEN` ends the condition its `IF` began and the - assignment that follows on the same line is an assignment like any other.""" +def reached_words(head: str) -> tuple[str, ...]: + """The words an assignment operator is reached through, which is everything since the last + word to open a block. A `THEN` ends the condition its `IF` began, so nothing ahead of it + describes what follows, and neither the name being written nor the keywords that would + mark a comparison ever sit further back than that.""" words = [word.group().upper() for word in FIRST_WORD.finditer(head)] opened = max((index for index, word in enumerate(words) if word in OPENS_A_BLOCK), default=-1) - reached = set(words[opened + 1 :]) - return not (reached & STATEMENT_KEYWORDS) and not (reached & GUARDS_A_CONDITION) + return tuple(words[opened + 1 :]) + + +def assigns_rather_than_compares(reached: tuple[str, ...]) -> bool: + """Whether a bare `=` reached through these words writes a variable or tests one. They are + all that tells the two apart: an assignment is reached with a name and perhaps a type, + while a comparison is reached either through a statement carrying its own keyword or + through a word that guards a condition.""" + words = set(reached) + return not (words & STATEMENT_KEYWORDS) and not (words & GUARDS_A_CONDITION) def executed_names(masked: str) -> frozenset[str]: """The variables handed to an `EXECUTE` by name. Reading these off the masked text keeps - an `EXECUTE` written inside a comment or a string from counting.""" + an `EXECUTE` written inside a comment or a string from counting. Masking blanks a literal + in place rather than removing it, so `EXECUTE '...'` can leave the word after it looking + like the name being run. Reaching that word means crossing no semicolon, which leaves only + the syntax `INTO`, `USING` and `END`, and no assignment ever writes one of those, so the + stray name has nothing to match.""" return frozenset(match.group(1).lower() for match in RUN_BY_NAME.finditer(masked)) diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 7ed8b9978a3..fe53e1ac823 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -572,6 +572,97 @@ class TestDynamicSql: ) assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_a_rewrite_assigned_past_an_earlier_comparison_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + " total int := 1;\n" + "BEGIN\n" + " IF total = 1 THEN stmt = 'DELETE FROM \"Foo\" WHERE true'; END IF;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("DELETE",) + assert _scan(tmp_path, sql)[0].line == 6 + + def test_a_rewrite_assigned_past_a_loop_comparison_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + " total int := 3;\n" + "BEGIN\n" + " WHILE total >= 1 LOOP stmt = 'UPDATE \"Foo\" SET \"a\" = 1'; END LOOP;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_rewrite_selected_into_a_target_a_line_down_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + "BEGIN\n" + " SELECT 'UPDATE \"Foo\" SET \"a\" = 1'\n" + " INTO\n" + " stmt;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_rewrite_selected_into_a_strict_target_a_line_down_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + "BEGIN\n" + " SELECT 'DELETE FROM \"Foo\"' INTO STRICT\n" + " stmt;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_a_rewrite_executed_by_a_name_a_line_down_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text := 'UPDATE \"Foo\" SET \"a\" = 1';\n" + "BEGIN\n" + " EXECUTE\n" + " stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_rewrite_compared_against_is_not_an_assignment(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text := 'ALTER TABLE \"Foo\" ADD COLUMN \"b\" TEXT';\n" + "BEGIN\n" + " IF stmt = 'UPDATE \"Foo\" SET \"a\" = 1' THEN\n" + " RAISE NOTICE 'the application owns that one';\n" + " END IF;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_rewrite_passed_to_execute_as_a_parameter_is_not_run(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text := 'UPDATE \"Foo\" SET \"a\" = 1';\n" + "BEGIN\n" + " EXECUTE 'INSERT INTO \"Log\" (\"sql\") VALUES ($1)' USING stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + def test_a_literal_selected_into_a_variable_nothing_runs_is_inert(self, tmp_path): sql = ( "DO $$\n" From fdcf867dbfe1fdab304e3c690eb0e50365b254a3 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 10:38:55 -0700 Subject: [PATCH 17/93] fix: read the one assignment a statement holds, and the loop that walks a query Reading every operator let a comparison beside an assignment look like one. `ok := n = 1 AND stmt = ''` registered `stmt` as written, which collided with the `EXECUTE stmt` further down and flagged a block that rewrites nothing. A statement holds one assignment at most, so the search now stops at the first operator that reads as one: everything after it is the expression being assigned, where an `=` only ever compares. Nine shapes were flagged this way, a cast, a `coalesce`, a `format`, a named-argument arrow and the rest, and all of them are valid PL/pgSQL that leaves the table untouched. `INTO` and `USING` no longer count as names an `EXECUTE` runs. Masking blanks a literal in place, so `EXECUTE '' INTO n` left `INTO` looking like the name being run, and an ordinary query reaching the same word collided with it. The docstring claiming that collision was impossible was wrong, and both words are now dropped instead. A loop is a fourth way a literal reaches a variable. `FOR stmt IN SELECT '' LOOP EXECUTE stmt` empties the table and the gate passed it, so the target of a `FOR` or a `FOREACH` is read as assigned too. Reading each statement once rather than once per operator also drops the cost of a statement with thousands of them from seconds to milliseconds. --- .../check_migrations_no_data_rewrites.py | 103 +++++++++----- .../test_check_migrations_no_data_rewrites.py | 134 ++++++++++++++++++ 2 files changed, 198 insertions(+), 39 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 9e6df46899b..dcb81593e3e 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -41,8 +41,9 @@ hide. The SQL an `EXECUTE` runs is scanned the same way, since a rewrite reads t same to Postgres whether it is spelled out or handed over as a string, and so is a literal parked in a variable some `EXECUTE` in the same body then runs by name, however it got there: an assignment with `:=`, the bare `=` PL/pgSQL takes as the -same operator, or a query returning it through `INTO`. So is the body of a `DO` -written in single quotes rather than dollar quotes. A literal nothing runs is text, however much it reads like a statement, so +same operator, a query returning it through `INTO`, or a loop walking the query it +came out of. So is the body of a `DO` written in single quotes rather than dollar +quotes. A literal nothing runs is text, however much it reads like a statement, so an error message naming a `DELETE` the application handles stays a message. Each literal is read on its own, so a keyword built by concatenating fragments that @@ -104,7 +105,8 @@ INTO_TARGETS = re.compile( r"([A-Za-z_][A-Za-z0-9_]*(?:\s*,\s*[A-Za-z_][A-Za-z0-9_]*)*)", re.IGNORECASE, ) -ASSIGNS = re.compile(r":=|(?!:=])=(?!=)") +LOOP_TARGET = re.compile(r"\bFOR(?:EACH)?\s+([A-Za-z_][A-Za-z0-9_]*)\s+IN\b", re.IGNORECASE) +WORD_OR_ASSIGN = re.compile(r"[A-Za-z_][A-Za-z0-9_]*|:=|(?!:=])=(?![=>])") PRECEDING_WORD = re.compile(r"([A-Za-z_][A-Za-z0-9_]*)[^A-Za-z0-9_]*$") EXPLAIN_OPTIONS = re.compile(r"\bEXPLAIN\b(?:\s+(?:ANALYZE|ANALYSE|VERBOSE)\b)+", re.IGNORECASE) @@ -143,6 +145,8 @@ GUARDS_A_CONDITION = frozenset({"IF", "ELSIF", "ELSEIF", "CASE", "WHEN", "WHILE" OPENS_A_BLOCK = frozenset({"BEGIN", "THEN", "ELSE", "LOOP"}) +NEVER_A_VARIABLE = frozenset({"INTO", "USING"}) + GUIDANCE = """ Migrations apply at proxy boot, before it serves traffic, so a statement whose cost scales with table size is downtime. Add the column and let the application backfill @@ -364,28 +368,21 @@ def hands_off_sql(statement: str, executed: frozenset[str]) -> bool: def assigned_names(statement: str) -> frozenset[str]: - """The candidate variable names a statement writes to, taken as every word ahead of the - assignment operator. A declaration carries its type and sometimes a leading `DECLARE` - alongside the name, and none of that is worth parsing when the only question is which - name is executed. PL/pgSQL spells that operator `:=` and takes a bare `=` as the same - thing, so both count, the second only where `assigns_rather_than_compares` reads it as - an assignment. Every operator in the statement is read rather than only the first, since - a comparison earlier on the line would otherwise claim the one slot and hide the - assignment after it: `IF n = 1 THEN stmt = '...'` writes `stmt` at its second `=`. A - query assigns through the target list after its `INTO` instead, which is how a rewrite - reaches a variable with neither operator appearing at all.""" - names: set[str] = set() - - for operator in ASSIGNS.finditer(statement): - reached = reached_words(statement[: operator.start()]) - if operator.group() == ":=" or assigns_rather_than_compares(reached): - names.update(word.lower() for word in reached) + """The candidate variable names a statement writes to. An assignment is read as every + word ahead of its operator, since a declaration carries its type and sometimes a leading + `DECLARE` alongside the name, and none of that is worth parsing when the only question + is which name is executed. A query assigns through the target list after its `INTO` + instead, and a loop through the variable it walks its query with, which is how a rewrite + reaches a variable with no operator appearing at all.""" + names = {word.lower() for word in assignment_reach(statement)} for targets in INTO_TARGETS.finditer(statement): if names_a_table(statement[: targets.start()]): continue names.update(word.group().lower() for word in FIRST_WORD.finditer(targets.group(1))) + names.update(loop.group(1).lower() for loop in LOOP_TARGET.finditer(statement)) + return frozenset(names) @@ -399,33 +396,61 @@ def names_a_table(before: str) -> bool: return word is not None and word.group(1).upper() == "INSERT" -def reached_words(head: str) -> tuple[str, ...]: - """The words an assignment operator is reached through, which is everything since the last - word to open a block. A `THEN` ends the condition its `IF` began, so nothing ahead of it - describes what follows, and neither the name being written nor the keywords that would - mark a comparison ever sit further back than that.""" - words = [word.group().upper() for word in FIRST_WORD.finditer(head)] - opened = max((index for index, word in enumerate(words) if word in OPENS_A_BLOCK), default=-1) - return tuple(words[opened + 1 :]) +def assignment_reach(statement: str) -> tuple[str, ...]: + """The words the statement's assignment is reached through, empty where it holds none. + PL/pgSQL spells the operator `:=` and takes a bare `=` as the same thing, so both count, + the second only where none of the words reached so far `marks_a_comparison`. The search + stops at the first operator that reads as an assignment, because a statement holds one + at most and everything after it is the expression being assigned, where an `=` only ever + compares: that is what keeps `ok := stmt = ''` from reading as a write to `stmt`. + What comes before can still be a comparison the assignment sits behind, as in + `IF n = 1 THEN stmt = ''`, and a word opening a block ends what it is reached + through, since nothing ahead of the `THEN` describes what follows it.""" + reached: list[str] = [] + compares = False + + for token in WORD_OR_ASSIGN.finditer(statement): + word = token.group().upper() + + if word == ":=": + return tuple(reached) + + if word == "=": + if not compares: + return tuple(reached) + continue + + if word in OPENS_A_BLOCK: + reached.clear() + compares = False + continue + + reached.append(word) + compares = compares or marks_a_comparison(word) + + return () -def assigns_rather_than_compares(reached: tuple[str, ...]) -> bool: - """Whether a bare `=` reached through these words writes a variable or tests one. They are - all that tells the two apart: an assignment is reached with a name and perhaps a type, - while a comparison is reached either through a statement carrying its own keyword or - through a word that guards a condition.""" - words = set(reached) - return not (words & STATEMENT_KEYWORDS) and not (words & GUARDS_A_CONDITION) +def marks_a_comparison(word: str) -> bool: + """Whether reaching a bare `=` through this word means the operator tests a variable + rather than writing one. These are all that tell the two apart: an assignment is reached + with a name and perhaps a type, while a comparison is reached either through a statement + carrying its own keyword or through a word that guards a condition.""" + return word in STATEMENT_KEYWORDS or word in GUARDS_A_CONDITION def executed_names(masked: str) -> frozenset[str]: """The variables handed to an `EXECUTE` by name. Reading these off the masked text keeps an `EXECUTE` written inside a comment or a string from counting. Masking blanks a literal - in place rather than removing it, so `EXECUTE '...'` can leave the word after it looking - like the name being run. Reaching that word means crossing no semicolon, which leaves only - the syntax `INTO`, `USING` and `END`, and no assignment ever writes one of those, so the - stray name has nothing to match.""" - return frozenset(match.group(1).lower() for match in RUN_BY_NAME.finditer(masked)) + in place rather than removing it, so `EXECUTE '...'` leaves whatever follows the literal + looking like the name being run. Only `INTO` and `USING` can sit there, since the syntax + allows nothing else between an `EXECUTE` and the semicolon ending it, and neither is ever + a variable, so both are dropped rather than left to collide with a query reaching one.""" + return frozenset( + match.group(1).lower() + for match in RUN_BY_NAME.finditer(masked) + if match.group(1).upper() not in NEVER_A_VARIABLE + ) def leads_with(statement: str, keyword: str) -> bool: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index fe53e1ac823..b93e8422933 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -663,6 +663,140 @@ class TestDynamicSql: ) assert _keywords(tmp_path, sql) == () + def test_a_rewrite_compared_beside_an_assignment_is_not_run(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + " total int := 1;\n" + " ok boolean;\n" + "BEGIN\n" + " stmt := 'CREATE INDEX IF NOT EXISTS \"ix_a\" ON \"Foo\" (\"a\")';\n" + " ok := total = 1 AND stmt = 'DELETE FROM \"Foo\"';\n" + " RAISE NOTICE 'purge script? %', ok;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_rewrite_named_as_an_argument_beside_an_assignment_is_not_run(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + " ok boolean;\n" + "BEGIN\n" + " stmt := 'CREATE INDEX IF NOT EXISTS \"ix_a\" ON \"Foo\" (\"a\")';\n" + " ok := probe_match(subject => stmt, wanted => 'DELETE FROM \"Foo\"');\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_rewrite_compared_after_a_wider_comparison_is_not_run(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + " total int := 1;\n" + " ok boolean;\n" + "BEGIN\n" + " stmt := 'CREATE INDEX IF NOT EXISTS \"ix_a\" ON \"Foo\" (\"a\")';\n" + " ok := total >= 1 AND stmt = 'DELETE FROM \"Foo\"';\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_rewrite_assigned_through_a_case_expression_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + " total int := 1;\n" + "BEGIN\n" + " stmt := CASE WHEN total = 1 THEN 'DELETE FROM \"Foo\"' ELSE 'SELECT 1' END;\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_a_query_reaching_into_past_an_execute_is_not_a_rewrite(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + " total int;\n" + "BEGIN\n" + " stmt := 'CREATE INDEX IF NOT EXISTS \"ix_a\" ON \"Foo\" (\"a\")';\n" + " EXECUTE 'SELECT count(*) FROM \"Foo\"'\n" + " INTO total;\n" + " SELECT (CASE WHEN total > 0 THEN 1 ELSE 2 END) INTO total\n" + " FROM \"Foo\"\n" + " WHERE \"a\" = 'DELETE FROM \"Foo\"';\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_query_reaching_using_past_an_execute_is_not_a_rewrite(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + " total int;\n" + "BEGIN\n" + " stmt := 'CREATE INDEX IF NOT EXISTS \"ix_a\" ON \"Foo\" (\"a\")';\n" + " EXECUTE 'SELECT count(*) FROM \"Foo\" WHERE \"a\" = $1'\n" + " USING 'k1';\n" + " SELECT (CASE WHEN true THEN 1 ELSE 2 END) INTO total\n" + " FROM \"Foo\" x JOIN \"Foo\" y USING (\"a\")\n" + " WHERE x.\"a\" = 'DELETE FROM \"Foo\"';\n" + " EXECUTE stmt;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_rewrite_walked_by_a_loop_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + "BEGIN\n" + " FOR stmt IN SELECT 'DELETE FROM \"Foo\"' LOOP\n" + " EXECUTE stmt;\n" + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("DELETE",) + assert _scan(tmp_path, sql)[0].line == 5 + + def test_a_rewrite_walked_by_a_foreach_loop_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " stmt text;\n" + "BEGIN\n" + " FOREACH stmt IN ARRAY ARRAY['UPDATE \"Foo\" SET \"a\" = 1'] LOOP\n" + " EXECUTE stmt;\n" + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_loop_over_a_query_running_nothing_is_inert(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE\n" + " rec record;\n" + "BEGIN\n" + " FOR rec IN SELECT \"a\" FROM \"Foo\" LOOP\n" + " RAISE NOTICE 'the DELETE FROM \"Foo\" path is the application''s: %', rec;\n" + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + def test_a_literal_selected_into_a_variable_nothing_runs_is_inert(self, tmp_path): sql = ( "DO $$\n" From 6d6a2fcfb892ff69f968b0ec8407a9d08618f84b Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 10:57:00 -0700 Subject: [PATCH 18/93] fix: match a marker to the statement it is written against, not to its line --- .../check_migrations_no_data_rewrites.py | 78 +++++++++++++------ .../test_check_migrations_no_data_rewrites.py | 13 ++++ 2 files changed, 67 insertions(+), 24 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index dcb81593e3e..7b0029c73cb 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -172,16 +172,38 @@ class Violation: @dataclass(frozen=True, slots=True) -class Markers: - lines: frozenset[int] - standalone: frozenset[int] +class Marker: + start: int + end: int + standalone: bool - def exempt(self, first: int, last: int) -> bool: - """Whether a statement spanning `first` to `last` carries a marker. A marker alone on - its line speaks for the statement below it, which is how one written above a rewrite - exempts it. A marker sharing its line with the statement it follows speaks for that - statement only, so the next statement down does not inherit the exemption.""" - return any(line in self.lines for line in range(first, last + 1)) or first - 1 in self.standalone + +@dataclass(frozen=True, slots=True) +class Markers: + sql: str + written: tuple[Marker, ...] + + def exempt(self, start: int, end: int) -> bool: + """Whether the statement spanning `start` to `end` carries a marker.""" + return any(self.speaks_for(marker, start, end) for marker in self.written) + + def speaks_for(self, marker: Marker, start: int, end: int) -> bool: + """Whether a marker is written against this statement. One alone on its line speaks for + the statement below it, which is how a marker written above a rewrite exempts it, and one + sharing its line with code speaks for the statement it follows. Either is matched by where + it sits rather than by the line it lands on, so a second statement sharing that line does + not inherit the exemption. A marker inside a statement speaks for it whichever kind it is, + which is how one on the opening line of a long statement still covers the whole of it.""" + if start <= marker.start < end: + return True + if marker.standalone: + return self.only_separators(marker.end, start) + return self.only_separators(end, marker.start) + + def only_separators(self, start: int, end: int) -> bool: + """Whether nothing but statement separators lie between two points, which is what makes a + marker and a statement adjacent whatever whitespace and line breaks sit between them.""" + return start <= end and not self.sql[start:end].strip(" \t\r\n;") def blank(text: str) -> str: @@ -463,16 +485,17 @@ def contains(statement: str, keyword: str) -> bool: def read_markers(sql: str) -> Markers: - lines: set[int] = set() - standalone: set[int] = set() + return Markers( + sql, + tuple( + Marker(match.start(), match.end(), alone_on_its_line(sql, match.start())) + for match in MARKER.finditer(sql) + ), + ) - for match in MARKER.finditer(sql): - line = sql.count("\n", 0, match.start()) + 1 - lines.add(line) - if not sql[sql.rfind("\n", 0, match.start()) + 1 : match.start()].strip(): - standalone.add(line) - return Markers(frozenset(lines), frozenset(standalone)) +def alone_on_its_line(sql: str, start: int) -> bool: + return not sql[sql.rfind("\n", 0, start) + 1 : start].strip() def scan(sql: str, migration: str, markers: Markers) -> Iterator[Violation]: @@ -482,16 +505,16 @@ def scan(sql: str, migration: str, markers: Markers) -> Iterator[Violation]: def scan_region( document: str, region: str, migration: str, markers: Markers, offset: int ) -> Iterator[Violation]: - """Violations in one region of `document`, whose text begins at `offset`. Lines are - always counted against the whole document, so a statement nested in a dollar-quoted - body reports its real file line and lines up with the markers read from that file.""" + """Violations in one region of `document`, whose text begins at `offset`. Positions are + always counted against the whole document, so a statement nested in a dollar-quoted body + reports its real file line and lines up with the markers read from that file.""" masked, bodies, literals = mask(region) executed = executed_names(masked) for match in STATEMENT.finditer(masked): - first = line_of(document, offset + keyword_start(match)) - last = line_of(document, offset + match.end()) - exempt = markers.exempt(first, last) + start = offset + statement_start(match) + end = offset + match.end() + exempt = markers.exempt(start, end) if hands_off_sql(match.group(), executed) and not exempt: for start, end in literals: @@ -501,12 +524,19 @@ def scan_region( keyword = offending_keyword(match.group()) if keyword is None or exempt: continue - yield Violation(migration, first, keyword) + yield Violation(migration, line_of(document, offset + keyword_start(match)), keyword) for start, end in bodies: yield from scan_region(document, region[start:end], migration, markers, offset + start) +def statement_start(statement: re.Match[str]) -> int: + """Where the statement's own text begins, past the whitespace and blanked comments it picked + up from whatever sat between it and the statement before it, one of which can be a marker.""" + text = statement.group() + return statement.start() + len(text) - len(text.lstrip()) + + def keyword_start(statement: re.Match[str]) -> int: word = leading_keyword(statement.group()) return statement.start() + (0 if word is None else word.start()) diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index b93e8422933..1573a1b5706 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -375,6 +375,19 @@ class TestEscapeHatch: assert _keywords(tmp_path, sql) == ("UPDATE",) assert _scan(tmp_path, sql)[0].line == 1 + def test_a_trailing_marker_exempts_only_the_statement_it_follows(self, tmp_path): + sql = 'DELETE FROM "Foo" WHERE "a" = 1; UPDATE "Bar" SET "b" = 2; -- data-migration-ok: one row' + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_a_marker_above_a_shared_line_exempts_only_the_first_statement_on_it(self, tmp_path): + sql = '-- data-migration-ok: one row\nUPDATE "Foo" SET "a" = 1; DELETE FROM "Bar" WHERE "b" = 2;' + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_a_marker_on_the_opening_line_of_a_statement_exempts_that_statement(self, tmp_path): + sql = 'UPDATE "Foo" -- data-migration-ok: one row\n SET "a" = 1;\nDELETE FROM "Bar";' + assert _keywords(tmp_path, sql) == ("DELETE",) + assert _scan(tmp_path, sql)[0].line == 3 + def test_marker_above_a_do_block_does_not_exempt_a_rewrite_inside_it(self, tmp_path): sql = ( "-- data-migration-ok: bounded, this belongs to the insert below\n" From dee93e2d4842d62531b17eeeb9ec9b37bc30508f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 11:12:26 -0700 Subject: [PATCH 19/93] fix: keep a marker on its own line bound to the statement directly below it --- .../check_migrations_no_data_rewrites.py | 10 ++++++++-- .../test_check_migrations_no_data_rewrites.py | 10 ++++++++++ 2 files changed, 18 insertions(+), 2 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 7b0029c73cb..3d57691cea4 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -197,12 +197,18 @@ class Markers: if start <= marker.start < end: return True if marker.standalone: - return self.only_separators(marker.end, start) + return self.on_the_line_below(marker.end, start) return self.only_separators(end, marker.start) + def on_the_line_below(self, start: int, end: int) -> bool: + """Whether a marker on its own line is written directly above the statement, which means + one line break and nothing else that carries meaning. A blank line between the two leaves + the marker reading as a note about the file rather than a bound on what follows it.""" + return self.only_separators(start, end) and self.sql[start:end].count("\n") == 1 + def only_separators(self, start: int, end: int) -> bool: """Whether nothing but statement separators lie between two points, which is what makes a - marker and a statement adjacent whatever whitespace and line breaks sit between them.""" + marker and the statement it follows adjacent however they are laid out.""" return start <= end and not self.sql[start:end].strip(" \t\r\n;") diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 1573a1b5706..5413f8cdc62 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -375,6 +375,11 @@ class TestEscapeHatch: assert _keywords(tmp_path, sql) == ("UPDATE",) assert _scan(tmp_path, sql)[0].line == 1 + def test_a_marker_a_blank_line_above_a_statement_does_not_exempt_it(self, tmp_path): + sql = '-- data-migration-ok: one row\n\nUPDATE "Foo" SET "a" = 1;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 3 + def test_a_trailing_marker_exempts_only_the_statement_it_follows(self, tmp_path): sql = 'DELETE FROM "Foo" WHERE "a" = 1; UPDATE "Bar" SET "b" = 2; -- data-migration-ok: one row' assert _keywords(tmp_path, sql) == ("DELETE",) @@ -410,6 +415,11 @@ class TestEscapeHatch: assert _keywords(tmp_path, sql) == ("UPDATE",) assert _scan(tmp_path, sql)[0].line == 5 + def test_marker_directly_above_a_one_line_do_block_does_not_exempt_its_body(self, tmp_path): + sql = '-- data-migration-ok: seeding one default row\nDO $$ BEGIN UPDATE "Foo" SET "a" = 1; END $$;' + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 2 + def test_marker_on_the_do_line_does_not_exempt_its_body(self, tmp_path): sql = ( "DO $$ -- data-migration-ok: bounded to one row\n" From e61baa6d0160991f7f7cd4fd1e361536cf4f80a3 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 11:25:31 -0700 Subject: [PATCH 20/93] fix: read one quoted run as one literal, and stop before bind values --- .../check_migrations_no_data_rewrites.py | 32 ++++++++++++++--- .../test_check_migrations_no_data_rewrites.py | 34 +++++++++++++++++++ 2 files changed, 61 insertions(+), 5 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 3d57691cea4..b4907425838 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -147,6 +147,8 @@ OPENS_A_BLOCK = frozenset({"BEGIN", "THEN", "ELSE", "LOOP"}) NEVER_A_VARIABLE = frozenset({"INTO", "USING"}) +BIND_VALUES = re.compile(r"\bUSING\b", re.IGNORECASE) + GUIDANCE = """ Migrations apply at proxy boot, before it serves traffic, so a statement whose cost scales with table size is downtime. Add the column and let the application backfill @@ -286,10 +288,20 @@ def skip_block_comment(sql: str, start: int) -> int: def skip_quoted(sql: str, start: int, quote: str) -> int: - """One quoted run, up to and including its closing quote. A doubled quote needs no - special case: closing on the first and reopening on the second masks the same span.""" - stop = sql.find(quote, start + 1) - return len(sql) if stop == -1 else stop + 1 + """One quoted run, up to and including its closing quote. A doubled quote is an escaped + quote sitting inside the run rather than the end of it. Closing on the first and reopening + on the second would mask the same span, which is why this looked like it needed no special + case, but the run is also handed on whole as one literal, and splitting it there offers the + tail of a string to be read as SQL in its own right.""" + index = start + 1 + while True: + stop = sql.find(quote, index) + if stop == -1: + return len(sql) + if sql[stop + 1 : stop + 2] == quote: + index = stop + 2 + continue + return stop + 1 def strip_parens(statement: str) -> str: @@ -523,8 +535,9 @@ def scan_region( exempt = markers.exempt(start, end) if hands_off_sql(match.group(), executed) and not exempt: + commands_end = match.start() + bind_values_start(match.group()) for start, end in literals: - if match.start() <= start and end <= match.end(): + if match.start() <= start and end <= commands_end: yield from scan_region(document, region[start:end], migration, markers, offset + start) keyword = offending_keyword(match.group()) @@ -536,6 +549,15 @@ def scan_region( yield from scan_region(document, region[start:end], migration, markers, offset + start) +def bind_values_start(statement: str) -> int: + """Where a statement stops handing commands to the server and starts listing bind values. + The expressions after `USING` are values substituted into the command, never commands in + their own right, so one that merely spells out a rewrite is not running it. Read off the + masked text, so a `USING` written inside the command string is not mistaken for this one.""" + keyword = BIND_VALUES.search(statement) + return len(statement) if keyword is None else keyword.start() + + def statement_start(statement: re.Match[str]) -> int: """Where the statement's own text begins, past the whitespace and blanked comments it picked up from whatever sat between it and the statement before it, one of which can be a marker.""" diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 5413f8cdc62..21fafcd205f 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -467,6 +467,40 @@ class TestDynamicSql: sql = "DO $$\nBEGIN\n EXECUTE 'UPDATE \"Foo\" SET \"a\" = date_trunc(''day'', \"t\")';\nEND $$;" assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_a_rewrite_quoted_as_data_inside_executed_sql_is_not_run(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'SELECT ''UPDATE \"Foo\" SET \"a\" = 1''';\nEND $$;" + assert _keywords(tmp_path, sql) == () + + def test_a_doubled_quote_does_not_split_the_literal_it_sits_in(self, tmp_path): + sql = "INSERT INTO \"Foo\" (\"note\") VALUES ('a''UPDATE \"Bar\" SET \"a\" = 1''b');" + assert _keywords(tmp_path, sql) == () + + def test_a_rewrite_following_a_doubled_quote_in_the_same_payload_is_flagged(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'SELECT ''x''; UPDATE \"Foo\" SET \"a\" = 1';\nEND $$;" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_rewrite_in_a_later_command_before_bind_values_is_flagged(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'SELECT 1; DELETE FROM \"Foo\" WHERE \"a\" = $1' USING 1;\nEND $$;" + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_a_bind_value_naming_a_rewrite_is_not_run(self, tmp_path): + sql = ( + "DO $$\nBEGIN\n EXECUTE 'INSERT INTO \"Audit\" (\"note\") VALUES ($1)'" + " USING 'DELETE FROM \"Foo\"';\nEND $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_rewrite_executed_with_bind_values_is_still_flagged(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'DELETE FROM \"Foo\" WHERE \"a\" = $1' USING 1;\nEND $$;" + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_using_written_inside_the_command_does_not_end_it(self, tmp_path): + sql = ( + "DO $$\nBEGIN\n EXECUTE 'DELETE FROM \"Foo\" USING \"Bar\"" + " WHERE \"Foo\".\"a\" = \"Bar\".\"a\"';\nEND $$;" + ) + assert _keywords(tmp_path, sql) == ("DELETE",) + def test_execute_of_ddl_passes(self, tmp_path): sql = "DO $$\nBEGIN\n EXECUTE 'ALTER TABLE \"Foo\" ADD COLUMN \"b\" TEXT';\nEND $$;" assert _keywords(tmp_path, sql) == () From 923da852fa1b245d5edc30bd9f563c52d7220cb7 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 12:03:49 -0700 Subject: [PATCH 21/93] fix: read a loop body as its own statement, not as part of the header A `FOR ... LOOP` header carries no semicolon of its own, so the first statement of the loop body is written into the same semicolon-delimited run. Reading the pair as one statement let the header's row source stand in as the keyword for both, which hid whatever the loop repeats: a plain `UPDATE` in a query-driven loop went unreported, and so did an `EXECUTE` of one. That is the shape a row-by-row backfill takes, and it is the shape this gate exists to stop. --- .../check_migrations_no_data_rewrites.py | 43 ++++-- .../test_check_migrations_no_data_rewrites.py | 134 ++++++++++++++++++ 2 files changed, 162 insertions(+), 15 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index b4907425838..c82d52167b4 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -106,6 +106,7 @@ INTO_TARGETS = re.compile( re.IGNORECASE, ) LOOP_TARGET = re.compile(r"\bFOR(?:EACH)?\s+([A-Za-z_][A-Za-z0-9_]*)\s+IN\b", re.IGNORECASE) +LOOP_HEADER = re.compile(r"\bFOR(?:EACH)?\b.*?\bLOOP\b", re.IGNORECASE | re.DOTALL) WORD_OR_ASSIGN = re.compile(r"[A-Za-z_][A-Za-z0-9_]*|:=|(?!:=])=(?![=>])") PRECEDING_WORD = re.compile(r"([A-Za-z_][A-Za-z0-9_]*)[^A-Za-z0-9_]*$") EXPLAIN_OPTIONS = re.compile(r"\bEXPLAIN\b(?:\s+(?:ANALYZE|ANALYSE|VERBOSE)\b)+", re.IGNORECASE) @@ -530,25 +531,37 @@ def scan_region( executed = executed_names(masked) for match in STATEMENT.finditer(masked): - start = offset + statement_start(match) - end = offset + match.end() - exempt = markers.exempt(start, end) + exempt = markers.exempt(offset + statement_start(match), offset + match.end()) - if hands_off_sql(match.group(), executed) and not exempt: - commands_end = match.start() + bind_values_start(match.group()) - for start, end in literals: - if match.start() <= start and end <= commands_end: - yield from scan_region(document, region[start:end], migration, markers, offset + start) + for clause, base in clauses(match.group(), match.start()): + if hands_off_sql(clause, executed) and not exempt: + commands_end = base + bind_values_start(clause) + for start, end in literals: + if base <= start and end <= commands_end: + yield from scan_region(document, region[start:end], migration, markers, offset + start) - keyword = offending_keyword(match.group()) - if keyword is None or exempt: - continue - yield Violation(migration, line_of(document, offset + keyword_start(match)), keyword) + keyword = offending_keyword(clause) + if keyword is None or exempt: + continue + yield Violation(migration, line_of(document, offset + keyword_start(clause, base)), keyword) for start, end in bodies: yield from scan_region(document, region[start:end], migration, markers, offset + start) +def clauses(statement: str, start: int) -> Iterator[tuple[str, int]]: + """The statements written inside one semicolon-delimited run, each with where it begins. A + `FOR ... LOOP` header takes no semicolon of its own, so the first statement of the loop body + is written into the same run, and reading the pair as one statement lets the header's row + source stand in as the keyword for both. That hides the statement the loop repeats, which is + the shape a row-by-row backfill takes. Splitting after each header, nested ones included, + reads the header and the body as the separate statements Postgres runs them as.""" + edges = (0, *(header.end() for header in LOOP_HEADER.finditer(statement)), len(statement)) + for opens, closes in zip(edges, edges[1:]): + if opens < closes: + yield statement[opens:closes], start + opens + + def bind_values_start(statement: str) -> int: """Where a statement stops handing commands to the server and starts listing bind values. The expressions after `USING` are values substituted into the command, never commands in @@ -565,9 +578,9 @@ def statement_start(statement: re.Match[str]) -> int: return statement.start() + len(text) - len(text.lstrip()) -def keyword_start(statement: re.Match[str]) -> int: - word = leading_keyword(statement.group()) - return statement.start() + (0 if word is None else word.start()) +def keyword_start(clause: str, base: int) -> int: + word = leading_keyword(clause) + return base + (0 if word is None else word.start()) def line_of(sql: str, offset: int) -> int: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 21fafcd205f..970bc9f5264 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -275,6 +275,140 @@ class TestDollarQuotedBlocks: assert _scan(tmp_path, sql)[0].line == 7 +class TestLoopBodies: + def test_a_rewrite_in_a_query_driven_loop_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE r record;\n" + "BEGIN\n" + ' FOR r IN SELECT "id" FROM "Bar" LOOP\n' + ' UPDATE "Foo" SET "a" = 1;\n' + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 5 + + def test_a_delete_in_a_query_driven_loop_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE r record;\n" + "BEGIN\n" + ' FOR r IN SELECT "id" FROM "Bar" LOOP\n' + ' DELETE FROM "Foo" WHERE "id" = r."id";\n' + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_a_join_using_in_the_loop_query_does_not_hide_the_body(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE r record;\n" + "BEGIN\n" + ' FOR r IN SELECT a."id" FROM "A" a JOIN "B" b USING ("id") LOOP\n' + " EXECUTE 'UPDATE \"Foo\" SET \"a\" = 1';\n" + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_rewrite_executed_in_a_query_driven_loop_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE r record;\n" + "BEGIN\n" + ' FOR r IN SELECT "id" FROM "Bar" LOOP\n' + " EXECUTE 'UPDATE \"Foo\" SET \"a\" = 1';\n" + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_rewrite_nested_under_a_guard_inside_a_loop_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE r record;\n" + "BEGIN\n" + ' FOR r IN SELECT "id" FROM "Bar" LOOP\n' + ' IF r."id" > 0 THEN\n' + ' UPDATE "Foo" SET "a" = 1;\n' + " END IF;\n" + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_rewrite_inside_a_nested_loop_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE a record;\n" + "DECLARE b record;\n" + "BEGIN\n" + ' FOR a IN SELECT "id" FROM "A" LOOP FOR b IN SELECT "id" FROM "B" LOOP\n' + ' UPDATE "Foo" SET "a" = 1;\n' + " END LOOP; END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_rewrite_supplying_a_nested_loop_is_flagged(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE a record;\n" + "DECLARE b record;\n" + "BEGIN\n" + ' FOR a IN SELECT "id" FROM "A" LOOP\n' + ' FOR b IN UPDATE "Foo" SET "x" = 1 RETURNING "id" LOOP\n' + " NULL;\n" + " END LOOP; END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 6 + + def test_a_loop_running_only_ddl_passes(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE r record;\n" + "BEGIN\n" + ' FOR r IN SELECT "id" FROM "Bar" LOOP\n' + ' CREATE INDEX "i" ON "Foo"("a");\n' + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_loop_over_a_rewrite_returning_rows_is_flagged_once(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE r record;\n" + "BEGIN\n" + ' FOR r IN UPDATE "Foo" SET "a" = 1 RETURNING "id" LOOP\n' + " NULL;\n" + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_marker_on_a_loop_exempts_the_rewrite_it_repeats(self, tmp_path): + sql = ( + "DO $$\n" + "DECLARE r record;\n" + "BEGIN\n" + " -- data-migration-ok: one row\n" + ' FOR r IN SELECT "id" FROM "Bar" LOOP\n' + ' UPDATE "Foo" SET "a" = 1;\n' + " END LOOP;\n" + "END $$;" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_select_for_update_lock_is_not_read_as_a_loop(self, tmp_path): + sql = 'DO $$\nBEGIN\n PERFORM 1 FROM "Foo" FOR UPDATE;\nEND $$;' + assert _keywords(tmp_path, sql) == () + + class TestQuotingAndComments: def test_update_inside_string_literal_passes(self, tmp_path): sql = 'ALTER TABLE "Foo" ADD COLUMN "note" TEXT NOT NULL DEFAULT \'UPDATE nothing\';' From 12c4652fc424800d14c6d942144064a9a09a09b5 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 14:43:51 -0700 Subject: [PATCH 22/93] fix: read bind values from the USING the command expression has closed A `JOIN ... USING` inside a subquery that helps build an EXECUTE's command was taken for the start of its bind values, so anything written after it went unscanned and a rewrite there was never reported. Only a `USING` with the parentheses closed can be the bind-values clause. --- .../check_migrations_no_data_rewrites.py | 11 +++++++--- .../test_check_migrations_no_data_rewrites.py | 21 +++++++++++++++++++ 2 files changed, 29 insertions(+), 3 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index c82d52167b4..b52289b4fe2 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -566,9 +566,14 @@ def bind_values_start(statement: str) -> int: """Where a statement stops handing commands to the server and starts listing bind values. The expressions after `USING` are values substituted into the command, never commands in their own right, so one that merely spells out a rewrite is not running it. Read off the - masked text, so a `USING` written inside the command string is not mistaken for this one.""" - keyword = BIND_VALUES.search(statement) - return len(statement) if keyword is None else keyword.start() + masked text, so a `USING` written inside the command string is not mistaken for this one, + and only once the parentheses have closed, so that the `USING` of a `JOIN` in a subquery + that helps build the command does not cut the command short and hide the rest of it.""" + for keyword in BIND_VALUES.finditer(statement): + preceding = statement[: keyword.start()] + if preceding.count("(") == preceding.count(")"): + return keyword.start() + return len(statement) def statement_start(statement: re.Match[str]) -> int: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 970bc9f5264..866a0d6bfe6 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -635,6 +635,27 @@ class TestDynamicSql: ) assert _keywords(tmp_path, sql) == ("DELETE",) + def test_a_join_using_in_a_subquery_building_the_command_does_not_end_it(self, tmp_path): + sql = ( + "DO $$\nDECLARE v text;\nBEGIN\n EXECUTE (SELECT v FROM \"A\" x JOIN \"A\" y" + " USING (\"id\")) || 'UPDATE \"Foo\" SET \"a\" = 1';\nEND $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_join_using_does_not_take_the_place_of_the_real_bind_values(self, tmp_path): + sql = ( + "DO $$\nDECLARE v text;\nBEGIN\n EXECUTE (SELECT v FROM \"A\" x JOIN \"A\" y" + " USING (\"id\")) || 'UPDATE \"Foo\" SET \"a\" = $1' USING 2;\nEND $$;" + ) + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_bind_value_naming_a_rewrite_after_a_subquery_join_is_not_run(self, tmp_path): + sql = ( + "DO $$\nDECLARE v text;\nBEGIN\n EXECUTE (SELECT v FROM \"A\" x JOIN \"A\" y" + " USING (\"id\")) USING 'UPDATE \"Foo\" SET \"a\" = 1';\nEND $$;" + ) + assert _keywords(tmp_path, sql) == () + def test_execute_of_ddl_passes(self, tmp_path): sql = "DO $$\nBEGIN\n EXECUTE 'ALTER TABLE \"Foo\" ADD COLUMN \"b\" TEXT';\nEND $$;" assert _keywords(tmp_path, sql) == () From a46f919e8d746b97b9d19fb6f6be0fdb2b37b626 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 15:05:10 -0700 Subject: [PATCH 23/93] fix: read a set-operated insert's row source term by term An `INSERT` whose `VALUES` list holds a scalar subquery was reported as a rewrite whenever that list was not the plain top-level one: joined to another term by `UNION`, `INTERSECT` or `EXCEPT`, or written inside parentheses, which Postgres accepts. Both shapes insert a fixed handful of rows, so the gate was rejecting migrations that do nothing wrong. A set operation is now split into its terms and each is read on its own, since the insert is a rewrite when any one term is a query. A row source kept in parentheses is read on its own terms too. The operators are found outside every parenthesis, so a set operation written inside a `VALUES` list does not cut the list in half. --- .../check_migrations_no_data_rewrites.py | 60 ++++++++++++++++--- .../test_check_migrations_no_data_rewrites.py | 33 ++++++++++ 2 files changed, 86 insertions(+), 7 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index b52289b4fe2..1cb61586060 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -115,6 +115,8 @@ REWRITES_ROWS = frozenset({"UPDATE", "DELETE", "MERGE"}) JOINS_QUERIES = ("UNION", "INTERSECT", "EXCEPT") +SET_OPERATION = re.compile(rf"\b(?:{'|'.join(JOINS_QUERIES)})\b", re.IGNORECASE) + STATEMENT_KEYWORDS = REWRITES_ROWS | frozenset( { "INSERT", @@ -378,20 +380,64 @@ def offending_keyword(statement: str) -> str | None: def row_source_keyword(statement: str) -> str | None: """Which keyword supplies an `INSERT` its rows, or `None` when a literal `VALUES` list does. A query outside every parenthesis is the row source outright. Failing that, a - `VALUES` outside every parenthesis is itself the row source, so the scalar subqueries - and helper CTEs nested within that list do not make the insert a rewrite, though only - while no set operation sits beside it at that same level: one that does joins the list - to a second query term, and that term is the row source however deeply it is - parenthesised. Failing both, the rows come from a parenthesised query, which Postgres + set operation at that same level joins several terms, and the insert is a rewrite when + any one of them is a query, so each term is read on its own rather than the statement + read whole. Failing that, a `VALUES` outside every parenthesis is itself the row source, + so the scalar subqueries and helper CTEs nested within that list do not make the insert + a rewrite. Failing all three, the rows come from a parenthesised group, which Postgres accepts and which reading only the unparenthesised text would let through: `INSERT INTO "t" ("a") (SELECT ...)` copies a whole table.""" outer = strip_parens(statement) joined = row_source_in(outer) if joined is not None: return joined - if contains(outer, "VALUES") and not any(contains(outer, word) for word in JOINS_QUERIES): + if SET_OPERATION.search(outer): + sources = (row_source_keyword(term) for term in set_operation_terms(statement, outer)) + return next((source for source in sources if source is not None), None) + if contains(outer, "VALUES"): return None - return row_source_in(statement) + wrapped = parenthesised_row_source(statement) + return row_source_in(statement) if wrapped is None else row_source_keyword(wrapped) + + +def set_operation_terms(statement: str, outer: str) -> Iterator[str]: + """The terms a top-level set operation joins. The operators are read from the text outside + every parenthesis, which `strip_parens` blanks in place rather than removing, so their + offsets are offsets into the statement itself and each term comes back from the original + text with its own parentheses intact. Reading them at that level is what keeps a set + operation written inside a `VALUES` list from cutting the list in half. An `ALL` or a + `DISTINCT` stays at the head of the term that follows, where it names no row source and + so reads as nothing.""" + edges = [0] + for operation in SET_OPERATION.finditer(outer): + edges += [operation.start(), operation.end()] + edges.append(len(statement)) + + for opens, closes in zip(edges[::2], edges[1::2]): + yield statement[opens:closes] + + +def parenthesised_row_source(statement: str) -> str | None: + """What the last group of parentheses closed at the statement's outermost level holds, + which is where an `INSERT` keeps a row source it has wrapped, the column list before it + being a group of its own. Postgres takes `INSERT INTO "t" ("a") (SELECT ...)` and + `... (VALUES (1))` alike, so reading the wrapped text on its own terms is what stops a + scalar subquery nested inside a wrapped `VALUES` list standing in for the rows.""" + depth = 0 + opens = None + wrapped = None + + for index, character in enumerate(statement): + if character == "(": + if depth == 0: + opens = index + depth += 1 + elif character == ")": + depth = max(depth - 1, 0) + if depth == 0 and opens is not None: + wrapped = statement[opens + 1 : index] + + return wrapped def row_source_in(text: str) -> str | None: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 866a0d6bfe6..c6c41a40c1a 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -161,6 +161,39 @@ class TestInsert: sql = 'INSERT INTO "Foo" ("id") VALUES ((SELECT 1 UNION SELECT 2 LIMIT 1));' assert _keywords(tmp_path, sql) == () + def test_a_scalar_subquery_in_a_set_operated_values_list_stays_bounded(self, tmp_path): + sql = ( + 'INSERT INTO "Foo" ("id") VALUES ((SELECT max("id") FROM "Bar"))' + " UNION ALL VALUES (2);" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_scalar_subquery_in_a_parenthesised_values_list_stays_bounded(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (VALUES ((SELECT max("id") FROM "Bar")));' + assert _keywords(tmp_path, sql) == () + + def test_a_scalar_subquery_in_set_operated_parenthesised_values_lists_stays_bounded(self, tmp_path): + sql = ( + 'INSERT INTO "Foo" ("id") (VALUES ((SELECT max("id") FROM "Bar")))' + " UNION ALL (VALUES (2));" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_set_operation_inside_a_values_list_does_not_split_the_terms(self, tmp_path): + sql = ( + 'INSERT INTO "Foo" ("id") VALUES ((SELECT max("id") FROM "Bar"' + ' UNION SELECT max("id") FROM "Bar")) UNION ALL VALUES (2);' + ) + assert _keywords(tmp_path, sql) == () + + def test_a_query_term_written_before_a_values_term_is_still_the_row_source(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (SELECT "id" FROM "Bar") UNION ALL (VALUES (2));' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_table_term_beside_parenthesised_values_is_still_the_row_source(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (VALUES (1)) UNION ALL (TABLE "Bar");' + assert _keywords(tmp_path, sql) == ("INSERT ... TABLE",) + def test_a_table_row_source_is_flagged(self, tmp_path): assert _keywords(tmp_path, 'INSERT INTO "Foo" TABLE "Bar";') == ("INSERT ... TABLE",) From 7e59f8c2095fb86d9d8186988ca4fc9e9d832623 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 15:20:24 -0700 Subject: [PATCH 24/93] fix: read every parenthesised group for the row source, not the last Taking the last group at the statement's outermost level assumed the row source was written there, and an insert is allowed to carry more after it: `(SELECT ...) ON CONFLICT ("id") DO NOTHING` ends on the conflict target and `... RETURNING ("id")` on the returning list, so the query supplying the rows was never reached and a full table copy passed the gate. Each group is now read on its own terms and the first to name a row source is the answer, since the others are the column list and the clauses an insert may carry, none of which names one. --- .../check_migrations_no_data_rewrites.py | 29 ++++++++++--------- .../test_check_migrations_no_data_rewrites.py | 18 ++++++++++++ 2 files changed, 34 insertions(+), 13 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 1cb61586060..17bc90bd0d6 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -386,7 +386,10 @@ def row_source_keyword(statement: str) -> str | None: so the scalar subqueries and helper CTEs nested within that list do not make the insert a rewrite. Failing all three, the rows come from a parenthesised group, which Postgres accepts and which reading only the unparenthesised text would let through: - `INSERT INTO "t" ("a") (SELECT ...)` copies a whole table.""" + `INSERT INTO "t" ("a") (SELECT ...)` copies a whole table. Each group at that level is + read on its own terms and the first to name a row source is the answer, since the ones + around it are the column list, the conflict target and the rest of the clauses an insert + is allowed to carry, and any of those can be the last group written.""" outer = strip_parens(statement) joined = row_source_in(outer) if joined is not None: @@ -396,8 +399,11 @@ def row_source_keyword(statement: str) -> str | None: return next((source for source in sources if source is not None), None) if contains(outer, "VALUES"): return None - wrapped = parenthesised_row_source(statement) - return row_source_in(statement) if wrapped is None else row_source_keyword(wrapped) + groups = list(parenthesised_groups(statement)) + if not groups: + return row_source_in(statement) + sources = (row_source_keyword(group) for group in groups) + return next((source for source in sources if source is not None), None) def set_operation_terms(statement: str, outer: str) -> Iterator[str]: @@ -417,15 +423,14 @@ def set_operation_terms(statement: str, outer: str) -> Iterator[str]: yield statement[opens:closes] -def parenthesised_row_source(statement: str) -> str | None: - """What the last group of parentheses closed at the statement's outermost level holds, - which is where an `INSERT` keeps a row source it has wrapped, the column list before it - being a group of its own. Postgres takes `INSERT INTO "t" ("a") (SELECT ...)` and - `... (VALUES (1))` alike, so reading the wrapped text on its own terms is what stops a - scalar subquery nested inside a wrapped `VALUES` list standing in for the rows.""" +def parenthesised_groups(statement: str) -> Iterator[str]: + """What each group of parentheses closed at the statement's outermost level holds, in the + order they are written. One of them is where an `INSERT` keeps a row source it has + wrapped, since Postgres takes `INSERT INTO "t" ("a") (SELECT ...)` and `... (VALUES (1))` + alike, and reading a group on its own terms is what stops a scalar subquery nested inside + a wrapped `VALUES` list standing in for the rows.""" depth = 0 opens = None - wrapped = None for index, character in enumerate(statement): if character == "(": @@ -435,9 +440,7 @@ def parenthesised_row_source(statement: str) -> str | None: elif character == ")": depth = max(depth - 1, 0) if depth == 0 and opens is not None: - wrapped = statement[opens + 1 : index] - - return wrapped + yield statement[opens + 1 : index] def row_source_in(text: str) -> str | None: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index c6c41a40c1a..b78b34e52c7 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -186,6 +186,24 @@ class TestInsert: ) assert _keywords(tmp_path, sql) == () + def test_a_conflict_target_after_a_parenthesised_row_source_does_not_hide_it(self, tmp_path): + sql = ( + 'INSERT INTO "Foo" ("id") (SELECT "id" FROM "Bar")' + ' ON CONFLICT ("id") DO NOTHING;' + ) + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_returning_list_after_a_parenthesised_row_source_does_not_hide_it(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (SELECT "id" FROM "Bar") RETURNING ("id");' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_conflict_target_beside_a_bounded_values_list_stays_bounded(self, tmp_path): + sql = ( + 'INSERT INTO "Foo" ("id") (VALUES ((SELECT max("id") FROM "Bar")))' + ' ON CONFLICT ("id") DO NOTHING;' + ) + assert _keywords(tmp_path, sql) == () + def test_a_query_term_written_before_a_values_term_is_still_the_row_source(self, tmp_path): sql = 'INSERT INTO "Foo" ("id") (SELECT "id" FROM "Bar") UNION ALL (VALUES (2));' assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) From 20e92d1e68c10c6e856b2618aa58341947abd587 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 22 Aug 2026 23:10:36 +0000 Subject: [PATCH 25/93] fix(anthropic/bedrock): request summarized adaptive thinking for reasoning_effort and use provider thinking token counts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/anthropic/chat/transformation.py | 9 +- .../bedrock/chat/converse_transformation.py | 24 ++++- litellm/llms/bedrock/chat/invoke_handler.py | 5 ++ litellm/types/llms/anthropic.py | 1 + .../test_reasoning_effort_translation.py | 2 +- .../test_anthropic_reasoning_effort.py | 12 +++ ...azure_anthropic_messages_transformation.py | 2 +- .../chat/test_converse_transformation.py | 90 +++++++++++++++++++ .../llms/bedrock/chat/test_invoke_handler.py | 23 +++++ .../test_anthropic_claude3_transformation.py | 4 +- ...artner_models_anthropic_messages_config.py | 2 +- 11 files changed, 165 insertions(+), 9 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index ef278c8f723..27caa9efc44 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1184,8 +1184,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if reasoning_effort is None or reasoning_effort == "none": return None if AnthropicConfig._is_adaptive_thinking_model(model, custom_llm_provider): + # without display, Anthropic defaults adaptive thinking to + # display="omitted" and returns a blank thinking block return AnthropicThinkingParam( type="adaptive", + display="summarized", ) elif reasoning_effort == "low": return AnthropicThinkingParam( @@ -2113,7 +2116,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) @staticmethod - def _thinking_tokens_from_usage(usage_object: Mapping[str, object]) -> int | None: + def thinking_tokens_from_usage(usage_object: Mapping[str, object]) -> int | None: details: Final = usage_object.get("output_tokens_details") if not isinstance(details, Mapping): return None @@ -2145,7 +2148,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): reported_thinking_tokens: Final = ( iteration_thinking_tokens if iteration_thinking_tokens is not None - else self._thinking_tokens_from_usage(usage_object) + else self.thinking_tokens_from_usage(usage_object) ) if reported_thinking_tokens is not None: capped_reported: Final = min(max(0, reported_thinking_tokens), completion_tokens) @@ -2168,7 +2171,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _sum_iteration_thinking_tokens(self, iterations: Sequence[object]) -> int | None: per_iteration: Final = tuple( - self._thinking_tokens_from_usage(iteration) if isinstance(iteration, Mapping) else None + self.thinking_tokens_from_usage(iteration) if isinstance(iteration, Mapping) else None for iteration in iterations ) reported: Final = tuple(tokens for tokens in per_iteration if tokens is not None) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index b437e25d24b..52366da8c35 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1617,6 +1617,8 @@ class AmazonConverseConfig(BaseConfig): } if additional_request_params: data["additionalModelRequestFields"] = additional_request_params + if "thinking" in additional_request_params: + data["additionalModelResponseFieldPaths"] = ["/usage/output_tokens_details"] if system_content_blocks: data["system"] = system_content_blocks @@ -1801,6 +1803,17 @@ class AmazonConverseConfig(BaseConfig): thinking_blocks_list.append(_redacted_block) return thinking_blocks_list + @staticmethod + def thinking_tokens_from_additional_fields(additional_fields: object) -> int | None: + """Converse omits thinking tokens from its usage block; they only arrive under + ``additionalModelResponseFields`` when ``/usage/output_tokens_details`` is requested.""" + if not isinstance(additional_fields, Mapping): + return None + usage: Final = additional_fields.get("usage") + if not isinstance(usage, Mapping): + return None + return AnthropicConfig.thinking_tokens_from_usage(usage) + @staticmethod def is_converse_usage_shape(usage_object: Mapping[str, object]) -> bool: """Converse-family models report camelCase token counts, not Anthropic's snake_case.""" @@ -1842,6 +1855,7 @@ class AmazonConverseConfig(BaseConfig): usage: ConverseTokenUsageBlock, reasoning_content: str | None = None, thinking_ran: bool = False, + provider_reasoning_tokens: int | None = None, ) -> Usage: input_tokens = usage["inputTokens"] output_tokens: Final = usage["outputTokens"] @@ -1862,9 +1876,14 @@ class AmazonConverseConfig(BaseConfig): cache_creation_tokens=cache_creation_input_tokens, text_tokens=raw_input_tokens, ) - reasoning_tokens: Final = ( + estimated_reasoning_tokens: Final = ( token_counter(text=reasoning_content, count_response_tokens=True) if reasoning_content else 0 ) + reasoning_tokens: Final = ( + min(max(0, provider_reasoning_tokens), output_tokens) + if provider_reasoning_tokens is not None + else estimated_reasoning_tokens + ) completion_tokens_details: Final = ( CompletionTokensDetailsWrapper( reasoning_tokens=reasoning_tokens, @@ -2272,6 +2291,9 @@ class AmazonConverseConfig(BaseConfig): completion_response["usage"], reasoning_content=chat_completion_message.get("reasoning_content"), thinking_ran=reasoningContentBlocks is not None, + provider_reasoning_tokens=self.thinking_tokens_from_additional_fields( + completion_response.get("additionalModelResponseFields") + ), ) ## HANDLE TOOL CALLS diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index ce89c6c23e2..3937b36aca0 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -331,6 +331,7 @@ class AWSEventStreamDecoder: self.json_mode = json_mode self._current_tool_name: str | None = None self._thinking_ran = False + self._provider_reasoning_tokens: int | None = None def check_empty_tool_call_args(self) -> bool: """ @@ -559,10 +560,14 @@ class AWSEventStreamDecoder: tool_use = self._handle_converse_stop_event(content_block_index) elif "stopReason" in chunk_data: finish_reason = map_finish_reason(chunk_data.get("stopReason", "stop")) + self._provider_reasoning_tokens = AmazonConverseConfig.thinking_tokens_from_additional_fields( + chunk_data.get("additionalModelResponseFields") + ) elif "usage" in chunk_data: usage = converse_config.transform_usage( chunk_data.get("usage", {}), thinking_ran=self._thinking_ran, + provider_reasoning_tokens=self._provider_reasoning_tokens, ) if thinking_blocks: self._thinking_ran = True diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index cc6eccbf3e0..d3b0f334163 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -685,6 +685,7 @@ ANTHROPIC_API_ONLY_HEADERS: Final = { # fails if calling anthropic on vertex ai class AnthropicThinkingParam(TypedDict, total=False): type: ReadOnly[Literal["enabled", "adaptive", "disabled"]] budget_tokens: int + display: ReadOnly[Literal["summarized", "omitted"]] class ANTHROPIC_HOSTED_TOOLS(str, Enum): diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py index f393a7b50b1..48a96d011d5 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py @@ -44,7 +44,7 @@ def test_reasoning_effort_maps_to_output_config_for_adaptive_model( ) assert "reasoning_effort" not in result - assert result.get("thinking") == {"type": "adaptive"} + assert result.get("thinking") == {"type": "adaptive", "display": "summarized"} assert result.get("output_config") == {"effort": expected_effort} diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_reasoning_effort.py b/tests/test_litellm/llms/anthropic/test_anthropic_reasoning_effort.py index ef74249ca8e..288817dff07 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_reasoning_effort.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_reasoning_effort.py @@ -5,6 +5,8 @@ Verifies that reasoning_effort=None returns None for all models, including Claude Opus 4.6. """ +import pytest + from litellm.llms.anthropic.chat.transformation import AnthropicConfig @@ -35,6 +37,16 @@ class TestMapReasoningEffort: ) assert result["type"] == "adaptive" + @pytest.mark.parametrize("effort", ["low", "medium", "high"]) + def test_adaptive_mapping_requests_summarized_display(self, effort): + """Regression LIT-5714: adaptive thinking without ``display`` makes Anthropic + return a blank thinking block, so reasoning_effort callers always got + ``reasoning_content: ""``.""" + result = AnthropicConfig._map_reasoning_effort( + reasoning_effort=effort, model="claude-opus-4-6", custom_llm_provider="anthropic" + ) + assert result["display"] == "summarized" + def test_other_model_low_returns_enabled_with_budget(self): result = AnthropicConfig._map_reasoning_effort( reasoning_effort="low", model="claude-4-sonnet-20250514", custom_llm_provider="anthropic" diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py index 53a432427d3..326edde743d 100644 --- a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py +++ b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py @@ -341,7 +341,7 @@ def test_messages_thinking_shape_follows_exact_azure_entry_flag(local_model_cost ) result = transform() - assert result.get("thinking") == {"type": "adaptive"} + assert result.get("thinking") == {"type": "adaptive", "display": "summarized"} assert result.get("output_config") == {"effort": "medium"} monkeypatch.setitem( diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index 604f3414775..b648b6322f7 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -366,6 +366,96 @@ def test_output_config_effort_forwarded_into_additional_request_fields(model): assert additional.get("output_config") == {"effort": "high"} +def test_reasoning_effort_requests_summarized_display_converse(): + """Regression LIT-5714: adaptive thinking synthesized from reasoning_effort must + request the summarized display, otherwise the provider returns a blank thinking + block and reasoning_content is always empty.""" + config = AmazonConverseConfig() + + optional_params = config.map_openai_params( + non_default_params={"reasoning_effort": "high"}, + optional_params={}, + model="bedrock/converse/us.anthropic.claude-opus-4-7", + drop_params=False, + ) + + assert optional_params["thinking"]["type"] == "adaptive" + assert optional_params["thinking"]["display"] == "summarized" + + +def test_thinking_request_adds_output_tokens_details_response_path(): + """Regression LIT-5714: the Converse usage block has no thinking-token field, so + thinking requests must ask for ``/usage/output_tokens_details`` via + ``additionalModelResponseFieldPaths``.""" + config = AmazonConverseConfig() + + result = config._transform_request( + model="bedrock/converse/us.anthropic.claude-opus-4-7", + messages=[{"role": "user", "content": "hi"}], + optional_params={ + "maxTokens": 256, + "thinking": {"type": "adaptive", "display": "summarized"}, + "output_config": {"effort": "high"}, + }, + litellm_params={}, + headers={}, + ) + + assert result["additionalModelResponseFieldPaths"] == ["/usage/output_tokens_details"] + + +def test_request_without_thinking_omits_response_field_paths(): + config = AmazonConverseConfig() + + result = config._transform_request( + model="bedrock/converse/us.anthropic.claude-opus-4-7", + messages=[{"role": "user", "content": "hi"}], + optional_params={"maxTokens": 256}, + litellm_params={}, + headers={}, + ) + + assert "additionalModelResponseFieldPaths" not in result + + +def test_transform_usage_prefers_provider_reasoning_tokens(): + """Regression LIT-5714: provider-reported thinking tokens must win over the + token_counter estimate derived from visible reasoning text.""" + config = AmazonConverseConfig() + + usage = config.transform_usage( + {"inputTokens": 40, "outputTokens": 3002, "totalTokens": 3042}, + reasoning_content="a short reasoning summary", + thinking_ran=True, + provider_reasoning_tokens=1033, + ) + + assert usage.completion_tokens_details.reasoning_tokens == 1033 + assert usage.completion_tokens_details.text_tokens == 3002 - 1033 + + +def test_transform_usage_falls_back_to_estimate_without_provider_tokens(): + config = AmazonConverseConfig() + + usage = config.transform_usage( + {"inputTokens": 40, "outputTokens": 300, "totalTokens": 340}, + reasoning_content="a short reasoning summary", + thinking_ran=True, + ) + + assert usage.completion_tokens_details.reasoning_tokens > 0 + assert usage.completion_tokens_details.reasoning_tokens < 300 + + +def test_thinking_tokens_parsed_from_additional_model_response_fields(): + parsed = AmazonConverseConfig.thinking_tokens_from_additional_fields( + {"usage": {"output_tokens_details": {"thinking_tokens": 92}}} + ) + assert parsed == 92 + assert AmazonConverseConfig.thinking_tokens_from_additional_fields(None) is None + assert AmazonConverseConfig.thinking_tokens_from_additional_fields({"usage": {}}) is None + + @pytest.mark.parametrize( "model,effort,expected_effort", [ diff --git a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py index e2892a6ccee..2c5ff118c85 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py @@ -206,6 +206,29 @@ def test_bedrock_converse_streaming_consistent_id(): ), "All chunk IDs must match the one captured from the messageStart event" +def test_converse_streaming_usage_uses_provider_thinking_tokens(): + """Regression LIT-5714: the messageStop event carries provider thinking tokens + under ``additionalModelResponseFields``; the usage chunk must report them instead + of a token_counter estimate.""" + chunks = [ + { + "contentBlockIndex": 0, + "delta": {"reasoningContent": {"text": "thinking about it"}}, + }, + { + "stopReason": "end_turn", + "additionalModelResponseFields": {"usage": {"output_tokens_details": {"thinking_tokens": 1033}}}, + }, + {"usage": {"inputTokens": 40, "outputTokens": 3002, "totalTokens": 3042}}, + ] + + decoder = AWSEventStreamDecoder(model="bedrock/anthropic.claude-opus-4-7") + parsed = [decoder.converse_chunk_parser(chunk) for chunk in chunks] + + usage = parsed[-1].usage + assert usage.completion_tokens_details.reasoning_tokens == 1033 + + @pytest.mark.asyncio async def test_make_call_does_not_rechunk_stream_by_default(): """Re-chunking the event stream into fixed 1024-byte blocks holds small diff --git a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index d3c28302bf9..1e09afd6919 100644 --- a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -1387,7 +1387,7 @@ def test_bedrock_messages_maps_reasoning_effort_for_adaptive_model( ) assert "reasoning_effort" not in result - assert result.get("thinking") == {"type": "adaptive"} + assert result.get("thinking") == {"type": "adaptive", "display": "summarized"} assert result.get("output_config") == {"effort": expected_effort} @@ -2935,7 +2935,7 @@ def test_bedrock_messages_thinking_shape_follows_exact_bedrock_entry_flag( ) result = transform() - assert result.get("thinking") == {"type": "adaptive"} + assert result.get("thinking") == {"type": "adaptive", "display": "summarized"} assert result.get("output_config") == {"effort": "medium"} monkeypatch.setitem(litellm.model_cost[model], "supports_adaptive_thinking", False) diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py index ba2f20e2337..f19e169dc9e 100644 --- a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py @@ -538,7 +538,7 @@ def test_messages_thinking_shape_follows_exact_vertex_entry_flag(local_model_cos ) result = transform() - assert result.get("thinking") == {"type": "adaptive"} + assert result.get("thinking") == {"type": "adaptive", "display": "summarized"} assert result.get("output_config") == {"effort": "medium"} monkeypatch.setitem( From 418e8ca5e8db8bd8a0e916d6579aec3279cd2395 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 23 Aug 2026 00:00:21 +0000 Subject: [PATCH 26/93] fix(bedrock): build response field paths as an immutable sequence to satisfy the type discipline gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/bedrock/chat/converse_transformation.py | 2 +- litellm/types/llms/bedrock.py | 3 ++- .../llms/bedrock/chat/test_converse_transformation.py | 2 +- type-discipline-budget.json | 2 +- 4 files changed, 5 insertions(+), 4 deletions(-) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 52366da8c35..767677cbcbf 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1618,7 +1618,7 @@ class AmazonConverseConfig(BaseConfig): if additional_request_params: data["additionalModelRequestFields"] = additional_request_params if "thinking" in additional_request_params: - data["additionalModelResponseFieldPaths"] = ["/usage/output_tokens_details"] + data["additionalModelResponseFieldPaths"] = ("/usage/output_tokens_details",) if system_content_blocks: data["system"] = system_content_blocks diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index 5665aa3277a..6ae2e31fe60 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -1,4 +1,5 @@ import json +from collections.abc import Sequence from enum import Enum from typing import TYPE_CHECKING, Any, Final, Literal @@ -396,7 +397,7 @@ class OutputConfigBlock(TypedDict, total=False): class CommonRequestObject(TypedDict, total=False): # common request object across sync + async flows additionalModelRequestFields: dict - additionalModelResponseFieldPaths: list[str] + additionalModelResponseFieldPaths: Sequence[str] inferenceConfig: InferenceConfig system: list[SystemContentBlock] toolConfig: ToolConfigBlock diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index b648b6322f7..4d2c077b548 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -401,7 +401,7 @@ def test_thinking_request_adds_output_tokens_details_response_path(): headers={}, ) - assert result["additionalModelResponseFieldPaths"] == ["/usage/output_tokens_details"] + assert result["additionalModelResponseFieldPaths"] == ("/usage/output_tokens_details",) def test_request_without_thinking_omits_response_field_paths(): diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 627811a7f1d..05098546325 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22805 + "limit": 22804 }, "LIT002": { "limit": 26873 From 528d358c0540970c8bdb3802f6c480b0c3f24a9e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 22 Aug 2026 17:30:18 -0700 Subject: [PATCH 27/93] fix: leave a bounded insert, a bounded writable CTE and an uncalled routine alone A parenthesised VALUES list ended the search for an insert's row source only when no group followed it, so a RETURNING or an ON CONFLICT DO UPDATE carrying a subquery was read as the rows the insert copies. A writable CTE bounded by its own VALUES list was handed the query the statement ends with for the same reason: the WITH branch read the whole statement rather than the part holding the insert. A CREATE FUNCTION or CREATE PROCEDURE body was scanned as if it ran at boot, but defining a routine only stores it. The body is now read when the same migration names the routine somewhere else, so a migration that defines a backfill and then runs it is still caught, and one whose name needed quoting is read either way since quoting is blanked at the call sites too. main() had no test, so neither its exit codes nor the branch the CI gate reads were pinned; a mutant returning 0 on a violation passed the whole suite. Its four outcomes now have tests, along with both directions of each fix above. --- .../check_migrations_no_data_rewrites.py | 120 +++++++++-- .../test_check_migrations_no_data_rewrites.py | 199 +++++++++++++++++- 2 files changed, 299 insertions(+), 20 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 17bc90bd0d6..fa966a214b1 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -8,10 +8,12 @@ plus a doubled heap that plain autovacuum will not give back. What is banned is the row-rewriting DML behind that, not everything whose cost scales that way. A non-concurrent `CREATE INDEX`, an `ALTER COLUMN ... TYPE` that is -not binary coercible, and a volatile `DEFAULT` on a new column all read the whole -table and all pass. That is deliberate: a rule wide enough to reach them fires on -most ordinary migrations, and a marker everyone adds by reflex stops carrying -information. The outage this was written for was a backfill. +not binary coercible, a volatile `DEFAULT` on a new column, a `CREATE TABLE ... AS +SELECT` or `SELECT ... INTO` filling a new table from an existing one, the rename +that pairs with one of those to swap a table out, and a `REFRESH MATERIALIZED VIEW` +all read the whole table and all pass. That is deliberate: a rule wide enough to +reach them fires on most ordinary migrations, and a marker everyone adds by reflex +stops carrying information. The outage this was written for was a backfill. Flagged, per statement, by its leading keyword: @@ -21,10 +23,15 @@ Flagged, per statement, by its leading keyword: INSERT only when its rows come from a query rather than a literal `VALUES` list. The query counts wherever it sits, since Postgres takes it parenthesised, and `TABLE t` is one as much as a `SELECT` is. An - insert bounded by a leading `VALUES` passes, scalar subqueries in that - list included, while a `VALUES` reached through a subquery or joined - to a query by a set operation bounds nothing - WITH a CTE-led statement containing any of the above + insert bounded by a `VALUES` list passes, written bare or in + parentheses, and so do the scalar subqueries in that list and the + `RETURNING` and `ON CONFLICT` clauses written after it, none of which + supply the rows. A `VALUES` reached through a subquery or joined to a + query by a set operation bounds nothing + WITH a CTE-led statement containing any of the above. An `INSERT` is read + against the part of the statement holding it, so a writable CTE + bounded by its own `VALUES` list is not handed the query the statement + ends with as the rows it copies Referential actions (`ON DELETE CASCADE`, `ON UPDATE CASCADE`) are schema, never a statement's leading keyword, so they pass. @@ -37,7 +44,12 @@ told not to run at boot, and a marker is a cheap answer if one ever does. Statements inside dollar-quoted bodies are scanned too. `DO $$ ... $$` is this repo's idiom for conditional DDL, so a body is where an `UPDATE` would otherwise -hide. The SQL an `EXECUTE` runs is scanned the same way, since a rewrite reads the +hide. A `CREATE FUNCTION` or `CREATE PROCEDURE` body is the exception, because +defining a routine only stores it: that body is read when the same migration names +the routine somewhere else, which is what defining a backfill and then running it +looks like, and left alone when nothing calls it. A routine whose name needed +quoting is read either way, since quoting is blanked at the call sites too and a +call written there could never be found. The SQL an `EXECUTE` runs is scanned the same way, since a rewrite reads the same to Postgres whether it is spelled out or handed over as a string, and so is a literal parked in a variable some `EXECUTE` in the same body then runs by name, however it got there: an assignment with `:=`, the bare `=` PL/pgSQL takes as the @@ -110,6 +122,11 @@ LOOP_HEADER = re.compile(r"\bFOR(?:EACH)?\b.*?\bLOOP\b", re.IGNORECASE | re.DOTA WORD_OR_ASSIGN = re.compile(r"[A-Za-z_][A-Za-z0-9_]*|:=|(?!:=])=(?![=>])") PRECEDING_WORD = re.compile(r"([A-Za-z_][A-Za-z0-9_]*)[^A-Za-z0-9_]*$") EXPLAIN_OPTIONS = re.compile(r"\bEXPLAIN\b(?:\s+(?:ANALYZE|ANALYSE|VERBOSE)\b)+", re.IGNORECASE) +DEFINES_A_ROUTINE = re.compile( + r"\bCREATE\b(?:\s+OR\s+REPLACE)?\s+(?:FUNCTION|PROCEDURE)\b", re.IGNORECASE +) +QUALIFIED_NAME = r"(?:\"[^\"]*\"|[A-Za-z_][A-Za-z0-9_$]*)" +ROUTINE_NAME = re.compile(rf"\s*(?:{QUALIFIED_NAME}\s*\.\s*)?({QUALIFIED_NAME})") REWRITES_ROWS = frozenset({"UPDATE", "DELETE", "MERGE"}) @@ -152,6 +169,8 @@ NEVER_A_VARIABLE = frozenset({"INTO", "USING"}) BIND_VALUES = re.compile(r"\bUSING\b", re.IGNORECASE) +WRITES_ROWS = re.compile(r"\bINSERT\b", re.IGNORECASE) + GUIDANCE = """ Migrations apply at proxy boot, before it serves traffic, so a statement whose cost scales with table size is downtime. Add the column and let the application backfill @@ -370,13 +389,30 @@ def offending_keyword(statement: str) -> str | None: if nested is not None: return f"WITH ... {nested}" if contains(statement, "INSERT"): - source = row_source_keyword(statement) + source = insert_row_source(statement) if source is not None: return f"WITH ... INSERT ... {source}" return None +def insert_row_source(statement: str) -> str | None: + """Which keyword supplies the rows to an `INSERT` written somewhere inside a `WITH` + statement. Only the parts that hold that insert are read, because a writable CTE sits + beside the query the statement ends with and reading the whole thing hands the insert + the outer `SELECT` as its row source: `WITH c AS (INSERT ... VALUES (1) RETURNING "x") + SELECT * FROM c` adds one literal row and copies nothing. A CTE keeps its insert in a + parenthesised group, and the statement's own insert, if it is the one writing, runs from + the keyword to the end, found in the text outside every parenthesis so a group's insert + is not counted twice.""" + inserts = [group for group in parenthesised_groups(statement) if contains(group, "INSERT")] + written = WRITES_ROWS.search(strip_parens(statement)) + if written is not None: + inserts.append(statement[written.start() :]) + sources = (row_source_keyword(insert) for insert in inserts) + return next((source for source in sources if source is not None), None) + + def row_source_keyword(statement: str) -> str | None: """Which keyword supplies an `INSERT` its rows, or `None` when a literal `VALUES` list does. A query outside every parenthesis is the row source outright. Failing that, a @@ -387,9 +423,12 @@ def row_source_keyword(statement: str) -> str | None: a rewrite. Failing all three, the rows come from a parenthesised group, which Postgres accepts and which reading only the unparenthesised text would let through: `INSERT INTO "t" ("a") (SELECT ...)` copies a whole table. Each group at that level is - read on its own terms and the first to name a row source is the answer, since the ones - around it are the column list, the conflict target and the rest of the clauses an insert - is allowed to carry, and any of those can be the last group written.""" + read on its own terms until one of them supplies the rows, since the ones before it are + the column list and the ones after it are the conflict target and the rest of the clauses + an insert is allowed to carry. A wrapped `VALUES` list is the row source as much as a + wrapped query is, so it ends the search rather than being skipped over: reading past it + reaches a `RETURNING (SELECT ...)` or a `DO UPDATE SET "a" = (SELECT ...)` written after + it and calls that scalar subquery the rows the insert copies.""" outer = strip_parens(statement) joined = row_source_in(outer) if joined is not None: @@ -402,8 +441,13 @@ def row_source_keyword(statement: str) -> str | None: groups = list(parenthesised_groups(statement)) if not groups: return row_source_in(statement) - sources = (row_source_keyword(group) for group in groups) - return next((source for source in sources if source is not None), None) + for group in groups: + if contains(strip_parens(group), "VALUES"): + return None + source = row_source_keyword(group) + if source is not None: + return source + return None def set_operation_terms(statement: str, outer: str) -> Iterator[str]: @@ -594,10 +638,54 @@ def scan_region( continue yield Violation(migration, line_of(document, offset + keyword_start(clause, base)), keyword) - for start, end in bodies: + for body in bodies: + if not runs_when_applied(masked, region, bodies, body): + continue + start, end = body yield from scan_region(document, region[start:end], migration, markers, offset + start) +def runs_when_applied( + masked: str, region: str, bodies: tuple[tuple[int, int], ...], body: tuple[int, int] +) -> bool: + """Whether a dollar-quoted body runs while the migration is being applied. A `DO` block runs + where it is written, and so does every other use of this quoting. A `CREATE FUNCTION` or a + `CREATE PROCEDURE` only stores its body, which runs when something calls the routine, so a + definition nothing calls rewrites no rows at boot and reporting it names a line that never + executes. Skipping every definition instead would let a migration define a backfill and then + run it unseen, which is the shape this check exists to catch, so the body is read whenever + the same migration names the routine anywhere outside the definition. The definition is + found in the masked text, where one written inside a comment has already been blanked, and + the name is read from the region at those same offsets, since masking blanks a quoted + identifier in place. A name that needed those quotes is blanked at its call sites too and + so can never be found there, which would read as uncalled however the migration runs it, + and the body is read rather than trusted.""" + start, end = body + opens = masked.rfind(";", 0, start) + 1 + defined = DEFINES_A_ROUTINE.search(masked, opens, start) + if defined is None: + return True + named = ROUTINE_NAME.match(region, defined.end(), start) + if named is None or named.group(1).startswith('"'): + return True + return contains(outside_definition(masked, region, bodies, opens, end), re.escape(named.group(1))) + + +def outside_definition( + masked: str, region: str, bodies: tuple[tuple[int, int], ...], opens: int, closes: int +) -> str: + """The migration's text with one routine definition blanked out and every dollar-quoted body + put back. Masking blanks the bodies alike, and a `DO` block is the ordinary way a migration + runs a routine it has just defined, so a call written inside one has to stay readable. The + definition is blanked after they are restored, which takes its own body with it, so a + routine that names itself recursively does not thereby count as called.""" + text = list(masked) + for start, end in bodies: + text[start:end] = region[start:end] + text[opens:closes] = blank(region[opens:closes]) + return "".join(text) + + def clauses(statement: str, start: int) -> Iterator[tuple[str, int]]: """The statements written inside one semicolon-delimited run, each with where it begins. A `FOR ... LOOP` header takes no semicolon of its own, so the first statement of the loop body diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index b78b34e52c7..6a55a7a7c8d 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -1,10 +1,9 @@ """Tests for tests/code_coverage_tests/check_migrations_no_data_rewrites.py. The checker reads migration.sql as SQL rather than as text, so the cases that matter -are the ones a grep would get wrong: `ON DELETE CASCADE` in a foreign key (60-odd -occurrences in the shipped migrations), an `UPDATE` inside a string literal or a -comment, and an `UPDATE` hidden in the `DO $$ ... $$` block this repo uses for -conditional DDL. +are the ones a grep would get wrong: the referential actions in a foreign key, of which +the shipped migrations carry 60, an `UPDATE` inside a string literal or a comment, and +an `UPDATE` hidden in the `DO $$ ... $$` block this repo uses for conditional DDL. """ import importlib.util @@ -218,6 +217,25 @@ class TestInsert: def test_a_table_named_in_the_insert_target_does_not_flag_it(self, tmp_path): assert _keywords(tmp_path, 'INSERT INTO "audit table" ("id") VALUES (1);') == () + def test_a_returning_subquery_after_a_wrapped_values_list_is_not_the_row_source(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (VALUES (1)) RETURNING (SELECT count(*) FROM "Bar");' + assert _keywords(tmp_path, sql) == () + + def test_a_conflict_update_after_a_wrapped_values_list_stays_bounded(self, tmp_path): + sql = ( + 'INSERT INTO "Foo" ("id") (VALUES (1))' + ' ON CONFLICT ("id") DO UPDATE SET "id" = (SELECT max("id") FROM "Bar");' + ) + assert _keywords(tmp_path, sql) == () + + def test_a_wrapped_values_list_of_several_rows_stays_bounded(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (VALUES (1), (2)) RETURNING (SELECT count(*) FROM "Bar");' + assert _keywords(tmp_path, sql) == () + + def test_the_row_source_names_its_own_keyword_not_a_later_subquery(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (TABLE "Bar") RETURNING (SELECT count(*) FROM "Baz");' + assert _keywords(tmp_path, sql) == ("INSERT ... TABLE",) + class TestCommonTableExpressions: def test_cte_led_update_is_flagged(self, tmp_path): @@ -248,6 +266,32 @@ class TestCommonTableExpressions: sql = 'WITH latest AS (SELECT max("id") AS "id" FROM "Bar")\nINSERT INTO "Config" ("k", "v") VALUES (\'rev\', (SELECT "id"::text FROM latest));' assert _keywords(tmp_path, sql) == () + def test_a_writable_cte_bounded_by_values_passes(self, tmp_path): + sql = 'WITH added AS (INSERT INTO "Foo" ("id") VALUES (1) RETURNING "id") SELECT * FROM added;' + assert _keywords(tmp_path, sql) == () + + def test_a_writable_cte_copying_a_query_is_flagged(self, tmp_path): + sql = ( + 'WITH added AS (INSERT INTO "Foo" ("id") SELECT "id" FROM "Bar" RETURNING "id")' + " SELECT * FROM added;" + ) + assert _keywords(tmp_path, sql) == ("WITH ... INSERT ... SELECT",) + + def test_a_bounded_writable_cte_does_not_hide_a_copying_one_beside_it(self, tmp_path): + sql = ( + 'WITH added AS (INSERT INTO "Foo" ("id") VALUES (1) RETURNING "id"),' + ' copied AS (INSERT INTO "Baz" ("id") SELECT "id" FROM "Bar" RETURNING "id")' + " SELECT * FROM added, copied;" + ) + assert _keywords(tmp_path, sql) == ("WITH ... INSERT ... SELECT",) + + def test_a_writable_cte_wrapping_its_row_source_is_flagged(self, tmp_path): + sql = ( + 'WITH added AS (INSERT INTO "Foo" ("id") (SELECT "id" FROM "Bar") RETURNING "id")' + " SELECT * FROM added;" + ) + assert _keywords(tmp_path, sql) == ("WITH ... INSERT ... SELECT",) + class TestDollarQuotedBlocks: def test_update_inside_do_block_is_flagged(self, tmp_path): @@ -326,6 +370,94 @@ class TestDollarQuotedBlocks: assert _scan(tmp_path, sql)[0].line == 7 +class TestStoredRoutines: + DEFINITION = ( + "CREATE FUNCTION backfill() RETURNS void AS $$\n" + "BEGIN\n" + ' UPDATE "Foo" SET "a" = 1;\n' + "END;\n" + "$$ LANGUAGE plpgsql;\n" + ) + PROCEDURE = ( + "CREATE OR REPLACE PROCEDURE sweep() AS $$\n" + "BEGIN\n" + ' DELETE FROM "Foo";\n' + "END;\n" + "$$ LANGUAGE plpgsql;\n" + ) + + def test_a_function_body_nothing_calls_passes(self, tmp_path): + assert _keywords(tmp_path, self.DEFINITION) == () + + def test_a_procedure_body_nothing_calls_passes(self, tmp_path): + assert _keywords(tmp_path, self.PROCEDURE) == () + + def test_a_function_the_migration_calls_is_flagged(self, tmp_path): + assert _keywords(tmp_path, self.DEFINITION + "SELECT backfill();\n") == ("UPDATE",) + + def test_a_procedure_the_migration_calls_is_flagged(self, tmp_path): + assert _keywords(tmp_path, self.PROCEDURE + "CALL sweep();\n") == ("DELETE",) + + def test_a_call_written_above_the_definition_still_counts(self, tmp_path): + assert _keywords(tmp_path, "SELECT backfill();\n" + self.DEFINITION) == ("UPDATE",) + + def test_a_call_from_inside_a_do_block_still_counts(self, tmp_path): + sql = self.DEFINITION + "DO $$ BEGIN PERFORM backfill(); END; $$;\n" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_trigger_wiring_the_function_up_counts_as_a_call(self, tmp_path): + sql = self.DEFINITION + 'CREATE TRIGGER t AFTER INSERT ON "Foo" EXECUTE FUNCTION backfill();\n' + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_schema_qualified_definition_nothing_calls_passes(self, tmp_path): + assert _keywords(tmp_path, self.DEFINITION.replace("backfill()", "public.backfill()")) == () + + def test_a_schema_qualified_function_the_migration_calls_is_flagged(self, tmp_path): + sql = self.DEFINITION.replace("backfill()", "public.backfill()") + "SELECT public.backfill();\n" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_the_name_written_only_in_a_comment_is_not_a_call(self, tmp_path): + sql = self.DEFINITION + "-- backfill() is run by hand after the deploy\n" + assert _keywords(tmp_path, sql) == () + + def test_a_recursive_call_does_not_count_as_the_migration_calling_it(self, tmp_path): + sql = ( + "CREATE FUNCTION backfill(n int) RETURNS void AS $$\n" + "BEGIN\n" + ' UPDATE "Foo" SET "a" = 1;\n' + " PERFORM backfill(n - 1);\n" + "END;\n" + "$$ LANGUAGE plpgsql;\n" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_quoted_routine_name_is_read_rather_than_trusted(self, tmp_path): + sql = self.DEFINITION.replace("backfill()", '"back fill"()') + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_do_block_is_not_a_routine_definition(self, tmp_path): + sql = 'DO $$ BEGIN UPDATE "Foo" SET "a" = 1; END; $$;\n' + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_definition_written_after_another_statement_is_still_recognised(self, tmp_path): + sql = 'ALTER TABLE "Foo" ADD COLUMN "a" INT;\n' + self.DEFINITION + assert _keywords(tmp_path, sql) == () + + def test_a_marker_exempts_a_rewrite_in_a_routine_the_migration_calls(self, tmp_path): + sql = ( + "CREATE FUNCTION backfill() RETURNS void AS $$\n" + "BEGIN\n" + ' UPDATE "Foo" SET "a" = 1; -- data-migration-ok: single config row\n' + "END;\n" + "$$ LANGUAGE plpgsql;\n" + "SELECT backfill();\n" + ) + assert _keywords(tmp_path, sql) == () + + def test_a_called_routine_reports_the_line_inside_its_body(self, tmp_path): + assert _scan(tmp_path, self.DEFINITION + "SELECT backfill();\n")[0].line == 3 + + class TestLoopBodies: def test_a_rewrite_in_a_query_driven_loop_is_flagged(self, tmp_path): sql = ( @@ -1309,3 +1441,62 @@ class TestGrandfathering: class TestShippedMigrations: def test_the_repo_is_clean(self): assert checker.main() == 0 + + +CLEAN = 'ALTER TABLE "Foo" ADD COLUMN "a" INT;' +DIRTY = 'UPDATE "Foo" SET "a" = 1;' +FIXTURE = "20260101000000_fixture" + + +def _tree(monkeypatch, tmp_path: Path, sql: str, grandfathered: frozenset = frozenset()) -> None: + """Stand a migrations directory holding one fixture migration in for the repo's own. The + root moves with it, since a rendered violation names the migration relative to the root and + the two are read off the same checkout everywhere but here.""" + directory = tmp_path / "migrations" / FIXTURE + directory.mkdir(parents=True) + (directory / "migration.sql").write_text(sql, encoding="utf-8") + monkeypatch.setattr(checker, "REPO_ROOT", tmp_path) + monkeypatch.setattr(checker, "MIGRATIONS_DIR", tmp_path / "migrations") + monkeypatch.setattr(checker, "GRANDFATHERED", grandfathered) + + +class TestExitCode: + def test_a_clean_tree_passes(self, tmp_path, monkeypatch): + _tree(monkeypatch, tmp_path, CLEAN) + assert checker.main() == 0 + + def test_a_violation_fails_the_check(self, tmp_path, monkeypatch): + _tree(monkeypatch, tmp_path, DIRTY) + assert checker.main() == 1 + + def test_a_stale_grandfather_alone_fails_the_check(self, tmp_path, monkeypatch): + _tree(monkeypatch, tmp_path, CLEAN, frozenset({FIXTURE})) + assert checker.main() == 1 + + def test_a_grandfathered_violation_passes(self, tmp_path, monkeypatch): + _tree(monkeypatch, tmp_path, DIRTY, frozenset({FIXTURE})) + assert checker.main() == 0 + + def test_a_missing_migrations_directory_is_an_error(self, tmp_path, monkeypatch): + monkeypatch.setattr(checker, "MIGRATIONS_DIR", tmp_path / "absent") + assert checker.main() == 2 + + def test_the_failure_names_the_migration_the_line_and_the_keyword( + self, tmp_path, monkeypatch, capsys + ): + _tree(monkeypatch, tmp_path, DIRTY) + checker.main() + printed = capsys.readouterr().out + assert f"migrations/{FIXTURE}/migration.sql:1" in printed + assert "UPDATE rewrites existing rows at boot" in printed + assert checker.GUIDANCE in printed + + def test_a_stale_grandfather_is_named(self, tmp_path, monkeypatch, capsys): + _tree(monkeypatch, tmp_path, CLEAN, frozenset({FIXTURE})) + checker.main() + assert f"{FIXTURE}: listed in GRANDFATHERED" in capsys.readouterr().out + + def test_a_missing_directory_is_reported_on_stderr(self, tmp_path, monkeypatch, capsys): + monkeypatch.setattr(checker, "MIGRATIONS_DIR", tmp_path / "absent") + checker.main() + assert "migrations directory not found" in capsys.readouterr().err From 20a82cad8a3a66dc716995d450f04bd976efa5b1 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 24 Aug 2026 10:23:17 -0700 Subject: [PATCH 28/93] fix: read a wrapped group before VALUES can end the search, and blank comments in restored bodies --- .../check_migrations_no_data_rewrites.py | 71 +++++++++++++++++-- .../test_check_migrations_no_data_rewrites.py | 32 +++++++++ 2 files changed, 96 insertions(+), 7 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index fa966a214b1..8539a258083 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -428,7 +428,10 @@ def row_source_keyword(statement: str) -> str | None: an insert is allowed to carry. A wrapped `VALUES` list is the row source as much as a wrapped query is, so it ends the search rather than being skipped over: reading past it reaches a `RETURNING (SELECT ...)` or a `DO UPDATE SET "a" = (SELECT ...)` written after - it and calls that scalar subquery the rows the insert copies.""" + it and calls that scalar subquery the rows the insert copies. The group is read on its + own terms before it is allowed to end the search, because a `VALUES` list joined to a + query by a set operation inside the group supplies every row the query does, and + stopping on the word `VALUES` alone would pass the whole copy.""" outer = strip_parens(statement) joined = row_source_in(outer) if joined is not None: @@ -442,11 +445,11 @@ def row_source_keyword(statement: str) -> str | None: if not groups: return row_source_in(statement) for group in groups: - if contains(strip_parens(group), "VALUES"): - return None source = row_source_keyword(group) if source is not None: return source + if contains(strip_parens(group), "VALUES"): + return None return None @@ -676,16 +679,70 @@ def outside_definition( ) -> str: """The migration's text with one routine definition blanked out and every dollar-quoted body put back. Masking blanks the bodies alike, and a `DO` block is the ordinary way a migration - runs a routine it has just defined, so a call written inside one has to stay readable. The - definition is blanked after they are restored, which takes its own body with it, so a - routine that names itself recursively does not thereby count as called.""" + runs a routine it has just defined, so a call written inside one has to stay readable. Each + body comes back with its comments blanked, since a name written in a comment is + documentation rather than a call, while its string literals stay readable because `EXECUTE` + runs one as SQL and the call can be written inside it. The definition is blanked after they + are restored, which takes its own body with it, so a routine that names itself recursively + does not thereby count as called.""" text = list(masked) for start, end in bodies: - text[start:end] = region[start:end] + text[start:end] = without_comments(region[start:end]) text[opens:closes] = blank(region[opens:closes]) return "".join(text) +def without_comments(sql: str) -> str: + """The text with its comments blanked in place and everything else kept, read with the same + lexing as `mask` so a `--` inside a string literal blanks nothing. A dollar-quoted body + nested within is read the same way on its own, which keeps a stray quote inside it from + reaching past its closing tag.""" + chunks: list[str] = [] + index = 0 + length = len(sql) + + while index < length: + pair = sql[index : index + 2] + + if pair == "--": + stop = sql.find("\n", index) + stop = length if stop == -1 else stop + chunks.append(blank(sql[index:stop])) + index = stop + continue + + if pair == "/*": + stop = skip_block_comment(sql, index) + chunks.append(blank(sql[index:stop])) + index = stop + continue + + character = sql[index] + + if character in "'\"": + stop = skip_quoted(sql, index, character) + chunks.append(sql[index:stop]) + index = stop + continue + + if character == "$": + tag = DOLLAR_TAG.match(sql, index) + if tag is not None: + closing = sql.find(tag.group(), tag.end()) + body_end = length if closing == -1 else closing + stop = length if closing == -1 else closing + len(tag.group()) + chunks.append(sql[index : tag.end()]) + chunks.append(without_comments(sql[tag.end() : body_end])) + chunks.append(sql[body_end:stop]) + index = stop + continue + + chunks.append(character) + index += 1 + + return "".join(chunks) + + def clauses(statement: str, start: int) -> Iterator[tuple[str, int]]: """The statements written inside one semicolon-delimited run, each with where it begins. A `FOR ... LOOP` header takes no semicolon of its own, so the first statement of the loop body diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 6a55a7a7c8d..932572a0264 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -236,6 +236,18 @@ class TestInsert: sql = 'INSERT INTO "Foo" ("id") (TABLE "Bar") RETURNING (SELECT count(*) FROM "Baz");' assert _keywords(tmp_path, sql) == ("INSERT ... TABLE",) + def test_a_select_term_wrapped_beside_a_values_term_is_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (VALUES (1) UNION ALL SELECT "id" FROM "Bar");' + assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",) + + def test_a_table_term_wrapped_beside_a_values_term_is_flagged(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (VALUES (1) UNION ALL TABLE "Bar");' + assert _keywords(tmp_path, sql) == ("INSERT ... TABLE",) + + def test_a_wrapped_set_operation_of_values_lists_stays_bounded(self, tmp_path): + sql = 'INSERT INTO "Foo" ("id") (VALUES (1) UNION ALL VALUES (2));' + assert _keywords(tmp_path, sql) == () + class TestCommonTableExpressions: def test_cte_led_update_is_flagged(self, tmp_path): @@ -420,6 +432,26 @@ class TestStoredRoutines: sql = self.DEFINITION + "-- backfill() is run by hand after the deploy\n" assert _keywords(tmp_path, sql) == () + def test_the_name_written_only_in_a_do_body_comment_is_not_a_call(self, tmp_path): + sql = self.DEFINITION + "DO $$ BEGIN\n-- backfill() is run by hand after the deploy\nPERFORM 1;\nEND; $$;\n" + assert _keywords(tmp_path, sql) == () + + def test_the_name_written_only_in_a_do_body_block_comment_is_not_a_call(self, tmp_path): + sql = self.DEFINITION + "DO $$ BEGIN /* backfill() runs later */ PERFORM 1; END; $$;\n" + assert _keywords(tmp_path, sql) == () + + def test_the_name_written_only_in_a_nested_body_comment_is_not_a_call(self, tmp_path): + sql = self.DEFINITION + "DO $$ BEGIN EXECUTE $q$SELECT 1 -- backfill() runs later\n$q$; END; $$;\n" + assert _keywords(tmp_path, sql) == () + + def test_the_name_written_in_an_executed_literal_counts_as_a_call(self, tmp_path): + sql = self.DEFINITION + "DO $$ BEGIN EXECUTE 'SELECT backfill()'; END; $$;\n" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_call_after_a_literal_holding_comment_dashes_still_counts(self, tmp_path): + sql = self.DEFINITION + "DO $$ BEGIN RAISE NOTICE '--'; PERFORM backfill(); END; $$;\n" + assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_a_recursive_call_does_not_count_as_the_migration_calling_it(self, tmp_path): sql = ( "CREATE FUNCTION backfill(n int) RETURNS void AS $$\n" From 03a676995aadef50179e982ddda19e8b39bbfdee Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 24 Aug 2026 11:08:24 -0700 Subject: [PATCH 29/93] feat(search): add Grounding with Bing Search (bing_grounding) as a search provider --- litellm/llms/azure/search/__init__.py | 3 + litellm/llms/azure/search/transformation.py | 353 ++++++++++++++++++ ...odel_prices_and_context_window_backup.json | 8 + .../bing_grounding_websearch_config.yaml | 40 ++ litellm/types/utils.py | 1 + litellm/utils.py | 2 + model_prices_and_context_window.json | 8 + .../test_bing_grounding_search.py | 187 ++++++++++ .../foundry_responses_web_search_fixture.json | 78 ++++ ...st_bing_grounding_search_transformation.py | 311 +++++++++++++++ .../search/test_base_search_transformation.py | 3 + .../public/assets/logos/bing.png | Bin 0 -> 31955 bytes .../_components/CreateSearchTools.tsx | 2 + 13 files changed, 996 insertions(+) create mode 100644 litellm/llms/azure/search/__init__.py create mode 100644 litellm/llms/azure/search/transformation.py create mode 100644 litellm/proxy/example_config_yaml/bing_grounding_websearch_config.yaml create mode 100644 tests/search_tests/test_bing_grounding_search.py create mode 100644 tests/test_litellm/llms/azure/search/foundry_responses_web_search_fixture.json create mode 100644 tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py create mode 100644 ui/litellm-dashboard/public/assets/logos/bing.png diff --git a/litellm/llms/azure/search/__init__.py b/litellm/llms/azure/search/__init__.py new file mode 100644 index 00000000000..2414ba2b1e8 --- /dev/null +++ b/litellm/llms/azure/search/__init__.py @@ -0,0 +1,3 @@ +from litellm.llms.azure.search.transformation import BingGroundingSearchConfig + +__all__ = ("BingGroundingSearchConfig",) diff --git a/litellm/llms/azure/search/transformation.py b/litellm/llms/azure/search/transformation.py new file mode 100644 index 00000000000..2caee9b50a0 --- /dev/null +++ b/litellm/llms/azure/search/transformation.py @@ -0,0 +1,353 @@ +""" +Calls the Microsoft Foundry Responses API with the `bing_grounding` or `web_search` +tool to search the web (Grounding with Bing Search). + +Microsoft docs: https://learn.microsoft.com/en-us/azure/ai-foundry/agents/how-to/tools/bing-grounding + +Setup: + 1. Set BING_GROUNDING_PROJECT_ENDPOINT to the Foundry project endpoint, e.g. + https://.services.ai.azure.com/api/projects/ + 2. Set BING_GROUNDING_MODEL to a model deployment in that project (e.g. gpt-4.1); + it runs the grounded search and its tokens are billed on that deployment + 3. Optional: set BING_GROUNDING_CONNECTION_ID to a Grounding with Bing Search + project connection id to use the `bing_grounding` tool; without it the + project's built-in `web_search` tool is used + 4. Auth: pass api_key, or set BING_GROUNDING_TOKEN to an Entra bearer token for + scope https://ai.azure.com/.default, or configure azure-identity + (AZURE_CLIENT_ID / AZURE_CLIENT_SECRET / AZURE_TENANT_ID, managed identity, + or any DefaultAzureCredential source) and the token is minted automatically + +Usage: + response = litellm.search( + query="latest AI developments", + search_provider="bing_grounding", + max_results=5, + ) +""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal + +import httpx +from pydantic import BaseModel, ConfigDict, ValidationError + +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.search.transformation import ( + BaseSearchConfig, + SearchResponse, + SearchResult, +) +from litellm.secret_managers.main import get_secret_str + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +_DOCS_URL: Final = "https://learn.microsoft.com/en-us/azure/ai-foundry/agents/how-to/tools/bing-grounding" + +PROJECT_ENDPOINT_ENV: Final = "BING_GROUNDING_PROJECT_ENDPOINT" +MODEL_ENV: Final = "BING_GROUNDING_MODEL" +CONNECTION_ID_ENV: Final = "BING_GROUNDING_CONNECTION_ID" +TOKEN_ENV: Final = "BING_GROUNDING_TOKEN" + +ENTRA_SCOPE: Final = "https://ai.azure.com/.default" + +_RESPONSES_PATH: Final = "/openai/v1/responses" +_SNIPPET_FALLBACK_LENGTH: Final = 300 + + +class _Annotation(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + type: str = "" + url: str | None = None + title: str | None = None + start_index: int | None = None + end_index: int | None = None + + +class _ContentPart(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + type: str = "" + text: str = "" + annotations: tuple[_Annotation, ...] = () + + +class _OutputItem(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + type: str = "" + content: tuple[_ContentPart, ...] = () + + +class _ResponsesEnvelope(BaseModel): + """A Foundry Responses API body. `output` is required: a body without it is not a + Responses API response and must not be reported as a successful empty search.""" + + model_config = ConfigDict(extra="ignore", frozen=True) + + output: tuple[_OutputItem, ...] + + +class _ErrorBody(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + message: str | None = None + + +class _ErrorEnvelope(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + error: _ErrorBody | None = None + + +def _unwrap_error_detail(error_message: str) -> str: + """ + Surface the human-readable message inside Foundry's error envelope. + + Tool failures nest a second JSON document as a string inside `error.message` + (observed live for `bing_grounding` connection errors), so the unwrap runs twice. + Falls back to the raw body for anything else. + """ + try: + envelope: Final = _ErrorEnvelope.model_validate_json(error_message) + except ValidationError: + return error_message + message: Final = envelope.error.message if envelope.error else None + if message is None: + return error_message + try: + nested: Final = _ErrorBody.model_validate_json(message) + except ValidationError: + return message + return nested.message or message + + +def _snippet(text: str, annotation: _Annotation) -> str: + """ + The text a citation supports, not the citation marker itself. + + A url_citation's start/end indices span the inline marker ("([host](url))"), + which follows the claim it backs, so the snippet is the marker's own line up + to where the marker starts. + """ + start: Final = annotation.start_index + marker_start: Final = start if start is not None and 0 <= start <= len(text) else len(text) + claim: Final = text[:marker_start].rsplit("\n", 1)[-1].strip() + if claim: + return claim[-_SNIPPET_FALLBACK_LENGTH:] + return text[:_SNIPPET_FALLBACK_LENGTH] + + +def _citation_results(envelope: _ResponsesEnvelope) -> tuple[SearchResult, ...]: + """One result per cited URL: first occurrence wins, order preserved as answered.""" + cited: Final = tuple( + SearchResult( + title=annotation.title or "", + url=annotation.url or "", + snippet=_snippet(part.text, annotation), + date=None, + last_updated=None, + ) + for item in envelope.output + if item.type == "message" + for part in item.content + if part.type == "output_text" + for annotation in part.annotations + if annotation.type == "url_citation" and annotation.url + ) + first_by_url: Final = MappingProxyType({result.url: result for result in reversed(cited)}) + return tuple(first_by_url[url] for url in dict.fromkeys(result.url for result in cited)) + + +class _SearchConfiguration(BaseModel): + model_config = ConfigDict(frozen=True) + + project_connection_id: str + count: int | None = None + + +class _BingGroundingParams(BaseModel): + model_config = ConfigDict(frozen=True) + + search_configurations: tuple[_SearchConfiguration, ...] + + +class _BingGroundingTool(BaseModel): + model_config = ConfigDict(frozen=True) + + type: Literal["bing_grounding"] = "bing_grounding" + bing_grounding: _BingGroundingParams + + +class _UserLocation(BaseModel): + model_config = ConfigDict(frozen=True) + + type: Literal["approximate"] = "approximate" + country: str + + +class _WebSearchTool(BaseModel): + model_config = ConfigDict(frozen=True) + + type: Literal["web_search"] = "web_search" + user_location: _UserLocation | None = None + + +class _ResponsesRequest(BaseModel): + model_config = ConfigDict(frozen=True) + + model: str + input: str + tools: tuple[_BingGroundingTool | _WebSearchTool, ...] + + +def _search_tool(optional_params: Mapping[str, object]) -> _BingGroundingTool | _WebSearchTool: + connection_id: Final = get_secret_str(CONNECTION_ID_ENV) + max_results: Final = optional_params.get("max_results") + country: Final = optional_params.get("country") + if connection_id: + configuration: Final = _SearchConfiguration( + project_connection_id=connection_id, + count=max_results if isinstance(max_results, int) else None, + ) + return _BingGroundingTool(bing_grounding=_BingGroundingParams(search_configurations=(configuration,))) + location: Final = _UserLocation(country=country.upper()) if isinstance(country, str) else None + return _WebSearchTool(user_location=location) + + +def _default_entra_token_minter() -> str: + from litellm.secret_managers.get_azure_ad_token_provider import get_azure_ad_token_provider + + return get_azure_ad_token_provider(azure_scope=ENTRA_SCOPE)() + + +class BingGroundingSearchConfig(BaseSearchConfig): + def __init__(self, entra_token_minter: Callable[[], str] | None = None) -> None: + super().__init__() + self._entra_token_minter = entra_token_minter + + @staticmethod + def ui_friendly_name() -> str: + return "Grounding with Bing Search" + + def validate_environment( + self, + headers: dict[str, str], # mutable-ok: BaseSearchConfig.validate_environment signature + api_key: str | None = None, + api_base: str | None = None, + **kwargs: object, # kwargs-ok: BaseSearchConfig.validate_environment signature + ) -> dict[str, str]: # mutable-ok: the http handler passes this straight to httpx as headers + """ + Validate environment and return headers. + + Returns a new dict rather than mutating ``headers``: the http handler calls this + a second time after ``litellm/search/main.py`` already did, so it has to be idempotent. + """ + resolved_token: Final = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=(TOKEN_ENV,), + base_env_var=PROJECT_ENDPOINT_ENV, + default_api_base=None, + ) or self._mint_entra_token(api_base) + return { # mutable-ok: httpx requires a plain dict of headers + **headers, + "Authorization": f"Bearer {resolved_token}", + "Content-Type": "application/json", + } + + def _mint_entra_token(self, caller_api_base: str | None) -> str: + self._assert_trusted_api_base_for_server_credential( + caller_api_base, None, PROJECT_ENDPOINT_ENV, "Azure AD token" + ) + minter: Final = self._entra_token_minter or _default_entra_token_minter + try: + return minter() + except Exception as e: + raise ValueError( + f"Grounding with Bing Search: no credential available. Pass api_key, set {TOKEN_ENV} " + f"to an Entra bearer token, or configure azure-identity (AZURE_CLIENT_ID / " + f"AZURE_CLIENT_SECRET / AZURE_TENANT_ID or any DefaultAzureCredential source) " + f"for scope {ENTRA_SCOPE}. Underlying error: {e}" + ) from e + + def get_complete_url( + self, + api_base: str | None, + optional_params: dict[str, object], # mutable-ok: BaseSearchConfig.get_complete_url signature + data: dict[str, object] | list[dict[str, object]] | None = None, # mutable-ok: base signature + **kwargs: object, # kwargs-ok: BaseSearchConfig.get_complete_url signature + ) -> str: + resolved_base: Final = api_base or get_secret_str(PROJECT_ENDPOINT_ENV) + if not resolved_base: + raise ValueError( + f"{PROJECT_ENDPOINT_ENV} is not set. Set it to your Microsoft Foundry project " + f"endpoint, e.g. https://.services.ai.azure.com/api/projects/." + ) + trimmed: Final = resolved_base.rstrip("/") + if trimmed.endswith(_RESPONSES_PATH): + return trimmed + return f"{trimmed}{_RESPONSES_PATH}" + + def transform_search_request( + self, + query: str | list[str], # mutable-ok: BaseSearchConfig.transform_search_request signature + optional_params: dict[str, object], # mutable-ok: base signature + **kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_request signature + ) -> dict[str, object]: # mutable-ok: the http handler passes this straight to httpx as the JSON body + """ + Transform Search request to the Foundry Responses API format. + + The unified params map as far as the API allows: + - max_results -> the bing_grounding search configuration's `count` (the built-in + web_search tool has no result-count knob, so it is dropped in that mode) + - country -> web_search's approximate `user_location` (bing_grounding's `market` + wants a full locale like en-US, which a bare country code cannot fill) + - search_domain_filter, max_tokens_per_page -> no API equivalent, dropped + """ + model: Final = get_secret_str(MODEL_ENV) + if not model: + raise ValueError( + f"{MODEL_ENV} is not set. Set it to a model deployment in the Foundry project " + f"that runs the grounded search, e.g. gpt-4.1." + ) + request: Final = _ResponsesRequest( + model=model, + input=" ".join(query) if isinstance(query, list) else query, + tools=(_search_tool(optional_params),), + ) + return request.model_dump(mode="json", exclude_none=True) + + def transform_search_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + **kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_response signature + ) -> SearchResponse: + try: + parsed: Final = _ResponsesEnvelope.model_validate_json(raw_response.content) + except ValidationError as e: + raise self.get_error_class( + error_message=f"response does not match the Foundry Responses API schema: {e}", + status_code=raw_response.status_code, + headers=dict(raw_response.headers), # mutable-ok: BaseSearchConfig.get_error_class signature + ) + results: Final = list(_citation_results(parsed)) # mutable-ok: SearchResponse.results is list[SearchResult] + return SearchResponse(results=results, object="search") + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: dict[str, str], # mutable-ok: BaseSearchConfig.get_error_class signature + ) -> Exception: + detail: Final = _unwrap_error_detail(error_message).rstrip(". ") + return BaseLLMException( + status_code=status_code, + message=f"Grounding with Bing Search: {detail}. See {_DOCS_URL} for details.", + headers=headers, + ) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 3af7d9e5019..9ba8b846e11 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -16890,6 +16890,14 @@ "notes": "Web Search on Amazon Bedrock AgentCore, billed by AWS on the gateway" } }, + "bing_grounding/search": { + "input_cost_per_query": 0.035, + "litellm_provider": "bing_grounding", + "mode": "search", + "metadata": { + "notes": "Grounding with Bing Search (G1 SKU): $35 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment." + } + }, "tinyfish/search": { "input_cost_per_query": 0.0, "litellm_provider": "tinyfish", diff --git a/litellm/proxy/example_config_yaml/bing_grounding_websearch_config.yaml b/litellm/proxy/example_config_yaml/bing_grounding_websearch_config.yaml new file mode 100644 index 00000000000..5ab18723f7f --- /dev/null +++ b/litellm/proxy/example_config_yaml/bing_grounding_websearch_config.yaml @@ -0,0 +1,40 @@ +# Web search via Microsoft Foundry: Grounding with Bing Search / the built-in +# web_search tool, called through the Foundry Responses API. +# See litellm/llms/azure/search/transformation.py for details. +# +# Required environment variables (the search router forwards only +# search_provider / api_key / api_base from the litellm_params block, so +# provider configuration rides env vars): +# BING_GROUNDING_PROJECT_ENDPOINT: the Foundry project endpoint, e.g. +# https://.services.ai.azure.com/api/projects/ +# BING_GROUNDING_MODEL: a model deployment in that project (e.g. gpt-4.1); +# it runs the grounded search, its tokens are billed on that deployment +# Optional: +# BING_GROUNDING_CONNECTION_ID: a Grounding with Bing Search project +# connection id; set it to use the bing_grounding tool ($35 per 1,000 +# transactions on the G1 SKU). Without it the project's built-in +# web_search tool is used +# BING_GROUNDING_TOKEN: an Entra bearer token for scope +# https://ai.azure.com/.default. Without it (and without api_key below) +# the token is minted via azure-identity (AZURE_CLIENT_ID / +# AZURE_CLIENT_SECRET / AZURE_TENANT_ID, managed identity, or any other +# DefaultAzureCredential source) + +model_list: + - model_name: claude-sonnet + litellm_params: + model: bedrock/us.anthropic.claude-sonnet-5 + aws_region_name: us-east-1 + +search_tools: + - search_tool_name: bing-grounding-search + litellm_params: + search_provider: bing_grounding + # Alternative to BING_GROUNDING_TOKEN / azure-identity: + # api_key: os.environ/BING_GROUNDING_TOKEN + +litellm_settings: + callbacks: ["websearch_interception"] + websearch_interception_params: + enabled_providers: ["bedrock"] + search_tool_name: bing-grounding-search diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 67eae2b4f21..371ec7d3375 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3844,6 +3844,7 @@ class SearchProviders(str, Enum): TINYFISH = "tinyfish" AGENTCORE = "agentcore" NIMBLE = "nimble" + BING_GROUNDING = "bing_grounding" # Create a set of all search provider values for quick lookup diff --git a/litellm/utils.py b/litellm/utils.py index e5ce7157e77..006717df187 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9111,6 +9111,7 @@ class ProviderConfigManager: from litellm.llms.apiserpent.search.transformation import ( APISerpentSearchConfig, ) + from litellm.llms.azure.search.transformation import BingGroundingSearchConfig from litellm.llms.bedrock.search.transformation import AgentCoreSearchConfig from litellm.llms.brave.search.transformation import BraveSearchConfig from litellm.llms.dataforseo.search.transformation import DataForSEOSearchConfig @@ -9152,6 +9153,7 @@ class ProviderConfigManager: SearchProviders.TINYFISH: TinyfishSearchConfig, SearchProviders.AGENTCORE: AgentCoreSearchConfig, SearchProviders.NIMBLE: NimbleSearchConfig, + SearchProviders.BING_GROUNDING: BingGroundingSearchConfig, } config_class: Final = PROVIDER_TO_CONFIG_MAP.get(provider, None) if config_class is None: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 3af7d9e5019..9ba8b846e11 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -16890,6 +16890,14 @@ "notes": "Web Search on Amazon Bedrock AgentCore, billed by AWS on the gateway" } }, + "bing_grounding/search": { + "input_cost_per_query": 0.035, + "litellm_provider": "bing_grounding", + "mode": "search", + "metadata": { + "notes": "Grounding with Bing Search (G1 SKU): $35 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment." + } + }, "tinyfish/search": { "input_cost_per_query": 0.0, "litellm_provider": "tinyfish", diff --git a/tests/search_tests/test_bing_grounding_search.py b/tests/search_tests/test_bing_grounding_search.py new file mode 100644 index 00000000000..00e5e382eef --- /dev/null +++ b/tests/search_tests/test_bing_grounding_search.py @@ -0,0 +1,187 @@ +""" +Tests for the Grounding with Bing Search (Microsoft Foundry) integration. +""" + +import json +from unittest.mock import AsyncMock, Mock, patch + +import pytest + +import litellm +from tests.search_tests.base_search_unit_tests import BaseSearchTest + +PROJECT_ENDPOINT = "https://acct.services.ai.azure.com/api/projects/proj" + +_ANSWER_TEXT = ( + "LiteLLM is an open source LLM gateway ([github.com](https://github.com/BerriAI/litellm))\n" + "The docs live on docs.litellm.ai ([docs.litellm.ai](https://docs.litellm.ai/))" +) + + +def _annotation(marker: str, url: str, title: str) -> dict: + start = _ANSWER_TEXT.index(marker) + return { + "type": "url_citation", + "url": url, + "title": title, + "start_index": start, + "end_index": start + len(marker), + } + + +MOCK_BING_GROUNDING_RESPONSE = { + "id": "resp_mock", + "object": "response", + "status": "completed", + "model": "gpt-4.1", + "output": [ + {"type": "web_search_call", "status": "completed"}, + { + "type": "message", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": _ANSWER_TEXT, + "annotations": [ + _annotation( + "([github.com](https://github.com/BerriAI/litellm))", + "https://github.com/BerriAI/litellm", + "BerriAI/litellm - GitHub", + ), + _annotation( + "([docs.litellm.ai](https://docs.litellm.ai/))", + "https://docs.litellm.ai/", + "LiteLLM Docs", + ), + ], + } + ], + }, + ], + "usage": {"input_tokens": 100, "output_tokens": 50}, +} + + +def _mock_response(): + response = Mock() + response.status_code = 200 + response.headers = {} + response.content = json.dumps(MOCK_BING_GROUNDING_RESPONSE).encode() + return response + + +@pytest.mark.skip(reason="Local only tested search providers") +class TestBingGroundingSearch(BaseSearchTest): + """ + E2E tests for Grounding with Bing Search that make real API calls. + Inherits from BaseSearchTest to run standard search tests. + """ + + def get_search_provider(self) -> str: + return "bing_grounding" + + +class TestBingGroundingSearchTransformation: + """ + Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked. + Transformation details are unit-tested in tests/test_litellm/llms/azure/search/. + """ + + @pytest.fixture(autouse=True) + def _server_env(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_PROJECT_ENDPOINT", PROJECT_ENDPOINT) + monkeypatch.setenv("BING_GROUNDING_MODEL", "gpt-4.1") + monkeypatch.setenv("BING_GROUNDING_TOKEN", "test-entra-token") + monkeypatch.delenv("BING_GROUNDING_CONNECTION_ID", raising=False) + + def test_bing_grounding_search_request_and_response(self): + with patch( # test-quality-ok: litellm.search has no client injection seam + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=_mock_response(), + ) as mock_post: + response = litellm.search( + query="what is litellm", + search_provider="bing_grounding", + max_results=5, + country="us", + ) + + assert mock_post.called + call_kwargs = mock_post.call_args.kwargs + assert call_kwargs["url"] == f"{PROJECT_ENDPOINT}/openai/v1/responses" + assert call_kwargs["headers"]["Authorization"] == "Bearer test-entra-token" + + request_body = call_kwargs["json"] + assert request_body["model"] == "gpt-4.1" + assert request_body["input"] == "what is litellm" + assert request_body["tools"] == [ + {"type": "web_search", "user_location": {"type": "approximate", "country": "US"}} + ] + + assert response.object == "search" + assert len(response.results) == 2 + assert response.results[0].url == "https://github.com/BerriAI/litellm" + assert response.results[0].title == "BerriAI/litellm - GitHub" + assert response.results[0].snippet == "LiteLLM is an open source LLM gateway" + assert response.results[1].url == "https://docs.litellm.ai/" + assert response.results[1].snippet == "The docs live on docs.litellm.ai" + + def test_connection_mode_sends_the_bing_grounding_tool(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv( + "BING_GROUNDING_CONNECTION_ID", + "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.CognitiveServices" + "/accounts/acct/projects/proj/connections/bing-conn", + ) + with patch( # test-quality-ok: litellm.search has no client injection seam + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=_mock_response(), + ) as mock_post: + litellm.search( + query="what is litellm", + search_provider="bing_grounding", + max_results=3, + ) + + request_body = mock_post.call_args.kwargs["json"] + assert request_body["tools"] == [ + { + "type": "bing_grounding", + "bing_grounding": { + "search_configurations": [ + { + "project_connection_id": ( + "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.CognitiveServices" + "/accounts/acct/projects/proj/connections/bing-conn" + ), + "count": 3, + } + ] + }, + } + ] + + @pytest.mark.asyncio + async def test_bing_grounding_asearch(self): + with patch( # test-quality-ok: litellm.asearch has no client injection seam + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new=AsyncMock(return_value=_mock_response()), + ) as mock_post: + response = await litellm.asearch( + query="what is litellm", + search_provider="bing_grounding", + ) + + assert mock_post.call_args.kwargs["json"]["tools"] == [{"type": "web_search"}] + assert len(response.results) == 2 + + def test_bing_grounding_search_tracks_cost(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + with patch( # test-quality-ok: litellm.search has no client injection seam + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=_mock_response(), + ): + response = litellm.search(query="pricing check", search_provider="bing_grounding") + + assert response._hidden_params["response_cost"] == pytest.approx(0.035) diff --git a/tests/test_litellm/llms/azure/search/foundry_responses_web_search_fixture.json b/tests/test_litellm/llms/azure/search/foundry_responses_web_search_fixture.json new file mode 100644 index 00000000000..2dba691d038 --- /dev/null +++ b/tests/test_litellm/llms/azure/search/foundry_responses_web_search_fixture.json @@ -0,0 +1,78 @@ +{ + "id": "resp_04bcf916b629f9be006a8c7c053e9881959a4efe7118cb3a34", + "object": "response", + "status": "completed", + "model": "gpt-4.1", + "output": [ + { + "type": "web_search_call", + "id": "ws_04bcf916b629f9be006a8c7c080e108195b3366ab048231129", + "response_id": "resp_04bcf916b629f9be006a8c7c053e9881959a4efe7118cb3a34", + "status": "completed", + "action": { + "type": "search", + "queries": [ + "latest LiteLLM release GitHub LiteLLM", + "LiteLLM github releases", + "LiteLLM repo LiteLLM" + ], + "query": "latest LiteLLM release GitHub" + } + }, + { + "type": "message", + "id": "msg_04bcf916b629f9be006a8c7c0c7b448195979bfafb3b28f1ba", + "response_id": "resp_04bcf916b629f9be006a8c7c053e9881959a4efe7118cb3a34", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "You asked: “What is the latest LiteLLM release on GitHub? Answer with the version number.”\n\nI verified the most recent release information from the official GitHub repository for LiteLLM (BerriAI/litellm) using multiple sources, including GitHub itself and secondary trackers:\n\n- On the GitHub **Releases** page for BerriAI/litellm, the topmost entry is **v1.99.0‑rc.1**, marked as a pre‑release, published “yesterday” (relative to today, August 24, 2026) ([github.com](https://github.com/BerriAI/litellm/releases)). This indicates that version **v1.99.0‑rc.1** is the most recent tag available.\n- An external release‑tracking site (ReleaseAlert) confirms: **Latest version: v1.99.0‑rc.1**, last published August 22, 2026 ([releasealert.dev](https://releasealert.dev/github/BerriAI/litellm)).\n- The GitHub API (via `releases/latest`) currently points to **v1.98.0** as the latest **stable** release, with published date August 23, 2026 ([api.github.com](https://api.github.com/repos/BerriAI/litellm/releases/latest)).\n\nTo summarize:\n\n- The absolute **latest** release tag on GitHub is **v1.99.0‑rc.1** (release candidate), published recently (August 22, 2026) ([github.com](https://github.com/BerriAI/litellm/releases)).\n- The most recent **stable** release is **v1.98.0**, published August 23, 2026 ([api.github.com](https://api.github.com/repos/BerriAI/litellm/releases/latest)).\n\nSince you asked for the “latest LiteLLM release on GitHub,” without specifying stable vs. pre‑release, the correct answer is:\n\n**v1.99.0‑rc.1**\n\nLet me know if you'd like details on what's new in that release, or if you'd prefer the latest stable version.", + "annotations": [ + { + "type": "url_citation", + "url": "https://github.com/BerriAI/litellm/releases", + "start_index": 456, + "end_index": 515, + "title": "Releases · BerriAI/litellm - GitHub" + }, + { + "type": "url_citation", + "url": "https://releasealert.dev/github/BerriAI/litellm", + "start_index": 722, + "end_index": 791, + "title": "BerriAI/litellm on GitHub | Release Alert" + }, + { + "type": "url_citation", + "url": "https://api.github.com/repos/BerriAI/litellm/releases/latest", + "start_index": 936, + "end_index": 1016, + "title": "api.github.com" + }, + { + "type": "url_citation", + "url": "https://github.com/BerriAI/litellm/releases", + "start_index": 1160, + "end_index": 1219, + "title": "Releases · BerriAI/litellm - GitHub" + }, + { + "type": "url_citation", + "url": "https://api.github.com/repos/BerriAI/litellm/releases/latest", + "start_index": 1300, + "end_index": 1380, + "title": "api.github.com" + } + ], + "logprobs": [] + } + ], + "status": "completed" + } + ], + "usage": { + "input_tokens": 15195, + "output_tokens": 467 + } +} diff --git a/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py b/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py new file mode 100644 index 00000000000..25dfa5fbe2b --- /dev/null +++ b/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py @@ -0,0 +1,311 @@ +import json +from pathlib import Path +from unittest.mock import Mock + +import pytest + +from litellm.llms.azure.search.transformation import BingGroundingSearchConfig + +REAL_FIXTURE = json.loads((Path(__file__).parent / "foundry_responses_web_search_fixture.json").read_text()) + +RESPONSES_URL = "https://acct.services.ai.azure.com/api/projects/proj/openai/v1/responses" + + +@pytest.fixture(autouse=True) +def _clean_env(monkeypatch: pytest.MonkeyPatch): + for var in ( + "BING_GROUNDING_PROJECT_ENDPOINT", + "BING_GROUNDING_MODEL", + "BING_GROUNDING_CONNECTION_ID", + "BING_GROUNDING_TOKEN", + ): + monkeypatch.delenv(var, raising=False) + + +def _config(entra_token_minter=None) -> BingGroundingSearchConfig: + return BingGroundingSearchConfig(entra_token_minter=entra_token_minter) + + +def _resp(payload, status_code: int = 200): + r = Mock() + r.status_code = status_code + r.headers = {} + r.content = (payload if isinstance(payload, str) else json.dumps(payload)).encode() + return r + + +def _message_response(text: str, annotations: list) -> dict: + return { + "output": [ + {"type": "web_search_call", "status": "completed"}, + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": text, "annotations": annotations}], + }, + ] + } + + +def _citation(url: str, title: str, start: int, end: int) -> dict: + return {"type": "url_citation", "url": url, "title": title, "start_index": start, "end_index": end} + + +def test_ui_friendly_name(): + assert _config().ui_friendly_name() == "Grounding with Bing Search" + + +def test_validate_environment_with_explicit_key(): + headers = _config().validate_environment({}, api_key="explicit-token") + assert headers["Authorization"] == "Bearer explicit-token" + assert headers["Content-Type"] == "application/json" + + +def test_validate_environment_reads_env_token(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_TOKEN", "env-token") + assert _config().validate_environment({})["Authorization"] == "Bearer env-token" + + +def test_validate_environment_falls_back_to_entra_minter(): + headers = _config(entra_token_minter=lambda: "entra-token").validate_environment({}) + assert headers["Authorization"] == "Bearer entra-token" + + +def test_validate_environment_api_key_beats_env_token(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_TOKEN", "env-token") + minter = Mock(return_value="entra-token") + headers = _config(entra_token_minter=minter).validate_environment({}, api_key="explicit-token") + assert headers["Authorization"] == "Bearer explicit-token" + minter.assert_not_called() + + +def test_validate_environment_env_token_beats_entra_minter(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_TOKEN", "env-token") + minter = Mock(return_value="entra-token") + assert _config(entra_token_minter=minter).validate_environment({})["Authorization"] == "Bearer env-token" + minter.assert_not_called() + + +def test_validate_environment_refuses_entra_token_for_caller_api_base(): + minter = Mock(return_value="entra-token") + with pytest.raises(ValueError, match="Refusing to send the server-configured"): + _config(entra_token_minter=minter).validate_environment({}, api_base="https://attacker.example.com") + minter.assert_not_called() + + +def test_validate_environment_entra_minter_failure_names_the_options(): + def failing_minter() -> str: + raise RuntimeError("no az login") + + with pytest.raises(ValueError, match="no credential available") as excinfo: + _config(entra_token_minter=failing_minter).validate_environment({}) + message = str(excinfo.value) + assert "BING_GROUNDING_TOKEN" in message + assert "https://ai.azure.com/.default" in message + assert "no az login" in message + + +def test_validate_environment_does_not_mutate_and_is_idempotent(): + config = _config() + caller_headers = {"X-Custom": "keep-me"} + + once = config.validate_environment(caller_headers, api_key="k") + twice = config.validate_environment(once, api_key="k") + + assert caller_headers == {"X-Custom": "keep-me"} + assert once == twice + assert once["X-Custom"] == "keep-me" + + +def test_get_complete_url_from_api_base(): + url = _config().get_complete_url("https://acct.services.ai.azure.com/api/projects/proj", {}) + assert url == RESPONSES_URL + + +def test_get_complete_url_reads_env_endpoint(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_PROJECT_ENDPOINT", "https://acct.services.ai.azure.com/api/projects/proj/") + assert _config().get_complete_url(None, {}) == RESPONSES_URL + + +def test_get_complete_url_missing_endpoint_raises(): + with pytest.raises(ValueError, match="BING_GROUNDING_PROJECT_ENDPOINT"): + _config().get_complete_url(None, {}) + + +@pytest.mark.parametrize( + "api_base", + [ + "https://acct.services.ai.azure.com/api/projects/proj", + "https://acct.services.ai.azure.com/api/projects/proj/", + "https://acct.services.ai.azure.com/api/projects/proj/openai/v1/responses", + "https://acct.services.ai.azure.com/api/projects/proj/openai/v1/responses/", + ], +) +def test_get_complete_url_appends_responses_path_exactly_once(api_base: str): + assert _config().get_complete_url(api_base, {}) == RESPONSES_URL + + +def test_transform_search_request_missing_model_raises(): + with pytest.raises(ValueError, match="BING_GROUNDING_MODEL"): + _config().transform_search_request("q", {}) + + +def test_transform_search_request_web_search_mode_exact_body(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_MODEL", "gpt-4.1") + body = _config().transform_search_request("latest AI developments", {"max_results": 5}) + assert body == { + "model": "gpt-4.1", + "input": "latest AI developments", + "tools": [{"type": "web_search"}], + } + + +def test_transform_search_request_web_search_mode_maps_country(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_MODEL", "gpt-4.1") + body = _config().transform_search_request("q", {"country": "us"}) + assert body["tools"] == [{"type": "web_search", "user_location": {"type": "approximate", "country": "US"}}] + + +def test_transform_search_request_connection_mode_exact_body(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_MODEL", "gpt-4.1") + monkeypatch.setenv("BING_GROUNDING_CONNECTION_ID", "conn-id") + body = _config().transform_search_request("q", {"max_results": 5}) + assert body == { + "model": "gpt-4.1", + "input": "q", + "tools": [ + { + "type": "bing_grounding", + "bing_grounding": {"search_configurations": [{"project_connection_id": "conn-id", "count": 5}]}, + } + ], + } + + +def test_transform_search_request_connection_mode_omits_count_without_max_results( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setenv("BING_GROUNDING_MODEL", "gpt-4.1") + monkeypatch.setenv("BING_GROUNDING_CONNECTION_ID", "conn-id") + body = _config().transform_search_request("q", {}) + assert body["tools"][0]["bing_grounding"]["search_configurations"] == [{"project_connection_id": "conn-id"}] + + +def test_transform_search_request_joins_list_query(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_MODEL", "gpt-4.1") + assert _config().transform_search_request(["foo", "bar"], {})["input"] == "foo bar" + + +def test_transform_search_response_real_fixture_dedupes_and_preserves_order(): + resp = _config().transform_search_response(_resp(REAL_FIXTURE), logging_obj=Mock()) + + assert resp.object == "search" + assert [r.url for r in resp.results] == [ + "https://github.com/BerriAI/litellm/releases", + "https://releasealert.dev/github/BerriAI/litellm", + "https://api.github.com/repos/BerriAI/litellm/releases/latest", + ] + assert resp.results[0].title == "Releases · BerriAI/litellm - GitHub" + assert resp.results[1].title == "BerriAI/litellm on GitHub | Release Alert" + + +def test_transform_search_response_real_fixture_snippets_are_the_cited_claims(): + resp = _config().transform_search_response(_resp(REAL_FIXTURE), logging_obj=Mock()) + + assert resp.results[0].snippet.startswith("- On the GitHub **Releases** page for BerriAI/litellm") + assert resp.results[1].snippet.startswith("- An external release") + assert resp.results[2].snippet.startswith("- The GitHub API (via `releases/latest`)") + for result in resp.results: + assert "url_citation" not in result.snippet + assert not result.snippet.startswith("([") + + +def test_transform_search_response_snippet_falls_back_to_text_head_for_leading_citation(): + text = "([example.com](https://example.com)) trailing prose" + payload = _message_response(text, [_citation("https://example.com", "Example", 0, 36)]) + resp = _config().transform_search_response(_resp(payload), logging_obj=Mock()) + assert resp.results[0].snippet == text + + +def test_transform_search_response_snippet_without_indices_uses_last_line(): + payload = _message_response( + "first line\nthe claim on the last line", + [{"type": "url_citation", "url": "https://example.com", "title": "Example"}], + ) + resp = _config().transform_search_response(_resp(payload), logging_obj=Mock()) + assert resp.results[0].snippet == "the claim on the last line" + + +def test_transform_search_response_ignores_non_citation_annotations(): + payload = _message_response("text", [{"type": "file_citation", "url": "https://example.com"}]) + assert _config().transform_search_response(_resp(payload), logging_obj=Mock()).results == [] + + +def test_transform_search_response_ignores_citation_without_url(): + payload = _message_response("text", [{"type": "url_citation", "title": "no url"}]) + assert _config().transform_search_response(_resp(payload), logging_obj=Mock()).results == [] + + +def test_transform_search_response_no_message_output(): + payload = {"output": [{"type": "web_search_call", "status": "completed"}]} + assert _config().transform_search_response(_resp(payload), logging_obj=Mock()).results == [] + + +@pytest.mark.parametrize( + "body", + [ + "502 Bad Gateway", + '{"output": "garbage"}', + '{"output": null}', + "{}", + ], +) +def test_transform_search_response_malformed_body_raises_instead_of_reporting_empty(body: str): + with pytest.raises(Exception, match="Grounding with Bing Search"): + _config().transform_search_response(_resp(body, status_code=502), logging_obj=Mock()) + + +def test_get_error_class_attributes_the_provider(): + error = _config().get_error_class(error_message="quota exceeded", status_code=429, headers={}) + assert error.status_code == 429 + assert "Grounding with Bing Search: quota exceeded" in str(error) + assert "learn.microsoft.com" in str(error) + + +def test_get_error_class_unwraps_the_nested_tool_error(): + nested_tool_error = json.dumps( + { + "error": "Tool_User_Error", + "message": ( + "The specified connection ID 'conn-id' in tool config input was not found " + "in the project or account connections." + ), + "code": "invalid_tool_input", + "tool": "bing_grounding", + } + ) + live_400_shape = json.dumps( + { + "error": { + "message": nested_tool_error, + "type": "invalid_request_error", + "param": None, + "code": "tool_user_error", + } + } + ) + error = _config().get_error_class(error_message=live_400_shape, status_code=400, headers={}) + assert ( + "Grounding with Bing Search: The specified connection ID 'conn-id' in tool config input " + "was not found in the project or account connections" in str(error) + ) + assert "Tool_User_Error" not in str(error) + + +def test_get_error_class_unwraps_a_plain_error_envelope(): + error = _config().get_error_class( + error_message='{"error":{"message":"The api key is invalid.","code":"401"}}', + status_code=401, + headers={}, + ) + assert "Grounding with Bing Search: The api key is invalid" in str(error) diff --git a/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py b/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py index e4402bbec49..e6aad7688d1 100644 --- a/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py +++ b/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py @@ -16,6 +16,7 @@ import pytest import litellm from litellm.llms.apiserpent.search.transformation import APISerpentSearchConfig +from litellm.llms.azure.search.transformation import BingGroundingSearchConfig from litellm.llms.base_llm.search.transformation import ( BaseSearchConfig, _is_trusted_search_api_base, @@ -59,6 +60,7 @@ _BASE_ENV_VARS = ( "TINYFISH_API_BASE", "CRW_API_BASE", "NIMBLE_API_BASE", + "BING_GROUNDING_PROJECT_ENDPOINT", ) @@ -99,6 +101,7 @@ PROVIDERS: Tuple[ProviderSpec, ...] = ( (TinyfishSearchConfig, {"TINYFISH_API_KEY": "srv"}, "caller-key", {}), (FastCRWSearchConfig, {"CRW_API_KEY": "srv"}, "caller-key", {}), (NimbleSearchConfig, {"NIMBLE_API_KEY": "srv"}, "caller-key", {}), + (BingGroundingSearchConfig, {"BING_GROUNDING_TOKEN": "srv"}, "caller-key", {}), ) _IDS = tuple(spec[0].__name__ for spec in PROVIDERS) diff --git a/ui/litellm-dashboard/public/assets/logos/bing.png b/ui/litellm-dashboard/public/assets/logos/bing.png new file mode 100644 index 0000000000000000000000000000000000000000..ab1f4359281421b7500b7f1d77109f69ce99a95d GIT binary patch literal 31955 zcmYIQbyQnVum*y=OL2FnxCeKNySrO)cef(Ntw3>y;_gmycXuhy3%~Q;d4J^O-Xwc> zcXoE>n{Q?lsiYu@1pfsd3=9lOT1xC27#KL{5*!Q`8uX^?Tx$M#=kiTb6s&5J;23nn zZKf%0E-w#82fBs@0}rJul9YlqkV2s7(6j8~q^#;sdE3SF+v4+O{pP=y&Z?85&J#g3 z3ytjh1^run?`q#k?~QG~bEyW--KO1{p;iCRQ^S|o?D(ds3*QRDs2N?K1FQ3|I9W6A z)4wbl9olcs9@bnGa<5_2rqd8@4WJb=Z=tV(n=c(UJ^S99KCa2n8|Ji4=mBx(F-KkZ zQT3}r9`Aq7C#JY5HwXEP}uQmZ@A&}^9^wad7j?>F_UDDAI7 z4^b_Ankz|xlmKNimv+G8A=jnNrpsLC{F4s-qyDR#~8iA+ zwQVY>N%I7J31h9-Yr`7OP~niS-5kJqBe7rU+s)h9diQT7u##t7i$?I6WO3HHx>7FD zoHJ>ve~~_sB>W1ts4NInL^lPE-!B7)R=Iq9G+u0SSJuA|W%tacB4$ewon@dbw0Tyb zryZE`4;>uW$)G2{?&kL|1{zN_074)7&I(eD*jtp%)@tA5aAMmo39YMc&^ZiF(sO z)1=HU`TaL|I^@_ovgOSmbOvl?;S7JzI2r1q-3(@1(0EE0IyC)_T` zyyhd1P^ZJ;fIUZ*#7y%n8#9dsg}<}_op@!^hTJH0)dREo;m&TLy3*Sts?T=C0n%&E zfY&?)>^!HuE2%0{Pkm;~DD?)qt@aud`Vco<7zv}awFPF@QOV;=|Fi1{6v6s7UwM9t zcDGHFkeUWcfZ_|{xD6n}wUwk^kXCiRJPLnJ!*Q>h(P8b?ZCFl(Z^3-ojgL4A1*ak2tX`Yi5&!x`TJhNF07B}n-s+!K!d%_}T5 zl?7?>vF@-}X3SgU*N$5bKF%61?8LsM0_{GvQh)Rn6sJB1Ff22%Q$nEjAP=^1&tS46xFMSa@Zzpcg z3H;Ini(kxgM2~iA*aq0;#GzbozmOy4-WZ z4jJXmosR!JKS)Mwx5`|%)kOMqG*%)rZz z?iFSMS0EZ!?oLYFx$;$Tn7k9uUjF25Cli6^_y3v4;?5H8)0ntJ|46FN>+9 zgb8rF3YJw)H_Whdd?JU|ii$)lwVO?hIpeBOQ8e3^+`#;6jXWd{I~0W#b*fs(G4E|4 z{+!@n`dwRA8e&%Wryw1biDZKq5dNac5XQAU*`N|Pp&d2o1t`zj^1}6~e6Oi3-n3KW zx9mKv@;!eMdVx&>nZ&CfriS|X!bK6i+j)72YP!na7!Bhj4g#))(#?5xaS0g;FQi4S z2O-fu$K6~2H(6qqE=um6iL-bPN86dsL-w|nYGxtfj1kRAYJNcd3K-3>et>e$Q$}mm z@CvGEK~>EtMN8)OjwGHecB#L}=u)0pEVyFAg|y4tS}rmH^#} zgg7Tdg7!Jf^h;hUyLykyjStP2w8hKRk^KhKEiv=Y+07He6?9tC?TGPJwkpEXh-Dxo z8Dmw#FL@lMRs7g-RsDEhCF|yLX;E?RW~JB?NOLhl{2SA%HQpN%vGBExCqS#t|4C1`y$>$n|U|8hOm!R?!Bu4hp@Q6Kt z#2le`&-I0HdX-yhf z?aIf{6)eCR&2h>HSWC&`uEW;!8Ya?_veDS|7DgYl8tMLS0TN4hLhSsF6Q2Cl2=-a4 z04K~0jW$k|QGcz7Umn`Q%p>E+f+s&s49+__QEOMD$uijR_H2_7z{0*OdiO9;$3f4x z&s0}q7tyS>OaV8m-FM1{u0T2I+sN$SbD#Ac45L)O1>IH=e55bT9&+zym9JQWgM7SO z`(VX^xw?L$F|=AMaKqha7hp8hX{Dr|Z1#Sg`Xv#Cdz?H8^>QbzceH(Q%^iu*ITNTK zR#55OoO|LA-zfLyd~P6`eFl><7Z|;X>|WBU;WgrLYO+CC2s+QX%$O6IzzOF?3UFv` z155NF4sbq#*JUA|q(oBtN(-qDp9Dw1y>oy8^H@8$yTo?!$k@{{X7HtKoY@fnzui+# z&GklrnP#-3%vjh7>hBMFhwKlEWEys(z60krTmph*-@Wx19M?19 zm5h${@SW`~5aO|Qfj(#i7h`)QRmwP_r0Z-|as~|RiLwFMUlbjS5o^zaxN@KE)q~Wtyhl)$ z4jraCo8zbpwI^kv6-@sFAa!r>6Q@q_ot=(_NFoh#x$uJ1ouo=@wh+mobq7J4EualsKAEB1oBNBF6Li4&6 z`B?Ecf7-yxeq-;duER+~;3LBa(LnN8%~E|!spCSmQ&J@9vKOjkIHUxuW_|*%)*wYV ztv~8S+xd7{IVe*!Dh4St)9NH|Sd5y>G=?=5mvM@-o!&Hok2XT+*m4*Pzbpj$64htv zNw{K2cbx0V8eCHm2@XFkxd(c1$rkN&NSYyX=Bh;0L}h9n3_p@M5Pr2}Z>$$yTVF_;dPV1gO zkgj8^VYOU9u(^r)Q+hkt9bDVWjKzOE7)3*$N4#l%cfuQ+^fCThS@nA74j_dNF^OpU z=|3|tK?s3&DHUfH?GDblo;&*YQ;{_O(|YW|=fWo44}AKp9fVNyt4>3Yza@`GYwq9_ zn~x`ze=3y^ZEpJEk@n>?L=m1jh%TsR2dQ=tKn4sn0d;)g1((zx*ZRWOV1g&q`YskZ zF#sz|X~_--%b6Eswiyazuka)VVBn}9?)d{y;Q_{TQ+1d~ZXK$$1PN*!KgunTqwkFI zQyG&3wd}x(x|{b-p|YVGw7E^t@VNN3dAZtl+L}}}h$5x*R==)sbTVa!b~z6iY%0%= zIP0_3fYx&!nx$hvV%aQGo@@C@A{&8?sWP#JW~gkX5ZhGdpU!u9txRFXE*eTob4`D; z7f(LspY}Gb0~!rznNR(HjH`DpE^tgox;Z~56Hxs=+UF(yZ#EQ-k-p^;#?>lJXItmG zF>$Flg?9Ga6`+^MWLm!d5!S?;c0LNnoEXSOf8wmoR-^Zy@uP^*7(~;!E)}ILrmu)% zpjU1DauhWxqrT{lxhnH^CqOeiZ8_~RvZ~zE-WOj$ywrKBSbq8rV}Q_VYM9C9i0+H& z!63Y8O?_9{ZZx0AydhKRFc28|@{rY=Rkd&=BQzZL)*EZtI0Mt-uZVht0)tmjH{*|| z-JKL@H{RWT0k&f;gEptC8rU|?c`%$6LigUXr%n8nug_H$%_jZ$6r4}I1#WfvP-zwM z!<0f`J8sB1jIt>JB>Rj{+5hbv%czKHsj3l4uPXbIG^zln+B2BurdvZ+l`Y{th_uhg z+xQV;zNFE3@}ri2JFMYL_4${o)!f$#s6Vc0hy+83p$cFBY*p&m;xr?;0(IAd68TuiOsLI@#8 zbC&3LwKOzAb|udLKLdHyDO%=N5K#Q=bWbt9a!8jgy-G6W?Hcjy%eVWRN#NA5r{(Qw zvRjF5H}1~#90Pv0XZkbv&HFTEB(zM^z2%+3<5*s;WrVZEpVVo@p+6VR^R#vK z?{9T=*3m2Q#U|7czxYJCeLKA9&4jq?e@jqSe1#EYm7B5qfiuS{W}h8%7mzv-b`eTNi@#MyXCm&&bK zTvh&GQ8bZ$;<&JKvd~mDT`SgVk~HfP)5yc+Rq@h1rnhhZ=813_FR|gHP=5(7$*%SU z6gNRHzW&pc*|;Qmj04ArV?FAXHD-uO$$0;rMrD0!W@Z6Bw z;JD^7LXko+Z2>MlqLUQeqm%rqWV=&W_VJw4+-K+ji6r$?!CIxO4gM-SL0!-vg)Od? zxkr^bmp~xsig!p0F7C}E)k+KaZcEz9Z@2=_QD9SB5P6_yqK&-C8696j+dt)+^7#;x z>(*&R>FJ1SoU1U%Dg>qYVk7!@vqh|SbD`mHtgEfU$JOblQy8aM5+|tyH=U0FtKH6o zAN%d5TaOn1AwX~eZMvJno@>qO3tWmzxRr{-BVcQVSm_eFmxDC!_XOm{T z%J!pDWX2V^9B%NCJW2IqqNeo|!-D7o5IjPRfdVG-GAfsp87tvyc&2wQB4U8mneP2aVQ8;28{;CJ7&ety1qY1`#? zx$%AIBKX`9m6`wL?x7%h%Bs5Uru|`FL+|g#bCB%Hufr^-_BE#OIbolZ?54i=aVk>! z^R3PBZF<~MV38L8+00OHvuXS?t|`I5zarOIXNcNcReza&t7>o?7n84Ub=uzR(eP#L zvZK*=Z^PFrTeuhJ;(`C!RFSh+{9O#B4O3uJ$WTROUkp*%9KtBMM$0*ZQP$4#5S_J5=pBE`9sRpN zs?Ds+T>auXJX=?UZmpxtz>~^2lw3Bn|2mvMw36%fkI`-@u#dHMRD;sE{q(wxq$a^d zM(BR4^MUcQt+L~_q}u9Zdk&=m==bQUTc8ueul5-U3D5LfpGP-T>#CpD!2KX-U=o_9 zDCdn&CV94jY03u=l^jAfz>kZ{Ulkc_KvNWIWrt-`;;!1s#XpF*czOTu^>6cB6#2=| zShdAjJyEh|?ZQ}s>3=Y3C43~+<~x=DU`O0p-*HwVwEjN#m2ahROu~`B_dKKJ0>!JR z(c0+8)ulk*n`~IGpf{d;Cir{7WvpAj3weU&s6<}GX5&A{mrUaJ!j9+j#irc%GoU~w z?Vqf}q`+l+uvXE*-w(mV?mJr2RV2VlhY4xs;T8T?a4|Ev;u{J#R_A0mVx8mDeQ-`- zR&e?6ZS&^GvaTo*4-XDRj*KkoXU`me#+=8?xHK@16bGIYqN5=Z11=s~$}gD^yjQJ0 zTg!0`nZ%LdQC&_pFLmy^iyPjD_Uzc)$%y0jfYQbl1gDcjB-afk;4e)K>*doZA_Sg1 z!KP@jtzavz2MNG?gRCkeAP6 z5cFSVuOWcXb3TpKV2tNv;9-hh$VI`etUa9rcSz{>{^ff!d431n%QhSL7r~8l&oc5vqJ$~L@J$BvIXk3Z)x}fWvgb({>q211~CPbnT zJR#WwdQi0k>U5mO+^U}ugGshTqBy>5f;9$F!(%SK#13k^6Y3vRqtH%Ch)Dy>4+k#6 zzxr=3iw*t{*WN1|efM0e=ca**ng79tg+cmH%Ee?~$DTeyZUlWOU^x-r=>j6*r==&15_XL?cnRwhoK?i*Tc_gM zs{XZ5m6g!-GR#Z8<}pJ_iTEejbAMl0zK#?$R@VA9@ruHY-NjpAf@7s^iHj5> zPm&Tu>5G{l!}C*OLvlKQz5a6Mp$TPuq4Nx<%8t0j0wpZ%fdiC#;0)+&A`9HIADkdI zusD6-jap0)RcFK@aAFmM9SaMzZ^-2kyOGL##P#JP3NbCjjL8%_p;16NR zo8qPoVOWj&$`yhI7q{TX#z1&jiql=ox@R8$PJ^S23dL2gX!&@|XgVL@ z6rthRFOZYiWLz@tBe;dY%r6P96jQ`_!Zrk5gb~Yvfx#ZykEd9d&_11<%@MYb`9nVD z*Hq9PXTV2j15hij1N_D!TQ9Lx0C6(J3j{IZJRo63MUAA`yb3^-$Mp!D4|zM6t{_EX53G37w;HSsVBacz*fzJkf3$ zAp(UCsOcggQBY$vqO7J+9UrL+{W~OK4=qCN2Uj^FOs)l(T3I|_+3ei*S;*m4`Gv3jammp6r5h&u$uI-sS_p-hZcv`3*{a}M4(%<<(GI^gwa(vZarIaAa z=EkD~X8mRCFKly_D=Lkr;r^KOZzg)9L#da0?|%(kud|Kk%tLL^&F2D;-d3yWpJ-V| zOx=*_rwMb55u7!THKQ@}gq0DMAj7P71(R)`b`>4b6cX`q${$9-mD6~_m%=V?g=|E_ z`s<~K>i`Fzg(tNV2cRLQ*$cnNZk2=W{!sdLGIQ9_3!*;@9VicyYT7Y8vcDJ`*IGxqn!+ib4TAng(H8 zOzMn#3C&DOFg9oeG@)=Saiq4p-Du&>%YjjM!UcHEXI?%_s-!>O{_^@A3O%kNVDXwv zpqsRBz@ne{kyws~D`>^R>*GVfMW+vU?OCfr6-5yF&T!C-7}H{eSJ$ZsLgJR$rrIA7 zM>w#)cYpoYnAUO2er_-X=<^+;h98-AvIE5d4NaQGo}p~xy^&_WVo7lb`LV5K6kD*X z_sHT2g9_4=T9!nPCb(IAj-zrN`_06{l==+=zrMtE4B;3@;WgS>8oY$!gwGQIvHf!0%v30uiBf=smg+~I!T3{*^uOyV8V1{BioCtlSL%A(& z9UK?1?{97>M`=-PN&Bx)=AYp*bh&7mIR0bJ;4|-E-m;)msA@L?1Vru~os~= zdI?r?d^pLq*BCX6Y}Arv0ZwuH4*`XLDz7csp43~*hq+kW1407^8gH;0`1;togTp`b zy=r=G@{<1LE#uA_QUOu~hoy-tsFB8g#Vg{xx3ESqBHsq|ke_!xel6^}o#6xR?)#q) zxeh)oe#0E(KhF2>cs}Xx?nnLFL!wf01;QXUWJyR2rD5!%qijju< zihJF~<#-Qd@?f$=Vd6?}=fxNVQ->c}+k&e1RpEPA6^n#!neuEvyAL@c%R+zLN2&}0la;1X6zu#wNa zdR2PfodO*0e1A2!zYY;X`>x)}qIqy#@+`gIh1**hc;6$PXXwG0n(+wy`Ghx0PHHm& z^{AYDniyqZ2+jj#gFefDOTl3SumUHaf-?jqVvemg0fHzQ*daHeDK9G>P z49h!*D4Wmi(V0Y};p8~@7i+UNsIZXGDO?Sa%m7#IN}>?`$07CnDmYz|M+0a)X5z*w&O2K)SL@s*eRP zW`l5(N>5&~bm^Z+)8jQ0+WlRwK5d*$CqV+^92P05$Z?_CFLKh-?BKJ_KI%6}rGY2) z_@#O@wC!X#x-9;cThMYjl+JJXxA0goSg7`MDjSFt4*h90>_n{jgMbQHtrR5e&ERNrOQmg@|5vwXamd;GEM1yJQ zb;?AS2m_mgh0j26KQ2KQkI(T&WOi194t$t@E98Ma>dAdR=Hg0hu=360GIgzu+=ySh z687%%uCmMZ^LWx^ovTsHT!ik9y6K5%=ykSalE5K|h9D?~ViUF_nlT7n3g%pbsvZTsT`>AsC3WT#?f_D_wHk#(y+z*1`UFvl)`~DE$u|) z(h$#EKzQwabp{2cIU*K21&&xATTIGXa;tsgAL2~6tRtdhMBgfyra36JK+64r*t3l%!$cFN$%4$E3Pkv3-)H^ADW`*_ zF+44?5<$&1S)ugBv0t3P@`70M(Zayl3KCcOniw~>wj~l$OY+7=7|Rwm0)Fo+pTKcu zoSsgmp|wR|?@;q=p#u8!3^+!|4|xPkkJ04^Q;Af?F7;K!{6ds+y<%N$TyzQBu@@k zIPE)4z)u|4R>3~bqM7Wj>5Wu~RDH+(6vrDd|1-;$VqFsNJcNfH!h`#tDWk zWk3;C-!@u7p_&E|rO63)pgF*WAyu(-&CP-5qI~+lr8(5zpG#DwMINnv5p3yL7m=6_N^#heR2#lVRWK=`nxUqL}?4 zV|wsljZEp4Af4xH7Fg}*`AL^Ru8zy+hwJS7!2GD`idxH3 znXbmeoW{t%iJs@6j9(c)>iDMF7xG{#dl$32s~*F)7k#$|AGQfg9;SXwSpy9 zt|5`nqKmQP_Yr4Uk=M=)#fuRI9Wg=WAo)mjO`#EYA_Y$;$OPAWgFMFH0E4xHIZVy5 z!_DnyQEI2E5#-YsFbN_R7X(E=eDU%OI9}kR zoK^R#Q1lpaH%KVzs}wWPLiFGN@jrNYpN8Ap5_=x~?Z5!>R|%n$WYJnNYNqM*QIYeQ z>=PsO5m|15+T%U!Vx6bPz^3FuuNmzxmRG-YKuC-F$U4k52C(aR*8Q5s>h5o4q%&pk zwidl|%fwGlu={lvy`suEWaz`gH@k3D^sW3a-(xkAl6+R zJXZXdF+I<&wxkFk?or(bo7`FZ?-HG^R2HUePbIvjpG*(8lE-rM2U5*l6UNo}8FRM9>u|vq57ev7o$7s!u500=Tw0lv=@h_>hWE|Zm zyGnXDWd`lqy+pXPH-WGWLQEGpSe#JbO!Iy&*<+f%m)1gP6gfmRok}3mHBI~%6R5Zf z1ZL{|)893TDRNl}(SjN4Rh|XE7^Uj;Z;ti{jU=uM%flchSiFy`uc{n2U?n@cFPHpF z-)cR09Lsz&9YOhh^V=c*gT%$X_AfBvHY`R}og|i65F!P5&{0U*D-hB2$RsJ~0b>%w zJ&2!4=&__VIc)J{dDj6s+!h}!O#44}Va%;ZSr2Mc+u@?tEZ`|4MoOYKc}#Rst#er9 zzWUZ`3cXRLxOg9KA^TslI8zGpb5M?>;39&N$bE-b8HH} z4ZlAy5_f>6!}b|>mCMdy$O0XXDKUxof$I4SqCHuM_ec;ww+u1X!Ksk(^jqb;=Q8kw z#!b4ymUR9J&34Zt9%GKYx5~-bpFZgr42}hiLgm2vMKA)#2nKi_F>mhVNMLlHlKj@i zxvkRYNbRtkH>IAt6Ms_-3x$1Cgx!W4QCXB+2!Fa;m+;BEw}6<2eskY=b9}*BT(=oX zCYrPXXlBz_^T{V2F1tuy2Lb96sE2nJQhpikXfOU5b&f>{=Nc_z5mlUE;Z36Tkh`(U zukN__7#h$oScQH!^`Y9Dt)IdrgIA5bQ_J$x??_QeA>#>&r%^TC=@N`do&n-O`z8cR z)W-893sUi-wU^H3IXXOri7*8nYc?gn+K!j)9Y$vqx%(Hnp*8hw$yc&!4;=@C@ejL zUI6z7pzP6$wH8;T`3g2m-&bct-3@NAjl|iynRjgy=fA((+z6u(*d+-n`5@6ypU6EW3ZzA?P@+PXK@3!3nPzUz_MPG zye>k6BT7faWtF+g6c^bQ9+vLoh}}6#h&DxNF9G^i-a zJbM6}sa7r>&F)y4 z9K>_jiakIvn|XCcIY7|jiLcM%V~n%1@Gw`WAw7uQK;*D<@DGdN6{{T~%VW=uxl43F zHmYuk9GW_U1y9H`9EBcLk7r;k3+4+FQ)~~6l#v_Wg&4vQEL@+P6RJ&4qaE1bMKcMW z-LNAP`f%r5^}KnKR>CoD9NZzH*E_xQ3jp=yA;7#5oP}Bk+pzL0R`OyJY4Rylu^Pdc zM?lk)Sb{htsDO0A$+pp6=6lS)$oohX9bv+LpP!C;IOh^>EQ)nWmH6}S)(m!M&g)M8!gch(Y`zwvMeoyoU&Y2WFiuLe zb*rpDb$org+3f23n~XeZ4Ip``K!{@)3?a7{q{q>v_voJ31Pdks;m@(nRsp%OCe$m@ z04KvS)Ysv}A^M9U4M&em!qVKM@9#fLnjQR%1QA`g_fpu%$9v5+wYCa)wGdOW#jv!o zcFW{$?D)(QrCQOA2wcA9kWguuGU!n4;c9a@KzGg-E(*!x`WT0pN5t`pUp_073i(g< zov&XzZ<*SH_y$d|3J5F?cA3OOAu1T#vEq4*wb@PB2>{t3ZT{5}Zm*Og-1)o`YsHg_ zt;j0d##hf}4|s)kM908g{ed!<{wnZ6ScpKP!&a7U61+(y3mc?V5 zNrUyaWGc=PGp6L02h~2MG~#T6m>PF;K_<8+51KL^Zyk#3%S5s6Rtim=T}a*Rr9qL) zkVIpFIc}!GzfHf|Zm8DuoIv5*i2ZsFT_N3a%BdjBUZ^|DweB5zM084Zgvw#&3At2%u(xN=Wf$F&DgsL%GMZ zDL94lC1riV73Sy5Fp=$cNr@Y;xb6?j0gnB1R=&=xaxjXyxQIhxkcpfcLn>v|1nrES zRO7uGy2=Dv{90@M(!};n3}_>eJp*uJ<|_o>6aj7!V?o_gfaYK~ue-(S$r}~tn=c+3 z@hDVo7)R(TkO-S>M;>x`_x->Hlb6HXj*@FW2x=8j{FSE_jS*e$NCJalQwkR!M~Wzr z!dyyy+jFKg$`GTCwN#83$K7yaK_eu>DYnM#*uVW7jR6TaycG&V@j=CDp|87svUJ-b ztG{R`f+EmHcMAyOiVQ8=zaIUS&R53HzZYEyZkr8e*834-GRW=R=p`}_foGB+h7Yk(0(6E^TTxe7v5Y_pofYA_|WA? z9EX@7%Z>VKeA$JRFw--J=^VFh)>~J_OC(jN=io&{9hzhW&j45-ux%8h^29AA{UYk>*pnjU~34jzg|-a2mz-KtW1 zHgv9oS?wVq@bWB@#ramqGEKu|(k#yAUzDDgu1NExseGwwUvA7Pd=+S`caYXt3+8u; z9dS4xDv~*8)fbuHCq3Y<$eRTN{v;Pn*Z}gRJ_BS&Cx}p>AB5-6S|r}VG5zsoOW|!m zNJ@EI-W#rH)g}Nej=Lp=(yQO7scHY$?aD3nAAW%WHhG!{sUD4I0;lAw1 zW8i@9nkRvIW9aXxzU#wn9gcl~D5;hY`l`f_!~Iq?KiIY-#9Nn!6IzADgE*U3#O z+`j+bB;==XW@-sC3SlueOa&-^(J``(oGmoQi)`D*Mr^Sek zl}J&kIk?Vd17_t3B1PaMODc|;0@9zD6L_ddWO^HZ6N@6!gN2(t2SmlamSxfW#<84I z`6Z<|iISoq71tf)hmH+7iB z=0P*KYR?fJ21t!*sgNghL>#eaAORT6q>q+Bk0V^Rv7^~p3C6r5a?k;f5zp(ouu$y(r#Nad|I(G$ok={DmD06=k9=G88UM#EBI`jj9}mrVSvOkRI0`1x6P;!7!Cm zS`e=rOfTjF4PEvMcPDy#Oke+SKDz|mh>Rao50UEvBJ(($Ml50y{O=x;cg;YdYe0Ht z1~VfzsDAXIKA}5MTUD<5AXSA40(hM2Q$r%V=-6QApHx`?>vaz(I!xLg+K!Elj#dN@ zcU{t=hj|J4DlUj8Rc`EBi z!4ba_(QtJM*%BPKy*zn*jZxzI>^5)XKB;fw9K{a zm1+`E`PH*5gBbkWFct{&FA;(ECirF&?Xq4y8m|SJso6}=s$^4wmlY532Dc29)8eI3 z5vsqr$+%|#bk)3S{fX>YH88@+6U6!ZdD0XKkJZf>^3im1_Y^HGYFT@gW4}VraP;`^ z_N4@IgY76flFbm>G7Nhf{#1PkU9R)IbaC^yAqkfUqzJPUI9c-U{N0WqCiic^Yrz&1 z*{(f8?2xVyKE@FKio-RHu#g~bzV<3IPnARWcwX(J?hUwv70a?sqVqr^mnqmON*7+& zs33LF!lKwMN~uzcZ}VW1xVuW@oQu4%({>ym6Gi^m(L#}`zyqFqyIcBk;5a#9SiqJ# ztI=YGVRn2tpC^Fn0M~gK=;<$Q35p=1$*4kte||kHk6V5x$g~Z$Z#HO|Smpp_aSXyG z>r3Xtzg4xWeysHwxclqASoD{D{N(D&69DESyRz;TOC=_9;4rmV!3T*sPV;w)NYnx1 zk;16FhReAgyF*%?f`smut;?$Z=fD$VfUzEk%kn-IL#^ zHrOxS0iTSvkE7Z2JSkO=5(YYm*9k7}+KT8~J?qAHUFEW0o*De)h!)TDldruPp7Jd?Vv8I52n@;!IFO~Qm&rn z=o(33;F!?3`$ko11fyZaY}t{XU#6o=VM-}>Qx;!vAYCw`v-r<&`(0)YoZ!Y?h()lT z6I7TD)!pvclHLGTHZ9MG-FYrI-83gVE1%q#@h|F~@zxW$m0r?JJU?e~%u~=Nn z1P;QE88mMXo=oWYkQHB(+nIJ-|6{s*3#_QK% ztoR@Nxs#PZ&xiebKa+LW(xw(dD-&r=s> zd*WAeW}8jz2bCg zm8UHcwn@_7$A!ksJ<0~P$Cwhe73tqJa_-JZ?>FmBI{ks53=SB3k=24KmOj0BJ7{>8 z3J6sC(Nx%q_UOZbX?UoP7+aChdNFV^ws~%ul->+CO?vOfv?IhUPFeiEtEWr?UA}jW z)t3GjPifYW@It-WVfs(Mk7uO13!h{%1?t+hw?;$ezk2>ow=|_rFn!U5BAQgox(Z#@ zxa$jfH5Ln-PH8aIzlQPUC(#a7sK`7CA8+*j*Zi_u-s|!Sj)z(3ZIj6lL|nHb+?RI% zmPxefHwy^HRpoF6EVIsb5esiAjAB;1-XG*SYuxP6+Ogsq1Ko{|6S9#d=Re({-cRuj zru}I3(c3AA(9Hu8w|OlP=k}yDiS^}!xY{cwD4ZvC#h7Nc4fYAM7Px{?19ps4Lg}1G zidmNv$Yy&kuVqNjN3_1g*SA9hAl~W0?t*FjYl1$ue+t5+FGHj9`M#NawcKW*CE(~2 z9-z16>0{y=uAUCcamjJ)r^It z>1diOw^G{Fm<5gH2v(SL&9lnWYHUW`Gb&Y!b5mp`j{)Nk?RJwjn`gy;Ob#gTm1+F- z>sOKn>wafQQx42sM{C12ViT-QG9f2)N0{;Abu5eoK^YmpyHz;HR9)8Zfh_%%%&Y56 z?Kh;*&(g^N%L?$kxEtbiTb&~tAg%t<$D~{r4k(Nch>ZH+O9dEZ&Y*}V(hsVCZ>It+1ZeR0(g*^*^_3m&%T#*U^5dYSK$HMToeOJdEY;NDW zdHC7fUe@uC-e9Hk$uq6-)?+9y zBppxD(QvD7Zp$Qq2-VxnX`tKsmV0K>u*$SxqlNzm?Qr02gXEOYilj#j`16U)a2A|} z!o0JyOA={+#Q;Q^R%-vrNIP2|N-+!*t~Nu+5qyDVd;fXc5C$3fhA%yII+h|=#LN@t z?CrdUeo-*->?^VRx$R|yP{-Rl43Skh?Had_jsjjXzAIaNAT1%S3xzY2ezsICXVP%W z9JF0)X8PR=WaVBX6#$UPtJvx8LVxN0=O*3$!7Qr28jlDf=YRAxf zGd_r9nI99BePAW)bCF{yyeAK3cZ{}I^&Law!gs8mmg zK|btb*5tZ7r3cPa(uyXRQXox;@wIkdk}9%yROHy3iCSm)fWNp!mB*?Uv1#<<-W%n| zme1#LV~$+6f$6iAzUPQu+saDBK z+m!VeU`d}SM{yNX2I~Wb9x7yASq>EYgBuKDVz<&Ryjy&Y|M~4p+{&)}{5|Sv7B?H8 zng+SQ?n9ZVB2^<8LS@hQ)5DUDOvmiE77?qGW&j|JxN;z2?9et&NSu!20Qvn^e3b*|GnNW$4Q46OB!aPUI z1_Je+lHn&=CY~-QpQ*8&!nQ#70f{RFC!ec)f%BQ+@EhVMX%W z=T0jfPa}x4Ct(f|_^~zh;L!h~GlUi%WAGL8L%SxcBli!{mNN!SmRuZ{xlCDq{JYKl zz1!nqe^vG3-ss+s;`sU;9Nz%E4p+DGA}ZgxD<-jK2pKi$P)obVZ;<)AEU#Ux4rvJ8RC)i@2MFDXOdZNf7aYrUMZ>qN*qWa)PUzq; zuzZTv$0|jn4f|2U2rQ-dsO0YYHGE!pr63H*_wu5R2Lt05`1~&b+rlcV61sN&9Q@b$ za;C-W>B8_;zMZNqi%ZW^-x@(*$%Tg6D!H8cUzO^~`K$hk%Ph{Kck3S+@0iBUBMsS3|zr47j7 zP|TrG^q6IX_ZQBn-smhigm*zik4QFw)alr9_&anm8z1JTJNM zta;nhLutyqf9AWX9%baVovr{AP^8hGRIPv)bNb??O)jT3;$?3$_={OTLRN?~v5Hvz zVL&I)d$;ERs_&2}jx(nuL&#Tj_jlYL?khW*;BY%t7~WPM9=!aQWNvKy&?7QVv3U>Y zKIhWtFub2ppwcj>t%&m7STfS@gShQAh+Ow5PF9NPfrmpwaUaLgIdk@pql*pyd`4tf zG_t+*eV(}w;^YGw!3dD|Gyf;~ItIml%#j~18lt8qm71@@Z9{X2Tok_b45c5u9dU;; zr#7!U02OnlmwC_}A7i#PdD2m}N{iimYQ_gBH2ZuQg@uza)&RX=0cpm8>X&W%NO?Ad z;+O7L^5f=EDht>m12^9FNc-{s@_X&h&+cefoquNA{mP5l?w4KI9zV+pk@ZNlYF`tO zeWo<8`w3Od%yIl%#cO`il@a18=TK?e38 zdaV7scU|9p>^(O$-W{fB&)RpmUH-A1?b6@h)pot=;IbuSf2TgFZ35dx_eH&O8 zS~JLG*8)H?Uk+eHRTi9u9Kund#)CJl#PA7Bb@}8zuNOKh0_bP%g&;%DI?Qyv~aWQe`4x9H?$i zoeMgz0G`g2nUv|3h5x?T z095diLqFVg{^6W*JoCgPI$O5!&Uy*BFm>r&YFPaU*Uq&!)gG4blt)?IxoBa|S4V=>WWAT}cAA9f3?ce>< zXWPE`$T0J!GS$RXT2;V&JtK<|nJo6x+S$7gx99%Ky>0i)FKRo!{Id4gS)na{*kK)! z@Um6C`m$1Rg2m?0?+^-SrZua59mZ7C77&1(Z7xs_lhA6!u6oODej{gN~EV(Ie1 zFpb$V-ue0vP)vHws1`}lkVhk1LM^nQ*}v>XuiGn)3agL9*-Qio?!-P;Zqat+Ma zP~L&0HWF37yK$h~grN|Vo2#MN!L0=Z zL9&Pxx6`U`y5S%uREspQT$E|Yv=ms^ICc7@`G5Y-o$X)z{B`Z&M-JvxDARcrrtUJY zQJNv#ulPZa|C(pd6>^fPiO5akw03&vUHRc1?d)BTw)@|7Sw0Yq`KH*lBtky*il+1j zF$~~ivMIsRi_K+NPf~fBpJufID3@fvBg|?i6up3pwY=C*&nuG`qof91$6nqy zmrv8D#F_W>f9O3owf}VW&F%2v)y-MCji>7r9#?Q>Hj&C3DGG*@8ohCIpXLUUnx~Rw zelEUgZ#(V1_qV%a6F6=9aKCJ5{_tCR+B?+;ypnwn-X)Y3+7N~3@CHE9*srOO?F)3*(Gn`~L z$28I1X3C`f#Onf-529>hj4x>TGHG3&#pCy{QTG+b3mw|Y-tK<*K>KI!xUT)$@873H zDd!Y2Nz*AHxORicCa^%}CTEiZh37egmluGPN+ZT&oychFVox~_{YyR-f9UgybK1_= zT$Dv`JIIGzFZgsQJ!o1N!RPh=#b+)r0i5aM4`Y&Q+Ky^FuOu9cY z4x}YNuQo*q0Mii=3wSCVO+ObmK2tsyDE8Cddc2lP-zV4kaLmb1eDL=6Z{B@l+jsDA z&uGsDIgS54N4)Poyd4XZ+;7Mcxzv$tvA?UQmVW3CCajd*%mswEgk+;DIao zLnk#~HfvZq6u+A2r?uyveU{wMX|`tf5GMoj+Bzg{VI%29lClZRe*}Q4OgJpdL zp;^Ir@w=cEWVJJM;?WtrZH9MR1TMS=V7&7OJ0}zR@mqJb@BHuAwCmzwA9LlR*SwKV zPiPG~3pP3(K_+J-b+aQpbhXizE;`PuNkCQ#baDNCehB!Euk$|p>igSm-*lBeI-CX& z4gu_^Vi2-f4ZqJjXFgv7a{Tk#09Yi^NIVNcP_^n}y~(sC z9OrBtx92V2&W~QmEB|q6?Q>UIu)azUyOU|9tzY4zXS=ub*nxWV5)TGOitJpJ+AQeeJvR$tupBxBNUS9 z#$J*p{fP9EPdyS}9J;(ceopX9q|eW&l4|YM=f&ScoLQUI2B4g4;%Yva)SmF1wy7}T zmH;V3&F6T z(Igbfp(MY8wR&zac!1&fdibKv>2&w5gYCP2>BjcIKJ$R`M0&Fgn_kUPr6MyHG;NmT zpNs`l{mKhA`kdFCNk8!X&`-(N20ZDP`!)U1mexq{`@AbY_hGbl(RH-}gj4)#QWZd< zBDVJ0xcxtS8Z?r#XtTNlAfvo%LnH5g2d5Lv2*_M&F#xNH=1yG=DEw>zg}Ps>AuSaz zB+c&kmvzZUPM81kV!MVo8F0s92z$io&!O?2{*V3EUG4kddu!v?pZxS9(*RaC)~L3A zo8Jo$xt`b95``BbI+6`($shw{Fy+d%?AlKje9FLAy%Mj4>96Y3JMZ2D z?aV!owZ|?xGshJ~NudW44YfC3ctLh%cAnJ+pu|^#jn^FX#o@)bVIH(_V%9Xt!DZHM z1Yz+D(q|P~hraRB=W3NEG6FA7`)4tkwwbx^dm6yx1BDs1kKVq!{i9#Jv0bb8^ed(0 z$GNI8c;9m&wt50j0-YyxdcHg_HAeti>IZ^;75?-?9oCm*&0pZ5ys96)u;2A7R?R1y zjnok+tpFFq?+5IP9~$+XiQ^b?5ri*aHUKiWy-kU)i4h*|4DsyZ+bjxhy^&qH1&=M< zx8U8-s=;n*Q~g*G*4KLQ0;*U{vk|W;`niGBZ7=eoP2+i@v7rA{wukp0Zr}H7x3r)A z7(dIEk@U1N=qgSXbTC{i>LZSD(dlkzeU7UjO8_z%m;R`Py_J^6Rpj9u|I(D*+Pah2 zZRORxh;o!03d-l+bs&CIFg5_UQMd3nUv#0lvx2j_17Jd(`=e!DKYYp@TuAM_h(_*Tv#xt5SZnz$=#uCiwN|bny$APaaKE?h2@yYw!_x|cF?V-K# zp1vA6WRwxGW9A8b6&sWAIP=yA7NQnMCjL*n+)gP`>6eZ+4*QUDdX~$0*@ZrDGupiz z)C6k754>E4ETP6pu>47%P~xRUH$q_XKV;9n|6t-l=${8bPG(MPFNzNUz3`m*eFlIm z&T0c7;FSr=S?pxZsd~&!UY)g=pqUAb=0(*vgmrlw>9SBO!nhz|u&B~UT3kJwXaDQ+ z;%i|#pbT*u->SW*x_9G_{p~w`>8AGE@ljqCr(-Aly5jd3DZcoxlp`mkb4x$x2qFDg zOR=ZCdCIbQ2kpvz@Ua_Df;5N1b*j0A$5lMPMqIO~&X=m@O z4=@GJfBB6QdK_PSiGNi{YO^}cY6HmW-8{cFl`*lF$^zAlp;Gf}C^;7A7V<~4mtMxj z2?Zzi4>sV%@=Fj{!b^mO#FZ=;BLP1>+QGw*w;%lQUG0ZHbVoZJzj?34uh44qqacPW zehfn?KrY9U1I+QESObn+Nj9MF0O)A^GQNfft^O?hiiI?gG2u5_nF&k*$d2fB_t@_? z5Y+!+4hR1^7ih*76K@BedGPUe_?&zMoo-VPw{KZC0Oje#oNCgwjx^il*0}j1*Q(ew zGV=@(n($n}ef{)>wPJq7u5}GpnAsVXjTcl|1n%?H~TiE$z1Wp1v9d zqC(2z_a?~@6gMZ2fizSqdhy5Ilkj4ZW8Za8n`K_HQ#TFEzBU=m=Z0a3Kio)PL2w(d zK$0G)^fjTSzlHzQ+UOrcr|*5Nop#O{@}FxjY=xckPd~kV-7}UufJnvmHzvI|qlPdB zJSVdjX0c2|)Hslas%h%ypPJF_DScg?il%OytV4gjndgvfyZZ}dga>2F({Q4+HujlOaqO@uzpKzQN)cKjx=yhGB#uw3axwJ)PBn3d0Zq=ykPQr6_%m^ z9HaOf;vFMPo7;WnrTkBQ;=cBMzkYk$eIUN4Pm%hHPjNXPdP^jDMN_1Tzovn=>4T`; zNe5B%ix!BcOFU;JbAm)KH2ll%8Vy5(obnQ)_)}?a;vYMNpqQkO7Ip^zof~w;mNMkw zG?Ih9LP9S!|J3D|3ZLs~UK>DWUAb2AFx||m+kj3B%*>-fX@n4v=Hjo(4v;pJdo1>C z2@T_heSuZ{qW49YK2GHM)A#Lb-~OIk+yCab_i3zjL9KfT+2D+it57Bd<>!z zckW77v^b?$YU(nMQCS8~O};{UD%FsZO)R8wp4Njao!x&T=s7-HmA=OqM1NT+*OjFv z=&`0=MtM>E8Srnobo$EhBt0K&UK;?}*9^1xBFTY!4*SHgT4=c*yvlOMOiPKCb1o9Z z(m;yll}hGHck+ktX#-+`O$!YT2U1Y_fdgrf%Vu_Sk&gbMZ+F{!=<)WiK742UkH33g z{1!&M{>5SCICB~l^`~5{1YQ9uZVk_as(g^J(5F7(WJ?;8XFaD&g_e#(Epse=6wB5I z|KYdiZIJ233?EzAAn#-1izYH0Vk-amb^oz5V}|8t(*M(!U)C;&O<>M#UK_xeX(ow@ z6a-5=cluxjdz$c4^em5d3YTo)hLA_x3)IlIN0|{XvlHMJlEn$683)oWZ-zB2@wV&q zzixQAedpD;x4U*9Oz#zo{8hNq_wf~E=9}Om>@3irtIIfQY#4~H3MOPRt1GlbkV!et z8gB>!Rd}T54G=n-yQ8eH|$j0ilvNmCCyDJu!zk# zY=x>KNNUoC`4Ayem{pP~6dtwxSM01?~ESe|&2?fML`6j$I~nk^u#qa}xca z=U9Z++m=BmD%pciKYj2~V+h_nkUME$zhzbd}@_oj>I`RTtLX`W}b0ZbF06XpfM zP`tnl$3cRp?6mVU-=dYEZ|{J;4fwh|3yNlVRQnZSY(d>TQM;|%H~n}*zPR(FAHTo- ztKYnL*y)jKVMCNH1S!i<bD_Ayx05iN^r#tS^Ag?U1d}@<3x6Q4x)=r2&ZdsXw?strhFxm{E;S~a+0otT{pc117`J~fcA)uS=8Je7PP-01*1q>cceS7T{T&KX<3Z&S6753EHeSfQv3WKNbg*6`7C4 z;T>7?kh7h$9)$G`VS_wIG{ZT~T7N;%nnQnWO-#$7Tp75>O4t6=Gz!R6deqH1SS|8$AX`|atZCC&t114^i5 zGnDQh))))yeCZ!p{Lk3&aikrJ4M3v_CzqdfM*G_*$?ZR*o7D#JnyW5qXT>J)=#ITP zxteu`5sex2+;@T*B7_V=##dpjm0UBhLROnd*RgR)_C>A5EqN9cdUTGVSo4`b{iS;V0>BUU&s6T zZ~E!$+XwF0lf$P7>u@2Y+vPm}A%ng=ulW}l7OmtIUnseGW&b_@2aQLe1=7mL=5PD=J_e+iQnzQd>G*X<$mWZH-BjE8*hzq8%A zGvCv9{DeDRt;tWOO3pPd;(+F8^U6njB@@PcHOv$vS!0E#9At)UTRufE|Al)@VHpOe zo*Y?KH;aGWzjPFx70Bf{3%8T61w9Wz^cNjaihQQs(9H6JD5psYfY$V(Qr`#)wGzzP*;mEXjvRuYZx^~YF5CoQYvM@gH2V;lSv9vcAuzWn0z+h2P2 z$?@c$F}uy`B>+Kx@ug?AOaH(NYZuQ2!J=TUJ47&M4Vq?;8DxSxOnGS1E4fi#YZDtm zOnidUD_d+BIzp&JDIhP4SQpRGn~WiJFIr&Y1iW>F`;*t~Y;X9P8`|4H8$Z$!jm52T zf)BnmNaCdtlENAbE%inn;h`d}u%(R_J+BziY)=NuXW+p@AKd`*r!f`2sej2Lka0Q| zC>e0ld;GF*TKVEBUO8OJ1y-3S`^Ta>C(DxVwGi-hJJHr12+B@=z)c zq8w+;nI{z~F*zrpB7N+X0TqRCl~j#61f*AXhUlDgyU%e%FZ>gu{S%=41JPfJZwvnw zmtwO2h*pS8jZO`-0=Li4WKX~4=+v_h1{3LAN z&u25VH^v6ACw`^o%CCM|%w~kmOtXjz*Ihv1EQ*yo0{Ty8lwLu$k1l=&J5OZKye;L4 zm!`NYHpdOw{w0*B{W0$E{>Xjp&F{FWeI&lUN6qMf&7!u=K?naLJ0SXu16|}em7z%v zW(fpHdTkTEKj~oxL-HT|G%RA`VE}Wo0~tZ&q6ri}jX4PRA2RFyyCP!ha9Z_`x|qQ7 zFDWS$q00mk3h4*E*sS8%^P&sdi_bc{ea{P@gX&4yW_JhhvWw1Ym&Y%&eDv@+?Um2F ztnIq}Ax%upb6*fz@2l-3=N8t@wRKYfa?3MJ<}vaNrV2N}O#If|-Ltr@>{x zvY{kj_Cjm2hmEQI>g{{l-;3A!cRUi`(`O(o*cD+H^+eazPv7vXSmY0z<#Q^@z)HU! z1~@js@Rfl%-Uy)k_w{q~k8FlmV~kt#j1?Jd<+!ZrXcU=!X>6gVxVG>gcL5AV`M@YR zC;JbBj$i(V@u!;%MEdd<`sM%8^Gro3)4k&sKbbiQ;bA@GrLEvl z#)djYPnw)TUTd5Bm%hWwXf~=Q#gV6d`%0K`C7#l2&ZFD?O90s*uKAZ1w4!&J=;OoV z?7j>c+`kW}5%CVH^$S;-EOw5Avps_~*)3IEp0~A%6+^M?SN&z5ajR(%yCJBeh^_ zYv3>DdSM4gx+#${6%Ipp0UBpz1zOKfVL`)&5)vej2Q!ofEU~M2_++5`R}3Rxf3&0g zr=8FbfwW(?@gbRw{xOMOMU(MAp?{1ckUA^+L38{Sl5A(##T&;EzGG z{PO22j+6Nm528Ia^BLw%dk(i(zW3e;m>G-7y5s)#$nV}#lN*E+(6qQ?@i={1XjyQZ zxqmY9hcAe?I^X=#cH1EqcP{WgD*Ac=7G6AcK@#76^C`2!lV3U|m$%MI6F+2CO~LOx z(QH51RFdu6-{@Z^D6UBsY~8>8bvky@9qrS#uVQ6B%06+LE?729#}>b23@@7Bb#YFx zHR-1=GB0_>W$nNJ{WrCX<72-kdz6(58orAKeho(W?Y5y4nzGgCRlMI(%r%frs57{-M1lB!dYZaq4iB0++t6mhNCP(6k2ocH}1g)Iz8WZU`1V2JA}T4-Q?s)j?nOsON1TXlDqgf}`~<-Ug{1Vs-4ahO@&xR0MK5sCBoL?2KkMu>+K+tO z8{6}rHP1iv>5)GfILns+4De$+54ErOzyp1EX98JxUwF%X?W~X9T&wWdQ^7hwA4p?y zaFx^|iM|#g3n=`EAHtxl!EDwxVcnwTR^gE-nmsRkW_#p?S89>Vx{QXeFa-)tk$~G_ z!13v|f`^8AWJ-^Pp87%Rx6{QBp7_yKTh}i;@RVKTpwW;W`Wfp+|CS$R6&&BDB*^ly zQQx-Gon-YT#H9DQt*?-*gI(-9P6ADQq=P%YGeY{`{Pizs|K}U8lIW>&n(YR_1ij{c z_qFTePdE2M)Ji|CUA}8iyWlsjZD$^MH0C!af4zX$Cm1QzYpE1AQM?%dKbN+d{e5w} zzx$=n(bYc-p%yR>s=yRKxM)QOr}cskT!pYEhtr^|Sdgc!>8LQFkNTtBs9*Ma2FjL; z7aQ<6Rtq$}p)nhfP14bmTu`4$}Aj zg)eG<{tv%Uyi@5k+m`@L(ziZ8Zq;H!HLuKQ-39Ep02!jLg&Xh)rSmc)p&u$7`qN$PlIZg?L73&;OaHLszhvO9`xk5Um!`XI z-oa2S|JCdA`5-{Tf;eV!E{hg2M;KgJIiyY_FH6qj=Qu5WtH zDOvo=O(aV>rZcy;BYxNJ1;2V zo!~L9i(~WsfB*6qx37Qw$@@vaOvcvdS-%8eSl|4yo$crDir=}5wZ(&?yh5@Hya@yr z^mBIXZkJwjS3CdC2kWoAuo6(J6=NJogKipq<4=YkSF3}UUeNYF@0pF;{m0M1pNz{s z<>fBuA9#X)?Dr_bCVjQ>f#|5w2hNgVApp-f_-MQM z=KI_Ex9@1@-M_npCVnm$WOFh7|LvV?%w|_r$2aXvJJacOm}#e_(AHAg(sl}MY1)e5 zrKq3+R-ut#5dt)7f*O^WqIjudkVM2F_{lGrh%rd~VElmbgAm1-egHy)D3oG}Eg)se zv~$0x|KESDwaR>hd~UnzHCMIu+Lf*Ka>`Et#L|r)|3G$&}fH-R0TX6%>lFcvvts>Wkg&S6g-9Pb)h^L+Sq0$q5KNFnmvM8a@QywOtDTJ;cLbB=Qf8cK-#F zn5A#K=F2w3CI5|wZCGuQb( z{;&edTTZ|;;39blK7lf{HN0+D(4wF=HmSvfZI3M%!E2}0hdu@}!jo+1nGaakpX^Wi z;a5*8WL2;q12pQ>j{c;-lCAXH+wMi3qx8k@6P4QGEBeQ>uU&zkX8I!Jd9R?`$0T{i zk2d~af9aO?!8>ng&)K#S&RlW?R>cHBoZg^az~}UqAX{-Kx@gd=?Z{yQAlUgVI{}kj zcn(Y=C6m%}sDlAe23SwDT$y?wSj3WUC3NgWV0hnQ97*Q4^sFx$WkDZWVs>9li{$Ti zQnzF~9fEjeKWfqJpFGih`tl@Y8R&krzPP3KB!QD)Qq_6dryM8vqPv(LvZEg>>Pw;g zta0(lCpLK9K6d4R?PDtasd~H~MZ5E5m$tip;;OceXE1YF7=cxB6U1fx%_}$SdS~!) zibM&Ok2KVRt-4PV&~(e-L_iRIXnJ}3TDsEcefX47&h3D=K}0MUSlI1q2$UVf3QP6q zzwMccK+3t{^)dj{hwvrrceCXP1fREItmpF&%O+wAQ$=_FkJxE95z>7$;?>8-=ap12 z!5IBWoVX0)r46K9YFs)YM8@#~F3e2&g#&KbbE9q1$9n(l=WcB8eC^e9^2d-IbzBuY z0M3ewf1P##pJx|9kU(?*PIUzyeElC1&OOsmbA&}SxeeEVio9qilStJ6SUH2%`S7fV zj;NL`{a|^uT*B;FNndG20}DH>f|Sy+Z3u&{&`QV1RE6%JyRT&WUm+L^2dlA z^SEku0CkqH|J1kIM;_Ymz^=f77XeSLK*?BffI5;>Ky6P1E3TBOh796aI}icSzIhA! z3SjYyANSysym3-r)WBZQ4_Ti80)yRUTgdp-r0+7whJLQ91I}^T2d`rV^wftZ8t15T zEc>21L##gou^zhmA5*Ik`nv*#4D92uf3CLt@B5i++pX7ZPsO?Vb|bLrCV&U`A89*2 z{&3TsfTPjz#`5A<3`eIq>clfaybQ9Q6o^+cN~pHT;}1$Y;4}Cmm?s?6C;g<$mhs{B zpiT$q!%v%FS6i*)vTyyu#YSNbK6XsmD*tUO?fX2Wx}tsi9}f`yU=G{?2b0*OkkDl- z>_8PS`J-O>QDYcm4K}az)9RtWfG5%|mu+nyc+(5ov$mXx)Lf2p1Xj&0K%MV%Hk{sm zOAiEz8pVYo$D!gL{gDPQ;|odi;K#N&DSWO#5Y~0AR!t`8r(Q}sQmtJ%AgcK6e#r8W(mNUJiYA4^2M0VOkNRc6vRFlKdd2z0);qN)41`lZY(bEA{8ttB zWxAhRAD@scb&qZNnx`L52p@Ww!L#(uO(Ba>m=+qx5E+c;>om#!@+&TDfAtGD&dJX} zKEbl;b^x6JVFk&H|7BFA9#nu#1L=eHOw*;wCw4f!AdH}VsD;To2^rDt- z@~U4o`rw~t5go8eMG>!X+&+pf50McaxxzAr`qSqLf>$;JS3TR7a_aX!x_Y_qI@-1) z9orn}$oXCyJirFLzyYex3;rOFj{BiXd*pXNpsV1+x3hntR{KA>~Rl80?R2kDrBfC^`F4KD2g*KFssmufC$a=_PCHhdXLmPISGhCjbWP^`G9=KJv8# z9yjhQ>y{uU0mnAP2sosMz$dUHcr{6gUTtNr6&?7Iybx2-Na=*vD2gut;)uqi9aCXM7RyVNV{b6&=o zuu3QN=i6|T_2SdfM?+GUp{9d=(&_ZKWppt)N|1GvoF24z2iqOYwPuWjJce+5m@yT07LbKJxAL2=qtlJ4=F)d zrCF0)MNRCwuNrbCLe-rH|ELLJQGS%IfQCjOyN@l(Qm{i>#$LWfaMC`GezyydCzr4Z zy-x`C0Reo8bG2`l4;tL^GbV|UgSKElCx%HsVBN2C^H{M15Y@KJE?8+F67tDk*5^Ur z`(Jl;d&&Hu&zI&zAJ)VKz@U9%=fU=hd%o2U(P1Y*twQa$-Dnet5GgFC>MX2=4pei8 zfDWB$RTUr5!s>xN>@X4sdnM<<&&0Kbl!0F|Zj0%JMkskSyM2i-BvIEFm9}LAbJ0)O zoD0b)A&VkMJIAH#C*S>)7ipy(NcRhV$8l!~JNVd(9YbrkUAMjc!EIOiYrJU8<;0J` znwbC?xlB(Oc5BQJlq+i?e(5dsAo_8$bhks*& zcfS2k{eBCPe%WUEG@j!`+jiDECH@udRoBj+=wrB^%(A8?00#11_w8wS>qT)Mknk&% zNR9|fXUdS2l~j>r3DgsM3S<(Z%N$D|`#rhwO3C9{D)eDjaxVDQZK9%wX`xEzsN}BO z383tgWjWC2RXXapElY#{p&v;d6A5ymj87eZppCf2xBon=NIB~9oL;H7_uixAKX=po z_8z1AWS2EJ0WhFU0KfOeJuLbWhY;22y8-$H=$I6e9ap>JRbW>j_f>F6p1>l|hlg>w zGIo7jkj>@q=|^&eLksqE!kFxVkN%LzeNqYA<2iM?oR}v6>#4k*877j#5WZQ19re4P z>`qh`WlGl$hM)L`Qrh39A6I$zt1fT1y>MRj8Na8hoN|92{6q}Q-Iw!o;@S;-X|IBj z=ymvN70ckQzB4K;63|W{y0jzYLzTV+Nam>oCVg@IB%3+GdNk4oU!pUfX#DeSMM#;9 z1|{f~bPxjH#}(`&-*0t{B*}1~Dx-25fh6nVA8j9QCTpHvoJMGg8O7Kku;8n?c2KS0 zBO309@gLW!|NF1JtZhI0j0Bsjr)mV&S0eT zk2Lx!=0mT^pmZ!d@DW~t78A{6tM3pjTM4S#B9k1s)qTTu+U+a-3bt%tvgLEvm%8<{ z#?|okIP8}WwhwW0e~)i2 zPu~cv`3Zob|GWP_&|d$UU5(EQMXr-Pne^;cGpU>~WWkUiuzK=|L4r)H_K+|=F_Xs1 z2((wo!VKOtL1Yr6>lfZiX0*ScKgc6!->6f+bm3VK)}WK}BaQR%Zf2A#eNI3)Hsor+ zLMHLbK7PMIpXPh7R{iUqeKrzvd74LHo&fR!eCE-^?bZ5a;fD?=N=X<;eg?QIcSVk1 zM+}Gt(UYnyyyZy4++Kj_;ioJJOK<>;KBX>)9%nt z;GZ8E{-~(<1TiI&m3BTVt6&fI6!-|6H|Qh=NkZaO;v2j7m+T-rV9^A(PdaKZ>RTYZ zV`Uz3v#shV`%@yPUuh?N4a{m+LYH6NKCY(?ctLM%opuRtzVZC_Yd>&td*+5U`OKLC zJZZ~30W4jbBMRi-{o?NSo-gfbv>oV@8CP-w-NgYfi41&bl_0emR-#y^@X$BJG!dXn zB3eJ;J1%|hpH)wIHWtYa!#wHd)*ieC{oJZs-;;v{5^usS`^YVx%Q5Wt^Ix?8(n~hC zcfaD&cGdQ^{pYtB>65k06TstKnScD|!S+UdlklMf`t#$K(UHVeaX;Ip1$i;y27c&E zkwh6}JqbX|1U4UZbXU?}O77B+~{We zPSCF2zOlXY`!8xQzheHO4u<$-F7pI%jH~kv{q^Cu|J!5j!}>n~u1J}~9yQF7Kr zeL9x~AD8T3qA%|J`uAScUbExe`97b4Jq5}<0UXCg`Um|1(zo6FXuDs(GUQJS@+m%^ zfK!wN&LEPZo{BqbgWQuRWf7@uItbOSS+5_+& zeY#y(_`&b%QHM>GRL11z9Fs;>_RrVHcHjQu3)>rSIKQpuySsBaWkz700FLug?b8E+ zcYR@Zd;k4=+QCecu1Yc~QB-AusF8^S)hm8aq|{QsWUM&$5kY?Ga{y!V#oKS=L^0V! zH&$w3mXkJkmXS*NhRJ0b_N>o|fvM7cQ2(OMXSBE7bU}NAKEAtg{ro+B56~%gm?wZI zc)`A`4-4G=?~k^>)NMh}c#_OXk{Lp>WYUl@kw|J(5}A(xcihO9e8~<;E8citE+yNl z#5T59=D-MBsn0}(JcqI=(kN4_c*%4gs_3W`ul&5~+zss)Z+cdH?KRux72hLt>K^6^ z;5)o_pMCUD`;E^%+W!9Q`;{T_wqX3d005OTPFyuJ@gmi8qCjLRAtCke)6+ITlXcj7 zI~fQ5WG8i0O{;p?tI1+PwqxXEQ3;9l)`#2L2K@-v=x%u@~IN-NMvQH=xoVz|v}*3GVHDwT07n1UaEpA}{X}pm$j(xs{#x zvCD{e0<53>OZDx%pS)sAyZy?o?V9s8oZyj};ZDX8m?wZI@v5*Z`0!Wuw!i+$p7tM) z9gbDEue`Y$4=efOC9}8a-mA-5>9ayp`n|s2TlDNA{;*3}uob%SjRAnl)mzVOKla>n z+AA;H++H+q@ncZd&~csso}!EWpPCFlrso77``Z5YiLdW#`}K3>tk6!zN?`fYlX@X} zW2ev)mn8PYfqtHVly6@+cvbDF4)*^zTm8O`Ib51pN30*(wdwnfRp6XJgn~=d{J*L-ltoI`yM;k z?tko1`OTI` rv)4P(FWJ1_JA#ws6lXV1!V&mC`?$)n8 = { google_pse: googlePseLogo.src, dataforseo: dataforseoLogo.src, nimble: nimbleLogo.src, + bing_grounding: bingLogo.src, }; interface SearchProviderLabelProps { From dd2e1cf7a8ab6cbbc934d78690e249cb5719ebec Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 24 Aug 2026 11:08:59 -0700 Subject: [PATCH 30/93] fix: read a single-quoted DO body for routine calls, not only dollar-quoted ones --- .../check_migrations_no_data_rewrites.py | 50 +++++++++++++++---- .../test_check_migrations_no_data_rewrites.py | 20 ++++++++ 2 files changed, 60 insertions(+), 10 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 8539a258083..c8c8dd315a8 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -625,6 +625,7 @@ def scan_region( reports its real file line and lines up with the markers read from that file.""" masked, bodies, literals = mask(region) executed = executed_names(masked) + runnable = executed_literals(masked, literals, executed) for match in STATEMENT.finditer(masked): exempt = markers.exempt(offset + statement_start(match), offset + match.end()) @@ -642,14 +643,37 @@ def scan_region( yield Violation(migration, line_of(document, offset + keyword_start(clause, base)), keyword) for body in bodies: - if not runs_when_applied(masked, region, bodies, body): + if not runs_when_applied(masked, region, bodies, runnable, body): continue start, end = body yield from scan_region(document, region[start:end], migration, markers, offset + start) +def executed_literals( + masked: str, literals: tuple[tuple[int, int], ...], executed: frozenset[str] +) -> tuple[tuple[int, int], ...]: + """The single-quoted literals a region runs as SQL, where a call to a routine the same + migration defines is as real as one written in the open. `DO '...'` runs its body and + `EXECUTE` runs the string it is handed, so a definition named inside one of those is called, + while a name in a message string or any literal nothing executes stays text. These are the + spans the direct scan already recurses into, read here so a call written in one is found when + the migration is searched for the routine's name.""" + return tuple( + (start, end) + for match in STATEMENT.finditer(masked) + for clause, base in clauses(match.group(), match.start()) + if hands_off_sql(clause, executed) + for start, end in literals + if base <= start and end <= base + bind_values_start(clause) + ) + + def runs_when_applied( - masked: str, region: str, bodies: tuple[tuple[int, int], ...], body: tuple[int, int] + masked: str, + region: str, + bodies: tuple[tuple[int, int], ...], + runnable: tuple[tuple[int, int], ...], + body: tuple[int, int], ) -> bool: """Whether a dollar-quoted body runs while the migration is being applied. A `DO` block runs where it is written, and so does every other use of this quoting. A `CREATE FUNCTION` or a @@ -671,22 +695,28 @@ def runs_when_applied( named = ROUTINE_NAME.match(region, defined.end(), start) if named is None or named.group(1).startswith('"'): return True - return contains(outside_definition(masked, region, bodies, opens, end), re.escape(named.group(1))) + return contains(outside_definition(masked, region, bodies, runnable, opens, end), re.escape(named.group(1))) def outside_definition( - masked: str, region: str, bodies: tuple[tuple[int, int], ...], opens: int, closes: int + masked: str, + region: str, + bodies: tuple[tuple[int, int], ...], + runnable: tuple[tuple[int, int], ...], + opens: int, + closes: int, ) -> str: - """The migration's text with one routine definition blanked out and every dollar-quoted body - put back. Masking blanks the bodies alike, and a `DO` block is the ordinary way a migration - runs a routine it has just defined, so a call written inside one has to stay readable. Each - body comes back with its comments blanked, since a name written in a comment is - documentation rather than a call, while its string literals stay readable because `EXECUTE` + """The migration's text with one routine definition blanked out and every runnable body put + back: the dollar-quoted bodies and the single-quoted literals `DO` and `EXECUTE` run as SQL. + Masking blanks all of them alike, and a `DO` block, dollar-quoted or single-quoted, is the + ordinary way a migration runs a routine it has just defined, so a call written inside one has + to stay readable. Each comes back with its comments blanked, since a name written in a comment + is documentation rather than a call, while its string literals stay readable because `EXECUTE` runs one as SQL and the call can be written inside it. The definition is blanked after they are restored, which takes its own body with it, so a routine that names itself recursively does not thereby count as called.""" text = list(masked) - for start, end in bodies: + for start, end in (*bodies, *runnable): text[start:end] = without_comments(region[start:end]) text[opens:closes] = blank(region[opens:closes]) return "".join(text) diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 932572a0264..e047218de5f 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -452,6 +452,26 @@ class TestStoredRoutines: sql = self.DEFINITION + "DO $$ BEGIN RAISE NOTICE '--'; PERFORM backfill(); END; $$;\n" assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_a_call_from_inside_a_single_quoted_do_block_still_counts(self, tmp_path): + sql = self.DEFINITION + "DO 'BEGIN PERFORM backfill(); END';\n" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_an_executed_literal_inside_a_single_quoted_do_counts_as_a_call(self, tmp_path): + sql = self.DEFINITION + "DO 'BEGIN EXECUTE ''SELECT backfill()''; END';\n" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_variable_run_by_execute_in_a_single_quoted_do_counts_as_a_call(self, tmp_path): + sql = self.DEFINITION + "DO 'DECLARE q text; BEGIN q := ''SELECT backfill()''; EXECUTE q; END';\n" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_the_name_written_only_in_a_single_quoted_do_comment_is_not_a_call(self, tmp_path): + sql = self.DEFINITION + "DO 'BEGIN\n-- backfill() runs later\nPERFORM 1; END';\n" + assert _keywords(tmp_path, sql) == () + + def test_the_name_written_only_in_a_non_runnable_string_is_not_a_call(self, tmp_path): + sql = self.DEFINITION + "SELECT 'backfill() runs after the deploy';\n" + assert _keywords(tmp_path, sql) == () + def test_a_recursive_call_does_not_count_as_the_migration_calling_it(self, tmp_path): sql = ( "CREATE FUNCTION backfill(n int) RETURNS void AS $$\n" From e43b496000ee1d46ac4a08661caef0f1d0b23b86 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 24 Aug 2026 11:45:04 -0700 Subject: [PATCH 31/93] fix: undouble single-quoted payloads before reading them for routine calls Call detection restored a single-quoted DO or EXECUTE payload through without_comments while its `''` escapes were still doubled. The first quote of a pair opened an empty string and closed it on the second, leaving a `--` or `/*` from a nested string bare, so it blanked the real call after it and the routine read as uncalled: its rewrite body then went unscanned at boot. Undouble each single-quoted payload before restoring it, and pad it back to the span it fills so the later offsets still land. Dollar-quoted bodies do not escape quotes and are left as they were. --- .../check_migrations_no_data_rewrites.py | 21 +++++++++++++++---- .../test_check_migrations_no_data_rewrites.py | 12 +++++++++++ 2 files changed, 29 insertions(+), 4 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index c8c8dd315a8..358b46db80b 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -240,6 +240,15 @@ def blank(text: str) -> str: return "".join(character if character == "\n" else " " for character in text) +def undouble(literal: str) -> str: + """The SQL a single-quoted literal stands for, with each doubled quote read back as the one it + escapes. `mask` hands the literal on raw, `''` and all, so re-lexing it as SQL needs the escapes + resolved first: left doubled, the first quote of a pair opens an empty string and closes it on + the second, and a `--` or `/*` in what was a nested string is then bare and blanks the code + after it.""" + return literal.replace("''", "'") + + def mask(sql: str) -> tuple[str, tuple[tuple[int, int], ...], tuple[tuple[int, int], ...]]: """Blank comments and quoted text, keeping offsets, and locate the spans that can still hold SQL: dollar-quoted bodies, and the single-quoted literals `EXECUTE` runs.""" @@ -712,12 +721,16 @@ def outside_definition( ordinary way a migration runs a routine it has just defined, so a call written inside one has to stay readable. Each comes back with its comments blanked, since a name written in a comment is documentation rather than a call, while its string literals stay readable because `EXECUTE` - runs one as SQL and the call can be written inside it. The definition is blanked after they - are restored, which takes its own body with it, so a routine that names itself recursively - does not thereby count as called.""" + runs one as SQL and the call can be written inside it. A single-quoted payload is undoubled as + it goes back, so a `--` or `/*` in one of its nested strings blanks nothing and the call after + it stays visible, and it is padded to the span it fills so the later offsets still land. The + definition is blanked after they are restored, which takes its own body with it, so a routine + that names itself recursively does not thereby count as called.""" text = list(masked) - for start, end in (*bodies, *runnable): + for start, end in bodies: text[start:end] = without_comments(region[start:end]) + for start, end in runnable: + text[start:end] = without_comments(undouble(region[start:end])).ljust(end - start) text[opens:closes] = blank(region[opens:closes]) return "".join(text) diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index e047218de5f..2987e981dde 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -464,6 +464,18 @@ class TestStoredRoutines: sql = self.DEFINITION + "DO 'DECLARE q text; BEGIN q := ''SELECT backfill()''; EXECUTE q; END';\n" assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_a_call_after_a_single_quoted_literal_holding_comment_dashes_still_counts(self, tmp_path): + sql = self.DEFINITION + "DO 'BEGIN RAISE NOTICE ''--''; PERFORM backfill(); END';\n" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_a_call_after_a_single_quoted_literal_opening_a_block_comment_still_counts(self, tmp_path): + sql = self.DEFINITION + "DO 'BEGIN RAISE NOTICE ''/*''; PERFORM backfill(); END';\n" + assert _keywords(tmp_path, sql) == ("UPDATE",) + + def test_an_executed_literal_after_a_single_quoted_comment_dash_string_still_counts(self, tmp_path): + sql = self.DEFINITION + "DO 'BEGIN RAISE NOTICE ''--''; EXECUTE ''SELECT backfill()''; END';\n" + assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_the_name_written_only_in_a_single_quoted_do_comment_is_not_a_call(self, tmp_path): sql = self.DEFINITION + "DO 'BEGIN\n-- backfill() runs later\nPERFORM 1; END';\n" assert _keywords(tmp_path, sql) == () From 4ddf1e5c6d7819cddd9316c7bb643037f5859ead Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 24 Aug 2026 11:57:45 -0700 Subject: [PATCH 32/93] test: guard the restore padding against a long escaped-quote run before a definition --- tests/test_litellm/test_check_migrations_no_data_rewrites.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 2987e981dde..280d1cb698c 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -476,6 +476,11 @@ class TestStoredRoutines: sql = self.DEFINITION + "DO 'BEGIN RAISE NOTICE ''--''; EXECUTE ''SELECT backfill()''; END';\n" assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_a_long_run_of_escaped_quotes_before_an_uncalled_definition_is_not_a_call(self, tmp_path): + escaped_quotes = "'" * 84 + sql = f"DO 'BEGIN RAISE NOTICE ''{escaped_quotes}''; PERFORM 1; END';\n" + self.DEFINITION + assert _keywords(tmp_path, sql) == () + def test_the_name_written_only_in_a_single_quoted_do_comment_is_not_a_call(self, tmp_path): sql = self.DEFINITION + "DO 'BEGIN\n-- backfill() runs later\nPERFORM 1; END';\n" assert _keywords(tmp_path, sql) == () From 79d0d7a48e4d3c90cee5a7350467fb44bf2791cc Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Mon, 24 Aug 2026 11:59:12 -0700 Subject: [PATCH 33/93] chore(codeowners): drop ryan-crabbe-berri from ui infra and generated files The /ui/ rule sweeps in the container plumbing (Dockerfile, nginx.conf) and the checked-in tsbuildinfo, none of which are dashboard code. Exempt them so review requests land on the people who actually own that surface. --- .github/CODEOWNERS | 3 +++ 1 file changed, 3 insertions(+) diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 7ae79aa666f..5eb94e561ba 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -1,6 +1,9 @@ /ui/ @yuneng-berri @ryan-crabbe-berri /litellm/proxy/_experimental/out/ @yuneng-berri @ryan-crabbe-berri +/ui/Dockerfile @yuneng-berri +/ui/nginx.conf @yuneng-berri /ui/litellm-dashboard/src/lib/http/schema.d.ts +/ui/litellm-dashboard/tsconfig.tsbuildinfo /model_prices_and_context_window.json @mateo-berri /litellm/model_prices_and_context_window_backup.json @mateo-berri /litellm-proxy-extras/litellm_proxy_extras/migrations/ @yuneng-berri @ryan-crabbe-berri From 5ed83942dcf583fa5e2ce93bed5f99ead2a224c4 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Mon, 24 Aug 2026 12:04:17 -0700 Subject: [PATCH 34/93] chore(codeowners): unown ui container plumbing entirely The Dockerfile and nginx.conf are deploy infra, so nobody on the dashboard side needs to gate them. Drop the remaining owner instead of making one person the sole required approver. --- .github/CODEOWNERS | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 5eb94e561ba..cfa0390e836 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -1,7 +1,7 @@ /ui/ @yuneng-berri @ryan-crabbe-berri /litellm/proxy/_experimental/out/ @yuneng-berri @ryan-crabbe-berri -/ui/Dockerfile @yuneng-berri -/ui/nginx.conf @yuneng-berri +/ui/Dockerfile +/ui/nginx.conf /ui/litellm-dashboard/src/lib/http/schema.d.ts /ui/litellm-dashboard/tsconfig.tsbuildinfo /model_prices_and_context_window.json @mateo-berri From a0c319878b3f0b10b287dc3d80ed9e36747ba513 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 24 Aug 2026 12:26:07 -0700 Subject: [PATCH 35/93] fix(search): harden bing_grounding auth, result cap, status, and cost - send a caller api_key via the Azure api-key header instead of Authorization: Bearer - cap web_search results to the requested max_results (the tool has no count knob) - surface a Foundry failed/incomplete response status as a 502 error - zero the per-query cost in web_search mode; keep the map price for connection mode - trim the example config to terse env-var pointers --- litellm/llms/azure/search/transformation.py | 130 ++++++++++++++---- litellm/llms/custom_httpx/llm_http_handler.py | 2 + .../bing_grounding_websearch_config.yaml | 34 ++--- .../test_bing_grounding_search.py | 14 +- ...st_bing_grounding_search_transformation.py | 63 ++++++++- 5 files changed, 190 insertions(+), 53 deletions(-) diff --git a/litellm/llms/azure/search/transformation.py b/litellm/llms/azure/search/transformation.py index 2caee9b50a0..7bc631e814d 100644 --- a/litellm/llms/azure/search/transformation.py +++ b/litellm/llms/azure/search/transformation.py @@ -12,10 +12,11 @@ Setup: 3. Optional: set BING_GROUNDING_CONNECTION_ID to a Grounding with Bing Search project connection id to use the `bing_grounding` tool; without it the project's built-in `web_search` tool is used - 4. Auth: pass api_key, or set BING_GROUNDING_TOKEN to an Entra bearer token for - scope https://ai.azure.com/.default, or configure azure-identity - (AZURE_CLIENT_ID / AZURE_CLIENT_SECRET / AZURE_TENANT_ID, managed identity, - or any DefaultAzureCredential source) and the token is minted automatically + 4. Auth: pass api_key (an Azure API key, sent in the api-key header), or set + BING_GROUNDING_TOKEN to an Entra bearer token for scope + https://ai.azure.com/.default, or configure azure-identity (AZURE_CLIENT_ID / + AZURE_CLIENT_SECRET / AZURE_TENANT_ID, managed identity, or any + DefaultAzureCredential source) and the token is minted automatically Usage: response = litellm.search( @@ -56,6 +57,8 @@ ENTRA_SCOPE: Final = "https://ai.azure.com/.default" _RESPONSES_PATH: Final = "/openai/v1/responses" _SNIPPET_FALLBACK_LENGTH: Final = 300 +_UPSTREAM_ERROR_STATUS: Final = 502 +_RESPONSE_COST_HEADER: Final = "llm_provider-x-litellm-response-cost" class _Annotation(BaseModel): @@ -83,21 +86,33 @@ class _OutputItem(BaseModel): content: tuple[_ContentPart, ...] = () -class _ResponsesEnvelope(BaseModel): - """A Foundry Responses API body. `output` is required: a body without it is not a - Responses API response and must not be reported as a successful empty search.""" - - model_config = ConfigDict(extra="ignore", frozen=True) - - output: tuple[_OutputItem, ...] - - class _ErrorBody(BaseModel): model_config = ConfigDict(extra="ignore", frozen=True) message: str | None = None +class _IncompleteDetails(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + reason: str | None = None + + +class _ResponsesEnvelope(BaseModel): + """A Foundry Responses API body. `output` is required: a body without it is not a + Responses API response and must not be reported as a successful empty search. + + A 200 body can still carry `status` `failed` or `incomplete`; those are surfaced as + errors rather than reported as a successful empty search.""" + + model_config = ConfigDict(extra="ignore", frozen=True) + + output: tuple[_OutputItem, ...] + status: str | None = None + error: _ErrorBody | None = None + incomplete_details: _IncompleteDetails | None = None + + class _ErrorEnvelope(BaseModel): model_config = ConfigDict(extra="ignore", frozen=True) @@ -163,6 +178,26 @@ def _citation_results(envelope: _ResponsesEnvelope) -> tuple[SearchResult, ...]: return tuple(first_by_url[url] for url in dict.fromkeys(result.url for result in cited)) +def _requested_max_results(response_kwargs: Mapping[str, object]) -> int | None: + """The unified `max_results` cap the caller asked for, if any. + + The built-in web_search tool has no server-side result-count knob, so the cap is + enforced here after the fact; connection mode also honors it as a hard ceiling on + top of the tool's `count` hint. + """ + optional_params: Final = response_kwargs.get("optional_params") + if not isinstance(optional_params, Mapping): + return None + max_results: Final = optional_params.get("max_results") + return ( + max_results if isinstance(max_results, int) and not isinstance(max_results, bool) and max_results > 0 else None + ) + + +def _capped(results: tuple[SearchResult, ...], max_results: int | None) -> tuple[SearchResult, ...]: + return results[:max_results] if max_results is not None else results + + class _SearchConfiguration(BaseModel): model_config = ConfigDict(frozen=True) @@ -247,18 +282,28 @@ class BingGroundingSearchConfig(BaseSearchConfig): Returns a new dict rather than mutating ``headers``: the http handler calls this a second time after ``litellm/search/main.py`` already did, so it has to be idempotent. """ - resolved_token: Final = self.resolve_server_api_key( - caller_api_key=api_key, + return { # mutable-ok: httpx requires a plain dict of headers + **headers, + **self._auth_header(api_key, api_base), + "Content-Type": "application/json", + } + + def _auth_header(self, api_key: str | None, api_base: str | None) -> Mapping[str, str]: + """ + A caller-supplied ``api_key`` is an Azure API key and rides the ``api-key`` header; + an Entra bearer token (``BING_GROUNDING_TOKEN`` or one minted via azure-identity) + rides ``Authorization: Bearer``. Foundry rejects the wrong scheme for each. + """ + if api_key: + return MappingProxyType({"api-key": api_key}) + token: Final = self.resolve_server_api_key( + caller_api_key=None, caller_api_base=api_base, key_env_vars=(TOKEN_ENV,), base_env_var=PROJECT_ENDPOINT_ENV, default_api_base=None, ) or self._mint_entra_token(api_base) - return { # mutable-ok: httpx requires a plain dict of headers - **headers, - "Authorization": f"Bearer {resolved_token}", - "Content-Type": "application/json", - } + return MappingProxyType({"Authorization": f"Bearer {token}"}) def _mint_entra_token(self, caller_api_base: str | None) -> str: self._assert_trusted_api_base_for_server_credential( @@ -303,8 +348,9 @@ class BingGroundingSearchConfig(BaseSearchConfig): Transform Search request to the Foundry Responses API format. The unified params map as far as the API allows: - - max_results -> the bing_grounding search configuration's `count` (the built-in - web_search tool has no result-count knob, so it is dropped in that mode) + - max_results -> the bing_grounding search configuration's `count`; the built-in + web_search tool has no result-count knob, so that mode instead caps the returned + results after the fact (see transform_search_response) - country -> web_search's approximate `user_location` (bing_grounding's `market` wants a full locale like en-US, which a bare country code cannot fill) - search_domain_filter, max_tokens_per_page -> no API equivalent, dropped @@ -336,8 +382,44 @@ class BingGroundingSearchConfig(BaseSearchConfig): status_code=raw_response.status_code, headers=dict(raw_response.headers), # mutable-ok: BaseSearchConfig.get_error_class signature ) - results: Final = list(_citation_results(parsed)) # mutable-ok: SearchResponse.results is list[SearchResult] - return SearchResponse(results=results, object="search") + if parsed.status == "failed": + detail: Final = ( + parsed.error.message if parsed.error and parsed.error.message else "the grounded search failed" + ) + raise self._upstream_error(detail, raw_response) + results: Final = _capped(_citation_results(parsed), _requested_max_results(kwargs)) + if not results and parsed.status == "incomplete": + reason: Final = ( + parsed.incomplete_details.reason + if parsed.incomplete_details and parsed.incomplete_details.reason + else "unknown reason" + ) + raise self._upstream_error(f"the grounded search was incomplete: {reason}", raw_response) + return self._priced(results) + + def _upstream_error(self, detail: str, raw_response: httpx.Response) -> Exception: + return self.get_error_class( + error_message=detail, + status_code=_UPSTREAM_ERROR_STATUS, + headers=dict(raw_response.headers), # mutable-ok: BaseSearchConfig.get_error_class signature + ) + + def _priced(self, results: tuple[SearchResult, ...]) -> SearchResponse: + """web_search mode runs no paid Grounding with Bing transaction, so it must not + inherit the connection-mode ``bing_grounding/search`` price; zero its per-query + cost while leaving connection mode to the cost map.""" + response: Final = SearchResponse( + results=list(results), # mutable-ok: SearchResponse.results is list[SearchResult] + object="search", + ) + if get_secret_str(CONNECTION_ID_ENV): + return response + response._hidden_params[ + "additional_headers" + ] = { # mutable-ok: response_cost_calculator writes into _hidden_params + _RESPONSE_COST_HEADER: 0.0 + } + return response def get_error_class( self, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index ed079197513..8fa025d6e4c 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1844,6 +1844,7 @@ class BaseLLMHTTPHandler: return provider_config.transform_search_response( raw_response=response, logging_obj=logging_obj, + optional_params=optional_params, ) async def async_search( @@ -1942,6 +1943,7 @@ class BaseLLMHTTPHandler: return provider_config.transform_search_response( raw_response=response, logging_obj=logging_obj, + optional_params=optional_params, ) async def _async_post_anthropic_messages_with_http_error_retry( diff --git a/litellm/proxy/example_config_yaml/bing_grounding_websearch_config.yaml b/litellm/proxy/example_config_yaml/bing_grounding_websearch_config.yaml index 5ab18723f7f..ac8a3db2d21 100644 --- a/litellm/proxy/example_config_yaml/bing_grounding_websearch_config.yaml +++ b/litellm/proxy/example_config_yaml/bing_grounding_websearch_config.yaml @@ -1,24 +1,14 @@ -# Web search via Microsoft Foundry: Grounding with Bing Search / the built-in -# web_search tool, called through the Foundry Responses API. -# See litellm/llms/azure/search/transformation.py for details. +# Web search via Microsoft Foundry (Grounding with Bing Search / the built-in +# web_search tool), called through the Foundry Responses API. # -# Required environment variables (the search router forwards only -# search_provider / api_key / api_base from the litellm_params block, so -# provider configuration rides env vars): -# BING_GROUNDING_PROJECT_ENDPOINT: the Foundry project endpoint, e.g. -# https://.services.ai.azure.com/api/projects/ -# BING_GROUNDING_MODEL: a model deployment in that project (e.g. gpt-4.1); -# it runs the grounded search, its tokens are billed on that deployment -# Optional: -# BING_GROUNDING_CONNECTION_ID: a Grounding with Bing Search project -# connection id; set it to use the bing_grounding tool ($35 per 1,000 -# transactions on the G1 SKU). Without it the project's built-in -# web_search tool is used -# BING_GROUNDING_TOKEN: an Entra bearer token for scope -# https://ai.azure.com/.default. Without it (and without api_key below) -# the token is minted via azure-identity (AZURE_CLIENT_ID / -# AZURE_CLIENT_SECRET / AZURE_TENANT_ID, managed identity, or any other -# DefaultAzureCredential source) +# Configure the provider with env vars (setup and pricing are in the LiteLLM docs; +# the code lives in litellm/llms/azure/search/transformation.py): +# BING_GROUNDING_PROJECT_ENDPOINT (required) the Foundry project endpoint +# BING_GROUNDING_MODEL (required) a model deployment in that project +# BING_GROUNDING_CONNECTION_ID (optional) a Grounding with Bing connection id; +# without it the built-in web_search tool is used +# BING_GROUNDING_TOKEN (optional) an Entra bearer token; without it (and +# without api_key) azure-identity mints one model_list: - model_name: claude-sonnet @@ -30,8 +20,8 @@ search_tools: - search_tool_name: bing-grounding-search litellm_params: search_provider: bing_grounding - # Alternative to BING_GROUNDING_TOKEN / azure-identity: - # api_key: os.environ/BING_GROUNDING_TOKEN + # Optional: an Azure API key instead of BING_GROUNDING_TOKEN / azure-identity + # api_key: os.environ/AZURE_AI_API_KEY litellm_settings: callbacks: ["websearch_interception"] diff --git a/tests/search_tests/test_bing_grounding_search.py b/tests/search_tests/test_bing_grounding_search.py index 00e5e382eef..3d1737477a1 100644 --- a/tests/search_tests/test_bing_grounding_search.py +++ b/tests/search_tests/test_bing_grounding_search.py @@ -175,7 +175,19 @@ class TestBingGroundingSearchTransformation: assert mock_post.call_args.kwargs["json"]["tools"] == [{"type": "web_search"}] assert len(response.results) == 2 - def test_bing_grounding_search_tracks_cost(self, monkeypatch: pytest.MonkeyPatch): + def test_web_search_mode_is_not_billed_the_g1_price(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + with patch( # test-quality-ok: litellm.search has no client injection seam + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=_mock_response(), + ): + response = litellm.search(query="pricing check", search_provider="bing_grounding") + + assert response._hidden_params["response_cost"] == 0.0 + + def test_connection_mode_tracks_the_g1_cost(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_CONNECTION_ID", "conn-id") monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) with patch( # test-quality-ok: litellm.search has no client injection seam diff --git a/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py b/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py index 25dfa5fbe2b..d33c9f92df6 100644 --- a/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py +++ b/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py @@ -55,15 +55,18 @@ def test_ui_friendly_name(): assert _config().ui_friendly_name() == "Grounding with Bing Search" -def test_validate_environment_with_explicit_key(): - headers = _config().validate_environment({}, api_key="explicit-token") - assert headers["Authorization"] == "Bearer explicit-token" +def test_validate_environment_api_key_uses_api_key_header_not_bearer(): + headers = _config().validate_environment({}, api_key="azure-api-key") + assert headers["api-key"] == "azure-api-key" + assert "Authorization" not in headers assert headers["Content-Type"] == "application/json" def test_validate_environment_reads_env_token(monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("BING_GROUNDING_TOKEN", "env-token") - assert _config().validate_environment({})["Authorization"] == "Bearer env-token" + headers = _config().validate_environment({}) + assert headers["Authorization"] == "Bearer env-token" + assert "api-key" not in headers def test_validate_environment_falls_back_to_entra_minter(): @@ -74,8 +77,9 @@ def test_validate_environment_falls_back_to_entra_minter(): def test_validate_environment_api_key_beats_env_token(monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("BING_GROUNDING_TOKEN", "env-token") minter = Mock(return_value="entra-token") - headers = _config(entra_token_minter=minter).validate_environment({}, api_key="explicit-token") - assert headers["Authorization"] == "Bearer explicit-token" + headers = _config(entra_token_minter=minter).validate_environment({}, api_key="azure-api-key") + assert headers["api-key"] == "azure-api-key" + assert "Authorization" not in headers minter.assert_not_called() @@ -265,6 +269,53 @@ def test_transform_search_response_malformed_body_raises_instead_of_reporting_em _config().transform_search_response(_resp(body, status_code=502), logging_obj=Mock()) +def test_transform_search_response_caps_results_to_max_results(): + annotations = [_citation(f"https://example.com/{i}", f"T{i}", 0, 5) for i in range(5)] + resp = _config().transform_search_response( + _resp(_message_response("claim", annotations)), logging_obj=Mock(), optional_params={"max_results": 2} + ) + assert [r.url for r in resp.results] == ["https://example.com/0", "https://example.com/1"] + + +def test_transform_search_response_without_max_results_returns_all_citations(): + annotations = [_citation(f"https://example.com/{i}", f"T{i}", 0, 5) for i in range(4)] + resp = _config().transform_search_response(_resp(_message_response("c", annotations)), logging_obj=Mock()) + assert len(resp.results) == 4 + + +def test_transform_search_response_failed_status_raises_with_error_message(): + payload = {"output": [], "status": "failed", "error": {"message": "content was filtered"}} + with pytest.raises(Exception, match="content was filtered") as excinfo: + _config().transform_search_response(_resp(payload), logging_obj=Mock()) + assert excinfo.value.status_code == 502 + + +def test_transform_search_response_incomplete_with_no_results_raises_with_reason(): + payload = {"output": [], "status": "incomplete", "incomplete_details": {"reason": "max_output_tokens"}} + with pytest.raises(Exception, match="incomplete: max_output_tokens"): + _config().transform_search_response(_resp(payload), logging_obj=Mock()) + + +def test_transform_search_response_incomplete_with_partial_results_returns_them(): + payload = _message_response("claim", [_citation("https://example.com", "Example", 0, 5)]) + payload["status"] = "incomplete" + resp = _config().transform_search_response(_resp(payload), logging_obj=Mock()) + assert [r.url for r in resp.results] == ["https://example.com"] + + +def test_transform_search_response_web_search_mode_zeroes_per_query_cost(): + payload = _message_response("claim", [_citation("https://example.com", "Example", 0, 5)]) + resp = _config().transform_search_response(_resp(payload), logging_obj=Mock()) + assert resp._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] == 0.0 + + +def test_transform_search_response_connection_mode_leaves_price_to_cost_map(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("BING_GROUNDING_CONNECTION_ID", "conn-id") + payload = _message_response("claim", [_citation("https://example.com", "Example", 0, 5)]) + resp = _config().transform_search_response(_resp(payload), logging_obj=Mock()) + assert "additional_headers" not in resp._hidden_params + + def test_get_error_class_attributes_the_provider(): error = _config().get_error_class(error_message="quota exceeded", status_code=429, headers={}) assert error.status_code == 429 From 422b24023c3a7d02d7f02ffebea1fd8ae059ca84 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 24 Aug 2026 12:39:19 -0700 Subject: [PATCH 36/93] fix: undouble an executed literal before scanning, so a doubled-quote comment can't hide a rewrite --- .../check_migrations_no_data_rewrites.py | 13 +++++++++++-- .../test_check_migrations_no_data_rewrites.py | 13 +++++++++++++ 2 files changed, 24 insertions(+), 2 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 358b46db80b..39fdfe716cf 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -631,7 +631,10 @@ def scan_region( ) -> Iterator[Violation]: """Violations in one region of `document`, whose text begins at `offset`. Positions are always counted against the whole document, so a statement nested in a dollar-quoted body - reports its real file line and lines up with the markers read from that file.""" + reports its real file line and lines up with the markers read from that file. A single-quoted + literal that `DO` or `EXECUTE` runs as SQL is undoubled before it is scanned, so a `--` or `/*` + in one of its nested strings blanks nothing and the statement after it stays visible, and it is + padded back to its span so the offsets still land.""" masked, bodies, literals = mask(region) executed = executed_names(masked) runnable = executed_literals(masked, literals, executed) @@ -644,7 +647,13 @@ def scan_region( commands_end = base + bind_values_start(clause) for start, end in literals: if base <= start and end <= commands_end: - yield from scan_region(document, region[start:end], migration, markers, offset + start) + yield from scan_region( + document, + undouble(region[start:end]).ljust(end - start), + migration, + markers, + offset + start, + ) keyword = offending_keyword(clause) if keyword is None or exempt: diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/test_litellm/test_check_migrations_no_data_rewrites.py index 280d1cb698c..f44e9d7e3aa 100644 --- a/tests/test_litellm/test_check_migrations_no_data_rewrites.py +++ b/tests/test_litellm/test_check_migrations_no_data_rewrites.py @@ -865,6 +865,19 @@ class TestDynamicSql: sql = "DO $$\nBEGIN\n EXECUTE 'SELECT ''x''; UPDATE \"Foo\" SET \"a\" = 1';\nEND $$;" assert _keywords(tmp_path, sql) == ("UPDATE",) + def test_a_comment_dash_inside_a_doubled_quote_does_not_hide_a_later_rewrite(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'SELECT ''--''; UPDATE \"Foo\" SET \"a\" = 1';\nEND $$;" + assert _keywords(tmp_path, sql) == ("UPDATE",) + assert _scan(tmp_path, sql)[0].line == 3 + + def test_a_block_comment_open_inside_a_doubled_quote_does_not_hide_a_later_rewrite(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'SELECT ''/*''; DELETE FROM \"Foo\"';\nEND $$;" + assert _keywords(tmp_path, sql) == ("DELETE",) + + def test_a_rewrite_genuinely_commented_out_inside_executed_sql_is_not_run(self, tmp_path): + sql = "DO $$\nBEGIN\n EXECUTE 'SELECT 1 -- UPDATE \"Foo\" SET \"a\" = 1';\nEND $$;" + assert _keywords(tmp_path, sql) == () + def test_a_rewrite_in_a_later_command_before_bind_values_is_flagged(self, tmp_path): sql = "DO $$\nBEGIN\n EXECUTE 'SELECT 1; DELETE FROM \"Foo\" WHERE \"a\" = $1' USING 1;\nEND $$;" assert _keywords(tmp_path, sql) == ("DELETE",) From aef742463743f833ea47da005e9ab8ddf102f997 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 24 Aug 2026 12:53:44 -0700 Subject: [PATCH 37/93] fix(azure/search): reject bool and non-positive max_results for bing_grounding count Extract a shared _valid_max_results predicate that rejects bools (an int subclass) and non-positive values, and reuse it from both the connection-mode request count and the response-side cap so both paths honor the same contract. --- litellm/llms/azure/search/transformation.py | 17 ++++++++++++----- ...est_bing_grounding_search_transformation.py | 18 ++++++++++++++++++ 2 files changed, 30 insertions(+), 5 deletions(-) diff --git a/litellm/llms/azure/search/transformation.py b/litellm/llms/azure/search/transformation.py index 7bc631e814d..0754c9b1fda 100644 --- a/litellm/llms/azure/search/transformation.py +++ b/litellm/llms/azure/search/transformation.py @@ -178,6 +178,16 @@ def _citation_results(envelope: _ResponsesEnvelope) -> tuple[SearchResult, ...]: return tuple(first_by_url[url] for url in dict.fromkeys(result.url for result in cited)) +def _valid_max_results(max_results: object) -> int | None: + """A positive-int `max_results`, else None. Rejects bools, an `int` subclass, and + non-positive values so neither the request-side `count` nor the response-side cap + forwards a value the other would silently ignore. + """ + if isinstance(max_results, bool) or not isinstance(max_results, int): + return None + return max_results if max_results > 0 else None + + def _requested_max_results(response_kwargs: Mapping[str, object]) -> int | None: """The unified `max_results` cap the caller asked for, if any. @@ -188,10 +198,7 @@ def _requested_max_results(response_kwargs: Mapping[str, object]) -> int | None: optional_params: Final = response_kwargs.get("optional_params") if not isinstance(optional_params, Mapping): return None - max_results: Final = optional_params.get("max_results") - return ( - max_results if isinstance(max_results, int) and not isinstance(max_results, bool) and max_results > 0 else None - ) + return _valid_max_results(optional_params.get("max_results")) def _capped(results: tuple[SearchResult, ...], max_results: int | None) -> tuple[SearchResult, ...]: @@ -247,7 +254,7 @@ def _search_tool(optional_params: Mapping[str, object]) -> _BingGroundingTool | if connection_id: configuration: Final = _SearchConfiguration( project_connection_id=connection_id, - count=max_results if isinstance(max_results, int) else None, + count=_valid_max_results(max_results), ) return _BingGroundingTool(bing_grounding=_BingGroundingParams(search_configurations=(configuration,))) location: Final = _UserLocation(country=country.upper()) if isinstance(country, str) else None diff --git a/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py b/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py index d33c9f92df6..fdc6f7bc239 100644 --- a/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py +++ b/tests/test_litellm/llms/azure/search/test_bing_grounding_search_transformation.py @@ -195,6 +195,24 @@ def test_transform_search_request_connection_mode_omits_count_without_max_result assert body["tools"][0]["bing_grounding"]["search_configurations"] == [{"project_connection_id": "conn-id"}] +@pytest.mark.parametrize("max_results", [True, False, 0, -1]) +def test_transform_search_request_connection_mode_omits_count_for_invalid_max_results( + monkeypatch: pytest.MonkeyPatch, max_results: object +): + monkeypatch.setenv("BING_GROUNDING_MODEL", "gpt-4.1") + monkeypatch.setenv("BING_GROUNDING_CONNECTION_ID", "conn-id") + body = _config().transform_search_request("q", {"max_results": max_results}) + assert body["tools"][0]["bing_grounding"]["search_configurations"] == [{"project_connection_id": "conn-id"}] + + +def test_transform_search_response_ignores_invalid_max_results_cap(): + annotations = [_citation(f"https://example.com/{i}", f"T{i}", 0, 5) for i in range(3)] + resp = _config().transform_search_response( + _resp(_message_response("claim", annotations)), logging_obj=Mock(), optional_params={"max_results": True} + ) + assert [r.url for r in resp.results] == [f"https://example.com/{i}" for i in range(3)] + + def test_transform_search_request_joins_list_query(monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("BING_GROUNDING_MODEL", "gpt-4.1") assert _config().transform_search_request(["foo", "bar"], {})["input"] == "foo bar" From 15d8f0b43ca6766752967cfeaec47bd9d148447e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 24 Aug 2026 12:54:57 -0700 Subject: [PATCH 38/93] refactor: drop the output-neutral pad on the executed-literal recursion The call-detection restore splices into a fixed-position list, so its .ljust(end - start) holds that length invariant and a test guards it. The DML-scan recursion instead hands the undoubled literal to a fresh scan_region as its own region, whose length feeds nothing, so the pad only appends trailing spaces that shift no keyword and change no reported line. Drop it and the docstring clause that claimed it kept the offsets landing --- .../code_coverage_tests/check_migrations_no_data_rewrites.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index 39fdfe716cf..b3f60febaba 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -633,8 +633,7 @@ def scan_region( always counted against the whole document, so a statement nested in a dollar-quoted body reports its real file line and lines up with the markers read from that file. A single-quoted literal that `DO` or `EXECUTE` runs as SQL is undoubled before it is scanned, so a `--` or `/*` - in one of its nested strings blanks nothing and the statement after it stays visible, and it is - padded back to its span so the offsets still land.""" + in one of its nested strings blanks nothing and the statement after it stays visible.""" masked, bodies, literals = mask(region) executed = executed_names(masked) runnable = executed_literals(masked, literals, executed) @@ -649,7 +648,7 @@ def scan_region( if base <= start and end <= commands_end: yield from scan_region( document, - undouble(region[start:end]).ljust(end - start), + undouble(region[start:end]), migration, markers, offset + start, From 0e96491554ea5b2bb51f1c8c79d4bb5c9718adaa Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Fri, 21 Aug 2026 15:28:22 -0700 Subject: [PATCH 39/93] feat(router): per-group supported reasoning efforts with max and ultra levels --- .../transformation.py | 21 ++--- .../llms/openai/chat/gpt_5_transformation.py | 6 +- litellm/main.py | 4 +- ...odel_prices_and_context_window_backup.json | 76 ++++++++++++----- litellm/router.py | 9 ++ .../reasoning_effort_capability.py | 53 ++++++++++++ litellm/types/llms/openai.py | 2 +- litellm/types/router.py | 1 + litellm/types/utils.py | 1 + litellm/utils.py | 1 + model_prices_and_context_window.json | 76 ++++++++++++----- model_prices_and_context_window.schema.json | 3 + ...responses_transformation_transformation.py | 13 ++- .../llms/openai/test_gpt5_transformation.py | 33 ++++++++ .../response_api_endpoints/test_endpoints.py | 4 +- .../test_reasoning_effort_capability.py | 79 +++++++++++++++++ tests/test_litellm/test_router.py | 84 +++++++++++++++++++ tests/test_litellm/test_utils.py | 1 + .../add_model/ComplexityRouterConfig.test.tsx | 41 ++++++++- .../add_model/ComplexityRouterConfig.tsx | 12 ++- .../add_model/TierModelEffortRows.tsx | 37 ++++---- .../add_model/complexity_router_tiers.ts | 12 ++- .../llm_calls/fetch_models.test.tsx | 37 +++++++- .../src/components/llm_calls/fetch_models.tsx | 45 ++++++---- ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 + 25 files changed, 551 insertions(+), 102 deletions(-) create mode 100644 litellm/router_utils/reasoning_effort_capability.py create mode 100644 tests/test_litellm/router_utils/test_reasoning_effort_capability.py diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index b94e91b3034..f94bf34e8d3 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -1113,22 +1113,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): litellm.reasoning_auto_summary or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true" ) - # If string is passed, map with optional summary based on flag/env var - if reasoning_effort == "none": - return Reasoning(effort="none", summary="detailed") if auto_summary_enabled else Reasoning(effort="none") - elif reasoning_effort == "high": - return Reasoning(effort="high", summary="detailed") if auto_summary_enabled else Reasoning(effort="high") - elif reasoning_effort == "xhigh": - return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh") - elif reasoning_effort == "medium": + # Level-agnostic: providers own effort validation, so an unknown level (max, ultra, future + # ones) passes through instead of being silently dropped here. + if reasoning_effort: return ( - Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium") - ) - elif reasoning_effort == "low": - return Reasoning(effort="low", summary="detailed") if auto_summary_enabled else Reasoning(effort="low") - elif reasoning_effort == "minimal": - return ( - Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal") + Reasoning(effort=reasoning_effort, summary="detailed") + if auto_summary_enabled + else Reasoning(effort=reasoning_effort) ) return None diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py index ffa3de0d5c6..c55640f6a1f 100644 --- a/litellm/llms/openai/chat/gpt_5_transformation.py +++ b/litellm/llms/openai/chat/gpt_5_transformation.py @@ -16,7 +16,7 @@ def _normalize_reasoning_effort_for_chat_completion( ) -> str | None: """Convert reasoning_effort to the string format expected by OpenAI chat completion API. - The chat completion API expects a simple string: 'none', 'low', 'medium', 'high', or 'xhigh'. + The chat completion API expects a simple effort string ('none' through 'ultra'). Config/deployments may pass the Responses API format: {'effort': 'high', 'summary': 'detailed'}. """ if value is None: @@ -222,8 +222,8 @@ class OpenAIGPT5Config(OpenAIGPTConfig): if "reasoning_effort" in optional_params: optional_params["reasoning_effort"] = normalized - if effective_effort == "xhigh": - # xhigh is an opt-in capability: only allow if model explicitly supports it. + if effective_effort in ("xhigh", "max", "ultra"): + # xhigh/max/ultra are opt-in capabilities: only allow if the model explicitly supports them. if not self._supports_reasoning_effort_level(model, effective_effort): if litellm.drop_params or drop_params: non_default_params.pop("reasoning_effort", None) diff --git a/litellm/main.py b/litellm/main.py index 84931c63544..1cac723c179 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -416,7 +416,7 @@ async def acompletion( logprobs: bool | None = None, top_logprobs: int | None = None, deployment_id=None, - reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh", "default"] | None = None, + reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh", "max", "ultra", "default"] | None = None, verbosity: Literal["low", "medium", "high"] | None = None, safety_identifier: str | None = None, service_tier: str | None = None, @@ -4920,7 +4920,7 @@ def completion( logit_bias: dict | None = None, user: str | None = None, # openai v1.0+ new params - reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh", "default"] | None = None, + reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh", "max", "ultra", "default"] | None = None, verbosity: Literal["low", "medium", "high"] | None = None, response_format: dict | type[BaseModel] | None = None, seed: int | None = None, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index c36018feb9d..4a275273fb2 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -6639,7 +6639,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/gpt-5.6-sol": { "cache_read_input_token_cost": 5e-07, @@ -6690,7 +6692,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/gpt-5.6-terra": { "cache_read_input_token_cost": 2e-07, @@ -6741,7 +6745,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/gpt-5.6-luna": { "cache_read_input_token_cost": 2e-08, @@ -6792,7 +6798,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/us/gpt-5.6": { "cache_read_input_token_cost": 5.5e-07, @@ -6839,7 +6847,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/us/gpt-5.6-sol": { "cache_read_input_token_cost": 5.5e-07, @@ -6887,7 +6897,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/us/gpt-5.6-terra": { "cache_read_input_token_cost": 2.2e-07, @@ -6935,7 +6947,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/us/gpt-5.6-luna": { "cache_read_input_token_cost": 2.2e-08, @@ -6983,7 +6997,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/eu/gpt-5.6": { "cache_read_input_token_cost": 5.5e-07, @@ -7030,7 +7046,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/eu/gpt-5.6-sol": { "cache_read_input_token_cost": 5.5e-07, @@ -7078,7 +7096,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/eu/gpt-5.6-terra": { "cache_read_input_token_cost": 2.2e-07, @@ -7126,7 +7146,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/eu/gpt-5.6-luna": { "cache_read_input_token_cost": 2.2e-08, @@ -7174,7 +7196,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/gpt-5.5": { "deprecation_date": "2027-10-26", @@ -26405,7 +26429,9 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": true, - "supports_xhigh_reasoning_effort": true + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "gpt-5.6-sol": { "cache_creation_input_token_cost": 5e-06, @@ -26469,7 +26495,9 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": true, - "supports_xhigh_reasoning_effort": true + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "gpt-5.6-terra": { "cache_creation_input_token_cost": 2.5e-06, @@ -26532,7 +26560,9 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": true, - "supports_xhigh_reasoning_effort": true + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "gpt-5.6-luna": { "cache_creation_input_token_cost": 2.5e-07, @@ -26595,7 +26625,9 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": true, - "supports_xhigh_reasoning_effort": true + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "gpt-5.6-cyber": { "cache_creation_input_token_cost": 1.5625e-05, @@ -49028,7 +49060,9 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "bedrock_mantle/openai.gpt-5.6-terra": { "input_cost_per_token": 2.2e-06, @@ -49060,7 +49094,9 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "bedrock_mantle/openai.gpt-5.6-luna": { "input_cost_per_token": 2.2e-07, @@ -49092,7 +49128,9 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "us.openai.gpt-5.6-sol": { "input_cost_per_token": 5.5e-06, diff --git a/litellm/router.py b/litellm/router.py index 045fd32847c..8cd1b804928 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -167,6 +167,10 @@ from litellm.router_utils.pre_call_checks.model_rate_limit_check import ( from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import ( PromptCachingDeploymentCheck, ) +from litellm.router_utils.reasoning_effort_capability import ( + intersect_supported_reasoning_efforts, + resolve_supported_reasoning_efforts, +) from litellm.router_utils.router_callbacks.track_deployment_metrics import ( increment_deployment_failures_for_current_minute, increment_deployment_successes_for_current_minute, @@ -9558,6 +9562,11 @@ class Router: if model_info.get("rpm", None) is not None and _deployment_rpm is None: _deployment_rpm = model_info.get("rpm") + model_group_info.supported_reasoning_efforts = intersect_supported_reasoning_efforts( + model_group_info.supported_reasoning_efforts, + resolve_supported_reasoning_efforts(model_info), + ) + if _deployment_tpm is not None: if total_tpm is None: total_tpm = 0 diff --git a/litellm/router_utils/reasoning_effort_capability.py b/litellm/router_utils/reasoning_effort_capability.py new file mode 100644 index 00000000000..f59c53e9287 --- /dev/null +++ b/litellm/router_utils/reasoning_effort_capability.py @@ -0,0 +1,53 @@ +"""Resolve which reasoning_effort values a deployment, and by intersection a model group, accepts. + +The model-map flags carry different polarity per level, mirroring the provider gates +(gpt_5_transformation.py restricts xhigh to explicit opt-in and treats minimal/low as opt-out; +anthropic/chat/transformation.py rejects only xhigh/max without an explicit flag): medium and high +are unconditional for any reasoning model, none/minimal/low are supported unless the map explicitly +says false, and xhigh/max require an explicit true. Shipping the resolved list keeps that polarity +in one place instead of re-encoding it in every consumer. +""" + +from collections.abc import Mapping, Sequence +from typing import Final + +REASONING_EFFORT_CAPABILITY_ORDER: Final = ("none", "minimal", "low", "medium", "high", "xhigh", "max", "ultra") + +_OPT_OUT_FLAGS: Final = ( + ("none", "supports_none_reasoning_effort"), + ("minimal", "supports_minimal_reasoning_effort"), + ("low", "supports_low_reasoning_effort"), +) +_OPT_IN_FLAGS: Final = ( + ("xhigh", "supports_xhigh_reasoning_effort"), + ("max", "supports_max_reasoning_effort"), + ("ultra", "supports_ultra_reasoning_effort"), +) +_UNCONDITIONAL_EFFORTS: Final = frozenset(("medium", "high")) + + +def resolve_supported_reasoning_efforts(model_info: Mapping[str, object]) -> tuple[str, ...] | None: + """None = no capability metadata for this deployment (e.g. a model absent from the model map, + whose stub info carries no supports_reasoning key at all); () = reasoning unsupported.""" + if "supports_reasoning" not in model_info: + return None + if model_info.get("supports_reasoning") is not True: + return () + opt_out: Final = frozenset(effort for effort, flag in _OPT_OUT_FLAGS if model_info.get(flag) is not False) + opt_in: Final = frozenset(effort for effort, flag in _OPT_IN_FLAGS if model_info.get(flag) is True) + allowed: Final = opt_out | _UNCONDITIONAL_EFFORTS | opt_in + return tuple(effort for effort in REASONING_EFFORT_CAPABILITY_ORDER if effort in allowed) + + +def intersect_supported_reasoning_efforts( + current: Sequence[str] | None, + resolved: Sequence[str] | None, +) -> tuple[str, ...] | None: + """Deployments without metadata (None) never narrow the group; an effort survives only when + every deployment with metadata accepts it, so the group offers nothing routing could reject.""" + if resolved is None: + return tuple(current) if current is not None else None + if current is None: + return tuple(resolved) + keep: Final = frozenset(current) & frozenset(resolved) + return tuple(effort for effort in REASONING_EFFORT_CAPABILITY_ORDER if effort in keep) diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index e7a3f825455..e3d28a0097f 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1840,7 +1840,7 @@ ResponsesAPIStreamingResponse = Annotated[ ] -REASONING_EFFORT = Literal["none", "minimal", "low", "medium", "high", "xhigh"] +REASONING_EFFORT = Literal["none", "minimal", "low", "medium", "high", "xhigh", "max", "ultra"] class OpenAIRealtimeStreamSession(TypedDict, total=False): diff --git a/litellm/types/router.py b/litellm/types/router.py index 9fd5cfa96ef..d4c735387a5 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -637,6 +637,7 @@ class ModelGroupInfo(BaseModel): supports_url_context: bool = Field(default=False) supports_reasoning: bool = Field(default=False) supports_function_calling: bool = Field(default=False) + supported_reasoning_efforts: tuple[str, ...] | None = Field(default=None) supported_openai_params: list[str] | None = Field(default=[]) configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 4d59650f410..6a5e4c24dce 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -164,6 +164,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False): supports_low_reasoning_effort: bool | None supports_xhigh_reasoning_effort: bool | None supports_max_reasoning_effort: bool | None + supports_ultra_reasoning_effort: bool | None # writable-ok: Pydantic warns on ReadOnly TypedDict fields supports_output_config: bool | None supports_image_size: bool | None bedrock_output_config_effort_ceiling: Literal["low", "medium", "high", "max", "xhigh"] | None diff --git a/litellm/utils.py b/litellm/utils.py index 5b2ef93edb3..a74cf23a4ae 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5764,6 +5764,7 @@ def _get_model_info_helper( supports_low_reasoning_effort=_model_info.get("supports_low_reasoning_effort", None), supports_xhigh_reasoning_effort=_model_info.get("supports_xhigh_reasoning_effort", None), supports_max_reasoning_effort=_model_info.get("supports_max_reasoning_effort", None), + supports_ultra_reasoning_effort=_model_info.get("supports_ultra_reasoning_effort", None), bedrock_output_config_effort_ceiling=_model_info.get("bedrock_output_config_effort_ceiling", None), bedrock_converse_supports_strict_tools=_model_info.get("bedrock_converse_supports_strict_tools", None), supports_computer_use=_model_info.get("supports_computer_use", None), diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index c36018feb9d..4a275273fb2 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -6639,7 +6639,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/gpt-5.6-sol": { "cache_read_input_token_cost": 5e-07, @@ -6690,7 +6692,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/gpt-5.6-terra": { "cache_read_input_token_cost": 2e-07, @@ -6741,7 +6745,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/gpt-5.6-luna": { "cache_read_input_token_cost": 2e-08, @@ -6792,7 +6798,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/us/gpt-5.6": { "cache_read_input_token_cost": 5.5e-07, @@ -6839,7 +6847,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/us/gpt-5.6-sol": { "cache_read_input_token_cost": 5.5e-07, @@ -6887,7 +6897,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/us/gpt-5.6-terra": { "cache_read_input_token_cost": 2.2e-07, @@ -6935,7 +6947,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/us/gpt-5.6-luna": { "cache_read_input_token_cost": 2.2e-08, @@ -6983,7 +6997,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/eu/gpt-5.6": { "cache_read_input_token_cost": 5.5e-07, @@ -7030,7 +7046,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/eu/gpt-5.6-sol": { "cache_read_input_token_cost": 5.5e-07, @@ -7078,7 +7096,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/eu/gpt-5.6-terra": { "cache_read_input_token_cost": 2.2e-07, @@ -7126,7 +7146,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/eu/gpt-5.6-luna": { "cache_read_input_token_cost": 2.2e-08, @@ -7174,7 +7196,9 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": false + "supports_minimal_reasoning_effort": false, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "azure/gpt-5.5": { "deprecation_date": "2027-10-26", @@ -26405,7 +26429,9 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": true, - "supports_xhigh_reasoning_effort": true + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "gpt-5.6-sol": { "cache_creation_input_token_cost": 5e-06, @@ -26469,7 +26495,9 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": true, - "supports_xhigh_reasoning_effort": true + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "gpt-5.6-terra": { "cache_creation_input_token_cost": 2.5e-06, @@ -26532,7 +26560,9 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": true, - "supports_xhigh_reasoning_effort": true + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "gpt-5.6-luna": { "cache_creation_input_token_cost": 2.5e-07, @@ -26595,7 +26625,9 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": true, - "supports_xhigh_reasoning_effort": true + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "gpt-5.6-cyber": { "cache_creation_input_token_cost": 1.5625e-05, @@ -49028,7 +49060,9 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "bedrock_mantle/openai.gpt-5.6-terra": { "input_cost_per_token": 2.2e-06, @@ -49060,7 +49094,9 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "bedrock_mantle/openai.gpt-5.6-luna": { "input_cost_per_token": 2.2e-07, @@ -49092,7 +49128,9 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_max_reasoning_effort": true, + "supports_ultra_reasoning_effort": true }, "us.openai.gpt-5.6-sol": { "input_cost_per_token": 5.5e-06, diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index f7f60c7666d..75bf3d47f35 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -702,6 +702,9 @@ "supports_tool_search": { "type": "boolean" }, + "supports_ultra_reasoning_effort": { + "type": "boolean" + }, "supports_url_context": { "type": "boolean" }, 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 1fb74b2b7bf..a0aecfafdbc 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 @@ -1585,10 +1585,15 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch): assert result_dict["summary"] == "custom_summary" print("✓ Dict input is passed through without modification") - # Test 5: None/unknown values return None - result_unknown = handler._map_reasoning_effort("unknown_value") - assert result_unknown is None - print("✓ Unknown reasoning_effort values return None") + # Test 5: levels this bridge does not enumerate (max, ultra, future ones) pass through so the + # provider can judge them, instead of being silently dropped before the request is built + from litellm.types.llms.openai import Reasoning + + for effort in ("max", "ultra", "unknown_value"): + result_passthrough = handler._map_reasoning_effort(effort) + assert result_passthrough == Reasoning(effort=effort) + assert handler._map_reasoning_effort("") is None + print("✓ Unenumerated reasoning_effort levels pass through to the provider") print( "✓ All reasoning_effort behaviors work correctly with flag/env var control" diff --git a/tests/test_litellm/llms/openai/test_gpt5_transformation.py b/tests/test_litellm/llms/openai/test_gpt5_transformation.py index d279b119efe..7595c64da07 100644 --- a/tests/test_litellm/llms/openai/test_gpt5_transformation.py +++ b/tests/test_litellm/llms/openai/test_gpt5_transformation.py @@ -1309,3 +1309,36 @@ def test_responses_gpt54_allow_temperature_effort_none( drop_params=False, ) assert params["temperature"] == 0.7 + + +@pytest.mark.parametrize("effort", ["max", "ultra"]) +def test_gpt5_6_allows_opt_in_reasoning_efforts(config: OpenAIConfig, effort: str): + params = config.map_openai_params( + non_default_params={"reasoning_effort": effort}, + optional_params={}, + model="gpt-5.6", + drop_params=False, + ) + assert params["reasoning_effort"] == effort + + +@pytest.mark.parametrize("effort", ["max", "ultra"]) +def test_gpt5_rejects_opt_in_reasoning_efforts_for_other_models(config: OpenAIConfig, effort: str): + with pytest.raises(litellm.utils.UnsupportedParamsError): + config.map_openai_params( + non_default_params={"reasoning_effort": effort}, + optional_params={}, + model="gpt-5.1", + drop_params=False, + ) + + +@pytest.mark.parametrize("effort", ["max", "ultra"]) +def test_gpt5_drops_opt_in_reasoning_efforts_when_requested(config: OpenAIConfig, effort: str): + params = config.map_openai_params( + non_default_params={"reasoning_effort": effort}, + optional_params={}, + model="gpt-5.1", + drop_params=True, + ) + assert "reasoning_effort" not in params diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index 9177944df2d..a906b0e638f 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -1352,7 +1352,9 @@ class TestParseCursorModelVariant: ("gemini-3.0-pro-thinking-low", "gemini-3.0-pro", "low"), ("claude-opus-5-fast", "claude-opus-5", None), ("gpt-5.6-sol", "gpt-5.6-sol", None), - ("foo-thinking-ultra-fast", "foo-thinking-ultra", None), + ("gpt-5.6-thinking-ultra-fast", "gpt-5.6", "ultra"), + ("gpt-5.6-thinking-max", "gpt-5.6", "max"), + ("foo-thinking-mega-fast", "foo-thinking-mega", None), ("-thinking-high", "-thinking-high", None), ], ) diff --git a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py new file mode 100644 index 00000000000..ebc7952081c --- /dev/null +++ b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py @@ -0,0 +1,79 @@ +from litellm.router_utils.reasoning_effort_capability import ( + intersect_supported_reasoning_efforts, + resolve_supported_reasoning_efforts, +) + + +class TestResolveSupportedReasoningEfforts: + def test_no_metadata_resolves_to_unknown(self): + assert resolve_supported_reasoning_efforts({}) is None + + def test_non_reasoning_model_supports_no_efforts(self): + assert resolve_supported_reasoning_efforts({"supports_reasoning": None}) == () + assert resolve_supported_reasoning_efforts({"supports_reasoning": False}) == () + + def test_reasoning_model_with_no_flags_gets_the_opt_out_levels_only(self): + # The kimi shape: supports_reasoning true, zero effort flags. medium/high are unconditional, + # none/minimal/low are opt-out so absence means supported, xhigh/max are opt-in so absence + # means unsupported. + assert resolve_supported_reasoning_efforts({"supports_reasoning": True}) == ( + "none", + "minimal", + "low", + "medium", + "high", + ) + + def test_explicit_false_removes_an_opt_out_level(self): + # The gpt-5.5-pro shape from the model map: only medium/high/xhigh are accepted upstream. + resolved = resolve_supported_reasoning_efforts( + { + "supports_reasoning": True, + "supports_none_reasoning_effort": False, + "supports_minimal_reasoning_effort": False, + "supports_low_reasoning_effort": False, + "supports_xhigh_reasoning_effort": True, + } + ) + assert resolved == ("medium", "high", "xhigh") + + def test_explicit_true_adds_the_opt_in_levels(self): + # The claude-opus shape: xhigh and max explicitly true, everything else absent. + resolved = resolve_supported_reasoning_efforts( + { + "supports_reasoning": True, + "supports_xhigh_reasoning_effort": True, + "supports_max_reasoning_effort": True, + } + ) + assert resolved == ("none", "minimal", "low", "medium", "high", "xhigh", "max") + + def test_ultra_is_opt_in(self): + without_flag = resolve_supported_reasoning_efforts({"supports_reasoning": True}) + with_flag = resolve_supported_reasoning_efforts( + {"supports_reasoning": True, "supports_ultra_reasoning_effort": True} + ) + assert without_flag is not None and "ultra" not in without_flag + assert with_flag is not None and with_flag[-1] == "ultra" + + def test_opt_in_flag_set_false_stays_excluded(self): + resolved = resolve_supported_reasoning_efforts( + {"supports_reasoning": True, "supports_xhigh_reasoning_effort": False} + ) + assert resolved is not None + assert "xhigh" not in resolved + + +class TestIntersectSupportedReasoningEfforts: + def test_unknown_never_narrows(self): + assert intersect_supported_reasoning_efforts(["medium", "high"], None) == ("medium", "high") + assert intersect_supported_reasoning_efforts(None, ["medium", "high"]) == ("medium", "high") + assert intersect_supported_reasoning_efforts(None, None) is None + + def test_intersection_keeps_canonical_order(self): + assert intersect_supported_reasoning_efforts( + ["max", "high", "medium", "xhigh"], ["xhigh", "medium", "minimal"] + ) == ("medium", "xhigh") + + def test_disjoint_sets_intersect_to_empty(self): + assert intersect_supported_reasoning_efforts(["max"], ["minimal"]) == () diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index d00fbf589e3..095c962328a 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -8878,3 +8878,87 @@ class TestAzureBaseModelFallbackLogging: deployment=None, received_model_name="my-group", id="azure-base-model-test-id" ) assert model_info["max_input_tokens"] == litellm.model_cost["azure/gpt-4o-mini"]["max_input_tokens"] + +def test_model_group_info_intersects_supported_reasoning_efforts(): + router = litellm.Router( + model_list=[ + { + "model_name": "smart-group", + "litellm_params": {"model": "anthropic/opus-like"}, + "model_info": {"id": "opus-like-deployment"}, + }, + { + "model_name": "smart-group", + "litellm_params": {"model": "openai/mini-like"}, + "model_info": {"id": "mini-like-deployment"}, + }, + ] + ) + + def _model_info(model_id: str, model_name: str): + if model_id == "opus-like-deployment": + return { + "key": model_name, + "litellm_provider": "anthropic", + "mode": "chat", + "supports_reasoning": True, + "supports_xhigh_reasoning_effort": True, + "supports_max_reasoning_effort": True, + } + return { + "key": model_name, + "litellm_provider": "openai", + "mode": "chat", + "supports_reasoning": True, + "supports_none_reasoning_effort": False, + "supports_minimal_reasoning_effort": True, + "supports_xhigh_reasoning_effort": False, + } + + with patch.object(router, "get_deployment_model_info", side_effect=_model_info): + result = router._set_model_group_info( + model_group="smart-group", + user_facing_model_group_name="smart-group", + ) + + assert result is not None + # opus-like offers all seven levels, mini-like lacks none/xhigh/max; only the common set survives, + # so the group never advertises an effort routing could hand to a deployment that rejects it. + assert result.supported_reasoning_efforts == ("minimal", "low", "medium", "high") + + +def test_model_group_info_reasoning_efforts_ignore_deployments_without_metadata(): + router = litellm.Router( + model_list=[ + { + "model_name": "smart-group", + "litellm_params": {"model": "anthropic/opus-like"}, + "model_info": {"id": "opus-like-deployment"}, + }, + { + "model_name": "smart-group", + "litellm_params": {"model": "openai/unmapped-model"}, + "model_info": {"id": "unmapped-deployment"}, + }, + ] + ) + + def _model_info(model_id: str, model_name: str): + if model_id == "opus-like-deployment": + return { + "key": model_name, + "litellm_provider": "anthropic", + "mode": "chat", + "supports_reasoning": True, + "supports_max_reasoning_effort": True, + } + return {"key": model_name, "litellm_provider": "openai", "mode": "chat"} + + with patch.object(router, "get_deployment_model_info", side_effect=_model_info): + result = router._set_model_group_info( + model_group="smart-group", + user_facing_model_group_name="smart-group", + ) + + assert result is not None + assert result.supported_reasoning_efforts == ("none", "minimal", "low", "medium", "high", "max") diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 27cf9067914..af69d117e04 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -997,6 +997,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "supports_none_reasoning_effort": {"type": "boolean"}, "supports_xhigh_reasoning_effort": {"type": "boolean"}, "supports_max_reasoning_effort": {"type": "boolean"}, + "supports_ultra_reasoning_effort": {"type": "boolean"}, "supports_adaptive_thinking": {"type": "boolean"}, "supports_legacy_thinking": {"type": "boolean"}, "thinking_always_on": {"type": "boolean"}, diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx index 0a848ed7deb..a26197aa243 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx @@ -8,7 +8,12 @@ vi.mock( ); const mockModelInfo = [ - { model_group: "gpt-4", mode: "chat", supports_reasoning: true }, + { + model_group: "gpt-4", + mode: "chat", + supports_reasoning: true, + supported_reasoning_efforts: ["medium", "high", "xhigh"], + }, { model_group: "gpt-3.5-turbo", mode: "chat" }, { model_group: "claude-3-opus", mode: "chat", supports_reasoning: true }, { model_group: "text-embedding-3-small", mode: "embedding" }, @@ -943,3 +948,37 @@ describe("ComplexityRouterConfig reasoning effort gating", () => { ).toHaveTextContent("low"); }); }); + +describe("ComplexityRouterConfig per-model effort filtering", () => { + it("offers only the efforts the model group supports", async () => { + renderWithProviders(); + const user = userEvent.setup(); + await user.click(screen.getByRole("combobox", { name: "Reasoning effort for gpt-4 in the Complex tier" })); + const options = (await screen.findAllByRole("option")).map((option) => option.textContent); + expect(options).toEqual(["Default", "medium", "high", "xhigh"]); + }); + + it("falls back to every effort when the group only reports supports_reasoning", async () => { + renderWithProviders(); + const user = userEvent.setup(); + await user.click( + screen.getByRole("combobox", { name: "Reasoning effort for claude-3-opus in the Reasoning tier" }), + ); + const options = (await screen.findAllByRole("option")).map((option) => option.textContent); + expect(options).toEqual(["Default", "none", "minimal", "low", "medium", "high", "xhigh"]); + }); + + // Hand-authored configs can carry a level outside the supported set (e.g. max); it must render + // and stay clearable rather than being masked as Default. + it("keeps showing a stored effort outside the supported set", () => { + renderWithProviders( + , + ); + expect(screen.getByRole("combobox", { name: "Reasoning effort for gpt-4 in the Complex tier" })).toHaveTextContent( + "max", + ); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index fc731e2c77f..e72905c2f73 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -13,6 +13,7 @@ import { ModelGroup } from "@/components/llm_calls/fetch_models"; import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig"; import ClassificationMethodConfig from "./ClassificationMethodConfig"; import { + REASONING_EFFORT_OPTIONS, ReasoningEffort, TierModelParamsByTier, pruneTierModelParams, @@ -251,8 +252,13 @@ const ComplexityRouterConfig: React.FC = ({ const defaultModel = resolveComplexityDefaultModel(value.tiers, value.default_model); // Embedding models can't serve a chat-completion role, so they're excluded here. - const reasoningModels = new Set( - modelInfo.filter((model) => model.supports_reasoning).map((model) => model.model_group), + // The backend list is the per-group intersection of accepted effort levels; when a proxy does not + // send it yet, fall back to the coarse supports_reasoning gate with every level offered. + const effortOptionsByModel: Record = Object.fromEntries( + modelInfo.map((model) => [ + model.model_group, + model.supported_reasoning_efforts ?? (model.supports_reasoning ? [...REASONING_EFFORT_OPTIONS] : []), + ]), ); const modelOptions = modelInfo @@ -365,7 +371,7 @@ const ComplexityRouterConfig: React.FC = ({ handleTierModelEffortChange(tier, model, effort)} /> diff --git a/ui/litellm-dashboard/src/components/add_model/TierModelEffortRows.tsx b/ui/litellm-dashboard/src/components/add_model/TierModelEffortRows.tsx index 67f583bd894..0c93b3228dc 100644 --- a/ui/litellm-dashboard/src/components/add_model/TierModelEffortRows.tsx +++ b/ui/litellm-dashboard/src/components/add_model/TierModelEffortRows.tsx @@ -2,20 +2,19 @@ import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@ import { SimpleTooltip } from "@/components/ui/tooltip"; import { Info } from "lucide-react"; import React from "react"; -import { REASONING_EFFORT_OPTIONS, ReasoningEffort, TierModelParams } from "./complexity_router_tiers"; +import { ReasoningEffort, TierModelParams } from "./complexity_router_tiers"; const PROVIDER_DEFAULT = "__provider_default__"; -const asEffort = (params: TierModelParams | undefined): ReasoningEffort | undefined => { +const storedEffort = (params: TierModelParams | undefined): ReasoningEffort | undefined => { const stored = params?.reasoning_effort; - if (typeof stored !== "string") return undefined; - return REASONING_EFFORT_OPTIONS.find((option) => option === stored); + return typeof stored === "string" && stored ? stored : undefined; }; interface TierModelEffortRowsProps { tierLabel: string; models: string[]; - reasoningModels: ReadonlySet; + effortOptionsByModel: Record; paramsByModel: Record | undefined; onEffortChange: (model: string, effort: ReasoningEffort | undefined) => void; } @@ -23,14 +22,21 @@ interface TierModelEffortRowsProps { const TierModelEffortRows: React.FC = ({ tierLabel, models, - reasoningModels, + effortOptionsByModel, paramsByModel, onEffortChange, }) => { - const shown = models.filter( - (model) => reasoningModels.has(model) || Object.keys(paramsByModel?.[model] ?? {}).length > 0, - ); - if (shown.length === 0) return null; + const rows = models + .map((model) => { + const effort = storedEffort(paramsByModel?.[model]); + const supported = effortOptionsByModel[model] ?? []; + // A stored effort outside the supported set (hand-authored, or capabilities changed since it + // was saved) stays listed so it renders and can be cleared. + const options = effort !== undefined && !supported.includes(effort) ? [...supported, effort] : supported; + return { model, effort, options }; + }) + .filter(({ model, options }) => options.length > 0 || Object.keys(paramsByModel?.[model] ?? {}).length > 0); + if (rows.length === 0) return null; return (
@@ -41,18 +47,17 @@ const TierModelEffortRows: React.FC = ({
- {shown.map((model) => ( + {rows.map(({ model, effort, options }) => (
{model} - handleClassifierContextPerTurnCharsChange(event.target.value === "" ? null : event.target.valueAsNumber) + handleClassifierContextBudgetCharsChange(event.target.value === "" ? null : event.target.valueAsNumber) } - min={1} + min={0} className="w-full" /> - Prior turns longer than this are truncated. + + Total characters of prior conversation sent to the classifier. Turns are taken newest first and quoted + whole while they fit, so a short conversation is never cut. + + {contextBudgetQuotesNothing && ( + + Under {MIN_QUOTED_CONTEXT_TURN_CHARS} characters there is no room to quote a turn that does not already + fit, so a long conversation reaches the classifier with no context at all. Set Context Window Size to 0 + to turn context off deliberately. + + )}
diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx index 0a848ed7deb..303ef3cdddc 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx @@ -108,7 +108,7 @@ describe("ComplexityRouterConfig", () => { classifier_type: "llm", classifier_llm_config: { model: "", timeout_ms: 3000, classification_rubric: "agentic" }, classifier_context_window_size: 3, - classifier_context_per_turn_chars: 200, + classifier_context_budget_chars: 8000, }; expect(onChange).toHaveBeenCalledWith(expectedValue); }); @@ -130,11 +130,10 @@ describe("ComplexityRouterConfig", () => { expect(screen.getByDisplayValue("750")).toBeInTheDocument(); expect(screen.getByText("Context Window Size")).toBeInTheDocument(); expect(screen.getByDisplayValue("5")).toBeInTheDocument(); - expect(screen.getByText("Context Per-Turn Character Limit")).toBeInTheDocument(); - expect(screen.getByDisplayValue("400")).toBeInTheDocument(); + expect(screen.queryByText("Context Per-Turn Character Limit")).not.toBeInTheDocument(); }); - it("should default classifier context fields to 3 and 200 when llm is selected without explicit values", () => { + it("should default the context window and budget when llm is selected", () => { const llmValue: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", @@ -147,8 +146,42 @@ describe("ComplexityRouterConfig", () => { const windowSizeSection = screen.getByText("Context Window Size").closest("div") as HTMLElement; expect(within(windowSizeSection).getByDisplayValue("3")).toBeInTheDocument(); - const perTurnCharsSection = screen.getByText("Context Per-Turn Character Limit").closest("div") as HTMLElement; - expect(within(perTurnCharsSection).getByDisplayValue("200")).toBeInTheDocument(); + const budgetSection = screen.getByText("Context Character Budget").closest("div") as HTMLElement; + expect(within(budgetSection).getByDisplayValue("8000")).toBeInTheDocument(); + }); + + it("should warn when the budget is too small to quote any turn that does not already fit", () => { + const llmValue: ComplexityRouterConfigValue = { + ...defaultValue, + classifier_type: "llm", + classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 }, + classifier_context_budget_chars: 50, + }; + renderWithProviders(); + + fireEvent.click(screen.getByText("Advanced: Classification Method")); + + expect(screen.getByText(/no room to quote a turn/i)).toBeInTheDocument(); + }); + + it("should not warn on a budget large enough to quote a turn, nor on a deliberate zero", () => { + for (const budget of [120, 8000, 0]) { + const { unmount } = renderWithProviders( + , + ); + fireEvent.click(screen.getByText("Advanced: Classification Method")); + expect(screen.queryByText(/no room to quote a turn/i)).not.toBeInTheDocument(); + unmount(); + } }); it("should show the assistant-turns switch with its configured value when classifier_type is llm", () => { @@ -229,26 +262,6 @@ describe("ComplexityRouterConfig", () => { }); }); - it("should call onChange with the updated classifier_context_per_turn_chars when edited", () => { - const onChange = vi.fn(); - const llmValue: ComplexityRouterConfigValue = { - ...defaultValue, - classifier_type: "llm", - classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 }, - }; - renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - - const perTurnCharsSection = screen.getByText("Context Per-Turn Character Limit").closest("div") as HTMLElement; - const input = within(perTurnCharsSection).getByRole("spinbutton"); - fireEvent.change(input, { target: { value: "500" } }); - - expect(onChange).toHaveBeenCalledWith({ - ...llmValue, - classifier_context_per_turn_chars: 500, - }); - }); - it("should render the custom technical keywords field", () => { renderWithProviders(); fireEvent.click(screen.getByText("Advanced: Classification Method")); diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index fc731e2c77f..c98dc20d3bc 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -31,7 +31,8 @@ export type { DimensionWeights, TierBoundaries, TokenThresholds }; export const DEFAULT_CLASSIFIER_TIMEOUT_MS = 3000; export const DEFAULT_TIER_DISTANCE_PENALTY = 0.5; export const DEFAULT_CLASSIFIER_CONTEXT_WINDOW_SIZE = 3; -export const DEFAULT_CLASSIFIER_CONTEXT_PER_TURN_CHARS = 200; +export const DEFAULT_CLASSIFIER_CONTEXT_BUDGET_CHARS = 8000; +export const MIN_QUOTED_CONTEXT_TURN_CHARS = 120; export const DEFAULT_SESSION_AFFINITY = false; export const DEFAULT_DEPLOYMENT_AFFINITY = true; @@ -137,6 +138,7 @@ export interface ComplexityRouterConfigValue { classifier_type: ClassifierType; classifier_llm_config?: ClassifierLLMConfig; classifier_context_window_size?: number; + classifier_context_budget_chars?: number; classifier_context_per_turn_chars?: number; classifier_context_include_assistant_turns?: boolean; classifier_fallback?: ClassifierFallback; diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index 9be5391040a..98ee2b7ae7c 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -352,7 +352,7 @@ const AddAutoRouterTab: React.FC = ({ classifierType: complexityRouterConfig.classifier_type, classifierLlmConfig: complexityRouterConfig.classifier_llm_config, classifierContextWindowSize: complexityRouterConfig.classifier_context_window_size, - classifierContextPerTurnChars: complexityRouterConfig.classifier_context_per_turn_chars, + classifierContextBudgetChars: complexityRouterConfig.classifier_context_budget_chars, classifierContextIncludeAssistantTurns: complexityRouterConfig.classifier_context_include_assistant_turns, classifierFallback: complexityRouterConfig.classifier_fallback, sessionAffinity: complexityRouterConfig.session_affinity ?? DEFAULT_SESSION_AFFINITY, diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts index cb55362c6a7..33f0fb8a539 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts @@ -23,7 +23,7 @@ const baseParams: BuildComplexityRouterConfigParams = { classifierType: "heuristic", classifierLlmConfig: undefined, classifierContextWindowSize: undefined, - classifierContextPerTurnChars: undefined, + classifierContextBudgetChars: undefined, classifierContextIncludeAssistantTurns: undefined, classifierFallback: undefined, sessionAffinity: false, @@ -93,39 +93,39 @@ describe("buildComplexityRouterConfig", () => { expect(config.classifier_llm_config).toBeUndefined(); }); - it("includes classifier_context_window_size and classifier_context_per_turn_chars only when classifier_type is llm", () => { + it("includes classifier_context_window_size and classifier_context_budget_chars only when classifier_type is llm", () => { const params: BuildComplexityRouterConfigParams = { ...baseParams, classifierType: "llm", classifierLlmConfig: { model: "gpt-4o-mini", timeout_ms: 3000 }, classifierContextWindowSize: 5, - classifierContextPerTurnChars: 300, + classifierContextBudgetChars: 4000, }; const config = buildComplexityRouterConfig(params); expect(config.classifier_context_window_size).toBe(5); - expect(config.classifier_context_per_turn_chars).toBe(300); + expect(config.classifier_context_budget_chars).toBe(4000); }); - it("omits classifier_context_window_size and classifier_context_per_turn_chars when classifier_type is heuristic even if values linger in state", () => { + it("omits classifier_context_window_size and classifier_context_budget_chars when classifier_type is heuristic even if values linger in state", () => { const params: BuildComplexityRouterConfigParams = { ...baseParams, classifierType: "heuristic", classifierContextWindowSize: 5, - classifierContextPerTurnChars: 300, + classifierContextBudgetChars: 4000, }; const config = buildComplexityRouterConfig(params); expect(config.classifier_context_window_size).toBeUndefined(); - expect(config.classifier_context_per_turn_chars).toBeUndefined(); + expect(config.classifier_context_budget_chars).toBeUndefined(); }); - it("omits classifier_context_window_size and classifier_context_per_turn_chars when classifier_type is llm but neither was set, leaving the backend default", () => { + it("omits classifier_context_window_size and classifier_context_budget_chars when classifier_type is llm but neither was set, leaving the backend default", () => { const config = buildComplexityRouterConfig({ ...baseParams, classifierType: "llm", classifierLlmConfig: { model: "gpt-4o-mini", timeout_ms: 3000 }, }); expect(config.classifier_context_window_size).toBeUndefined(); - expect(config.classifier_context_per_turn_chars).toBeUndefined(); + expect(config.classifier_context_budget_chars).toBeUndefined(); }); it("allows classifier_context_window_size of 0, distinct from unset, to send no prior-turn context", () => { diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts index 677cbe7063f..bd95bea226a 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts @@ -80,7 +80,7 @@ export interface BuildComplexityRouterConfigParams { classifierType: ClassifierType; classifierLlmConfig: ClassifierLLMConfig | undefined; classifierContextWindowSize: number | undefined; - classifierContextPerTurnChars: number | undefined; + classifierContextBudgetChars: number | undefined; classifierContextIncludeAssistantTurns: boolean | undefined; classifierFallback: ClassifierFallback | undefined; sessionAffinity: boolean; @@ -111,6 +111,7 @@ export interface ComplexityRouterConfigPayload { classifier_type: ClassifierType; classifier_llm_config?: ClassifierLLMConfig; classifier_context_window_size?: number; + classifier_context_budget_chars?: number; classifier_context_per_turn_chars?: number; classifier_context_include_assistant_turns?: boolean; classifier_fallback?: ClassifierFallback; @@ -219,7 +220,7 @@ export const buildComplexityRouterConfig = ({ classifierType, classifierLlmConfig, classifierContextWindowSize, - classifierContextPerTurnChars, + classifierContextBudgetChars, classifierContextIncludeAssistantTurns, classifierFallback, sessionAffinity, @@ -270,8 +271,8 @@ export const buildComplexityRouterConfig = ({ classifier_context_window_size: classifierContextWindowSize, }), ...(classifierType === "llm" && - classifierContextPerTurnChars !== undefined && { - classifier_context_per_turn_chars: classifierContextPerTurnChars, + classifierContextBudgetChars !== undefined && { + classifier_context_budget_chars: classifierContextBudgetChars, }), ...(classifierType === "llm" && classifierContextIncludeAssistantTurns !== undefined && { diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts index 59254dbbe6f..1a8f5f34909 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts @@ -113,18 +113,30 @@ describe("buildUpdatedComplexityRouterConfig classifier context window", () => { expect(result.classifier_context_per_turn_chars).toBe(300); }); - it("persists an edited classifier context window size and per-turn char limit", () => { + it("persists an edited classifier context window size", () => { const formValue = { tiers: STORED_LLM.tiers, classifier_type: "llm" as const, classifier_llm_config: STORED_LLM.classifier_llm_config, classifier_context_window_size: 10, - classifier_context_per_turn_chars: 500, }; const result = buildUpdatedComplexityRouterConfig(STORED_LLM, formValue); expect(result.classifier_context_window_size).toBe(10); - expect(result.classifier_context_per_turn_chars).toBe(500); + }); + + it("carries a stored per-turn cap through untouched now that no control sets it", () => { + // The modal stopped rendering a per-turn control, so the key left MANAGED_COMPLEXITY_ROUTER_KEYS. + // Had it stayed managed, every open-and-save would have silently dropped an operator's cap. + const formValue = { + tiers: STORED_LLM.tiers, + classifier_type: "llm" as const, + classifier_llm_config: STORED_LLM.classifier_llm_config, + classifier_context_window_size: 10, + }; + const result = buildUpdatedComplexityRouterConfig(STORED_LLM, formValue); + + expect(result.classifier_context_per_turn_chars).toBe(300); }); it("omits classifier context fields when classifier_type is heuristic even if values linger in state", () => { @@ -132,12 +144,12 @@ describe("buildUpdatedComplexityRouterConfig classifier context window", () => { tiers: STORED_LLM.tiers, classifier_type: "heuristic" as const, classifier_context_window_size: 5, - classifier_context_per_turn_chars: 300, + classifier_context_budget_chars: 4000, }; const result = buildUpdatedComplexityRouterConfig(STORED_LLM, formValue); expect(result.classifier_context_window_size).toBeUndefined(); - expect(result.classifier_context_per_turn_chars).toBeUndefined(); + expect(result.classifier_context_budget_chars).toBeUndefined(); }); it("does not resurrect a stale stored classifier_context_window_size once the form's own value is unset", () => { @@ -151,7 +163,6 @@ describe("buildUpdatedComplexityRouterConfig classifier context window", () => { const result = buildUpdatedComplexityRouterConfig(STORED_LLM, formValue); expect(result.classifier_context_window_size).toBeUndefined(); - expect(result.classifier_context_per_turn_chars).toBeUndefined(); }); }); diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.test.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.test.tsx index a124bd08476..6b49ebe740f 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.test.tsx +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.test.tsx @@ -245,7 +245,7 @@ describe("EditAutoRouterModal classifier context window", () => { await user.click(await screen.findByText("Advanced: Classification Method")); await screen.findByText("Context Window Size"); expect(screen.getByDisplayValue("5")).toBeInTheDocument(); - expect(screen.getByDisplayValue("300")).toBeInTheDocument(); + expect(screen.queryByText("Context Per-Turn Character Limit")).not.toBeInTheDocument(); await user.click(screen.getByRole("button", { name: /save changes/i })); diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx index 8e5317d9c1f..bce96ec76f5 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx @@ -78,7 +78,7 @@ const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([ "classifier_type", "classifier_llm_config", "classifier_context_window_size", - "classifier_context_per_turn_chars", + "classifier_context_budget_chars", "classifier_context_include_assistant_turns", "classifier_fallback", "session_affinity", @@ -175,8 +175,8 @@ export const buildUpdatedComplexityRouterConfig = ( classifier_context_window_size: value.classifier_context_window_size, }), ...(value.classifier_type === "llm" && - value.classifier_context_per_turn_chars !== undefined && { - classifier_context_per_turn_chars: value.classifier_context_per_turn_chars, + value.classifier_context_budget_chars !== undefined && { + classifier_context_budget_chars: value.classifier_context_budget_chars, }), ...(value.classifier_type === "llm" && value.classifier_context_include_assistant_turns !== undefined && { @@ -368,9 +368,9 @@ const EditAutoRouterModal: React.FC = ({ typeof parsedConfig.classifier_context_window_size === "number" ? parsedConfig.classifier_context_window_size : undefined, - classifier_context_per_turn_chars: - typeof parsedConfig.classifier_context_per_turn_chars === "number" - ? parsedConfig.classifier_context_per_turn_chars + classifier_context_budget_chars: + typeof parsedConfig.classifier_context_budget_chars === "number" + ? parsedConfig.classifier_context_budget_chars : undefined, classifier_context_include_assistant_turns: typeof parsedConfig.classifier_context_include_assistant_turns === "boolean" diff --git a/ui/litellm-dashboard/src/lib/autorouter_presets.ts b/ui/litellm-dashboard/src/lib/autorouter_presets.ts index 7ef0e06e2e6..f20bd3adb8a 100644 --- a/ui/litellm-dashboard/src/lib/autorouter_presets.ts +++ b/ui/litellm-dashboard/src/lib/autorouter_presets.ts @@ -268,6 +268,7 @@ export const buildPresetPrefill = ( model: resolve(config.classifier_llm_config.model), }, classifier_context_window_size: config.classifier_context_window_size, + classifier_context_budget_chars: config.classifier_context_budget_chars, classifier_context_per_turn_chars: config.classifier_context_per_turn_chars, classifier_context_include_assistant_turns: config.classifier_context_include_assistant_turns, session_affinity: config.session_affinity ?? DEFAULT_SESSION_AFFINITY, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 87f050e417f..78329a1e53c 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -32487,18 +32487,23 @@ export interface components { * @description Replaces the opening instructions of the LLM classifier rubric (the judging-criteria prose) for a custom tier set. The per-tier bullets and the trust-boundary paragraph telling the classifier to ignore tier requests embedded in quoted caller text are always appended after it and cannot be overridden. Requires tier_definitions; a built-in-tier router customizes its prompt via classifier_llm_config.system_prompt or classification_rubric instead. */ classification_prompt?: string | null; + /** + * Classifier Context Budget Chars + * @description Maximum characters of prior-turn text quoted to the LLM classifier, across the whole context window, per classification call. Turns are taken newest first and quoted whole while they fit, so a conversation small enough to quote entirely is never cut; once the budget runs out the older turns are dropped whole and only the turn straddling the boundary is truncated, into whatever space is left. The current ask and the caller's system prompt sit outside this budget and are always sent in full, as does the numbering each quoted turn carries. A budget under 120 leaves no room to quote a turn and suppresses the block; set classifier_context_window_size to 0 to turn context off deliberately. Only applies when classifier_type is 'llm'. + * @default 8000 + */ + classifier_context_budget_chars: number; /** * Classifier Context Include Assistant Turns - * @description Include assistant turns in the classifier context window, so difficulty stated by the model rather than by the user stays visible: a plan the assistant calls complex, which the user approves with 'yes', is classified on the work being approved instead of on the word 'yes'. When enabled, classifier_context_window_size counts the last N turns of the conversation across both roles rather than the last N user turns, and assistant text is sent to the classifier model, which may be a different deployment or provider than the routed completion model. Assistant replies share classifier_context_per_turn_chars with user turns, so raise it if replies are truncated before the part that carries the difficulty. Off by default because enabling it shifts tier decisions, and therefore spend, for an already-deployed router. Only applies when classifier_type is 'llm'. + * @description Include assistant turns in the classifier context window, so difficulty stated by the model rather than by the user stays visible: a plan the assistant calls complex, which the user approves with 'yes', is classified on the work being approved instead of on the word 'yes'. When enabled, classifier_context_window_size counts the last N turns of the conversation across both roles rather than the last N user turns, and assistant text is sent to the classifier model, which may be a different deployment or provider than the routed completion model. Assistant replies spend classifier_context_budget_chars alongside user turns, so raise it if the oldest turns stop being quoted once replies join the window. Off by default because enabling it shifts tier decisions, and therefore spend, for an already-deployed router. Only applies when classifier_type is 'llm'. * @default false */ classifier_context_include_assistant_turns: boolean; /** * Classifier Context Per Turn Chars - * @description Maximum character length for each prior turn's text in the classifier context window. Turns exceeding this are truncated. Only applies when classifier_type is 'llm'. - * @default 200 + * @description Optional cap on each individual prior turn's text, applied before classifier_context_budget_chars bounds the block. Unset by default, so one long turn may spend the whole budget, which is usually what a follow-up needs; set it when no single turn should dominate the context the classifier sees. A capped turn keeps its opening and its ending with the middle elided. Only applies when classifier_type is 'llm'. */ - classifier_context_per_turn_chars: number; + classifier_context_per_turn_chars?: number | null; /** * Classifier Context Window Size * @description Number of prior user turns (tool output and harness reminders excluded) to include as context in the LLM classifier prompt, so a follow-up like 'now do the same for the streaming path' is classified against what it refers to. Counts turns of both roles when classifier_context_include_assistant_turns is enabled. These turns are sent to the classifier model, which may be a different deployment or provider than the routed completion model; that call already carries the current user ask and the caller's system prompt in full. Set to 0 to send neither prior turns nor any conversation context beyond the current ask. Only applies when classifier_type is 'llm'. From bb27bfd9a7457e69b79ff2901f165f3a0e4c8ef0 Mon Sep 17 00:00:00 2001 From: Anmol Jaiswal <68013660+anmolg1997@users.noreply.github.com> Date: Tue, 25 Aug 2026 20:42:10 +0530 Subject: [PATCH 69/93] fix(http_handler): dispose aiohttp session when AsyncHTTPHandler is finalized without a running loop (#36670) * fix(http_handler): dispose aiohttp session when finalized without a running loop AsyncHTTPHandler.__del__ can only schedule an async close when a running event loop exists at finalization time; in any other context (worker threads whose loop has closed, sync contexts, interpreter shutdown) the RuntimeError from get_running_loop() is swallowed and the underlying aiohttp ClientSession is abandoned to GC, emitting 'Unclosed client session' / 'Unclosed connector' warnings. This is the disposal gap left after the recycle-time fix: clients created for short-lived event loops (the loop-id-keyed LLM client cache mints one handler per loop) are never recycled - they live and die with their loop, and their finalization is precisely the loop-less case. Fix: - no running loop: fall back to the connector's synchronous teardown via LiteLLMAiohttpTransport._mark_connector_closed - the same finalizer-safe path used for dead-loop recycles - honoring _owns_session so a shared session is never closed. - running loop: keep the async close, but hold a strong reference to the scheduled task until it completes (a bare create_task() result may be collected before running), mirroring _background_close_tasks. Tests: loop-less finalization closes a dead-loop session; running-loop finalization registers and drains the close task; the sync fallback respects session ownership. All three fail without the fix. * lint: conform new finalizer code to the type-discipline budget Final on the five never-rebound locals (LIT010); the class-level task registry keeps its mutable set with the sanctioned mutable-ok reason, mirroring the aiohttp transport's registry (LIT001). * lint: reasoned pyright ignore on the cross-class teardown call The handler deliberately reuses the transport's finalizer-safe connector teardown; no public seam exists and an async close can never run at loop-less finalization. Clears the net-new reportPrivateUsage the basedpyright budget gate flagged once the LIT stage passed. * fix(http_handler): retrieve exceptions from finalizer close tasks A bare discard done-callback dropped the task without consuming its exception, so a failing aclose() emitted "Task exception was never retrieved" at GC, the same noise class this path exists to remove. Mirror the transport's _on_close_task_done: discard, early-return on cancellation, retrieve and debug-log the exception. * fix(http_handler): dispose foreign-loop sessions instead of scheduling aclose on the live loop GC on a live loop (e.g. the app's) of a handler whose session belongs to another, possibly dead, loop scheduled aclose() on the current loop, the cross-loop path the transport refuses. Route both that case and the loop-less case through the transport's lifecycle-aware _close_recycled_session, which picks async close on the session's own loop, threadsafe handoff, or the synchronous connector teardown. Regression test: a dead-loop session collected while another loop runs is disposed without scheduling anything on that loop. * chore: retrigger CI (test_mcp_logging payload-order flake, also failed on litellm_spendlogs_fallback_metadata minutes earlier) * test(mcp): select the MCP tool-call payload instead of the last-delivered one TestMCPLogger kept a single last-writer slot; an async success event from another call (a mocked acompletion whose log task lands late) races the MCP event for it, so the cost assertions intermittently read the wrong payload. This PR's finalizer change shifts task interleaving on the loop and tips that latent race over (also seen on an unrelated PR minutes earlier). Collect call_type=call_mcp_tool payloads in their own list and assert on those. * test(mcp): MCPLoggerHook inherits the order-independent payload capture It duplicated TestMCPLogger's init and success handler verbatim; the hook test reads the same MCP payload selection, so subclass instead. --- litellm/llms/custom_httpx/http_handler.py | 76 ++++++++- tests/mcp_tests/test_mcp_logging.py | 48 ++++-- .../llms/custom_httpx/test_http_handler.py | 161 +++++++++++++++--- 3 files changed, 246 insertions(+), 39 deletions(-) diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 52f30e31641..777ab576de2 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -9,7 +9,7 @@ import threading import time from collections.abc import AsyncIterable, Callable, Iterable, Mapping from http.cookiejar import CookieJar, DefaultCookiePolicy -from typing import TYPE_CHECKING, Any, Final, Optional, TypeAlias, TypedDict +from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional, TypeAlias, TypedDict import certifi import httpx @@ -933,11 +933,83 @@ class AsyncHTTPHandler: response.raise_for_status() return response + # Strong references to finalizer-scheduled client-close tasks. A bare + # create_task() result may be garbage-collected before it runs, leaving + # the underlying aiohttp session unclosed ("Unclosed client session"). + # Mirrors LiteLLMAiohttpTransport._background_close_tasks. + _finalizer_close_tasks: ClassVar[set["asyncio.Task[None]"]] = set() # mutable-ok: strong refs for pending closes + + @classmethod + def _on_finalizer_close_done(cls, task: "asyncio.Task[None]") -> None: + cls._finalizer_close_tasks.discard(task) + if task.cancelled(): + return + exc: Final = task.exception() + if exc is not None: + verbose_logger.debug("Error closing client at finalization: %s", exc) + + def _aiohttp_session_bound_elsewhere(self, loop: asyncio.AbstractEventLoop) -> bool: + """True when the wrapped aiohttp session is bound to a loop other than + ``loop`` — awaiting ``aclose()`` here would touch that loop's internals.""" + from litellm.llms.custom_httpx.aiohttp_transport import ( + LiteLLMAiohttpTransport, + ) + + transport: Final = getattr(self._client, "_transport", None) + if not isinstance(transport, LiteLLMAiohttpTransport): + return False + session: Final = transport.client + if not isinstance(session, ClientSession) or session.closed: + return False + return getattr(session, "_loop", None) is not loop + + def _dispose_wrapped_aiohttp_session(self) -> None: + """Dispose the wrapped aiohttp session when ``aclose()`` cannot run here. + + Finalization either has no running loop, or a loop the session is not + bound to. Delegating to the transport's lifecycle-aware disposal picks + the safe path per session state (async close on its own loop, threadsafe + handoff to a loop running elsewhere, or the synchronous connector + teardown that flips the flags ``ClientSession.__del__`` checks), so no + "Unclosed client session" / "Unclosed connector" warnings fire at + garbage collection. + """ + from litellm.llms.custom_httpx.aiohttp_transport import ( + LiteLLMAiohttpTransport, + ) + + transport: Final = getattr(self._client, "_transport", None) + if not isinstance(transport, LiteLLMAiohttpTransport): + return + # A shared session (e.g. the proxy's) is never this handler's to close. + if not getattr(transport, "_owns_session", False): + return + session: Final = transport.client + if isinstance(session, ClientSession) and not session.closed: + transport._close_recycled_session(session) # pyright: ignore[reportPrivateUsage] # deliberate reuse of the transport's lifecycle-aware disposal; an async close can never run in this context + def __del__(self) -> None: try: if not _handler_may_close_client(sys.getrefcount(self._client), self._owns_client): return - asyncio.get_running_loop().create_task(self._client.aclose()) + try: + loop: Final = asyncio.get_running_loop() + except RuntimeError: + # No running loop at finalization time (worker threads after + # their loop closed, interpreter/worker shutdown, GC in a + # sync context). An async close can never run here. + self._dispose_wrapped_aiohttp_session() + return + if self._aiohttp_session_bound_elsewhere(loop): + # GC ran on a live loop (e.g. the app's) but the session + # belongs to another, possibly dead, loop — awaiting aclose() + # here is the cross-loop path the transport refuses. + self._dispose_wrapped_aiohttp_session() + return + task: Final = loop.create_task(self._client.aclose()) + cls: Final = type(self) + cls._finalizer_close_tasks.add(task) + task.add_done_callback(cls._on_finalizer_close_done) except Exception: pass diff --git a/tests/mcp_tests/test_mcp_logging.py b/tests/mcp_tests/test_mcp_logging.py index 7ee745b311e..1903f29001f 100644 --- a/tests/mcp_tests/test_mcp_logging.py +++ b/tests/mcp_tests/test_mcp_logging.py @@ -24,12 +24,20 @@ from mcp.types import Tool as MCPTool, CallToolResult, TextContent class TestMCPLogger(CustomLogger): def __init__(self): self.standard_logging_payload = None + self.mcp_tool_call_payloads = [] super().__init__() async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): print("success event") - self.standard_logging_payload = kwargs.get("standard_logging_object", None) - print(f"Captured standard_logging_payload: {self.standard_logging_payload}") + payload = kwargs.get("standard_logging_object", None) + self.standard_logging_payload = payload + # Async success events from other calls (e.g. a mocked acompletion whose + # log task is delivered late) race with the MCP event for the single + # last-writer slot; keep MCP tool calls in their own list so assertions + # are order-independent. + if payload is not None and payload.get("call_type") == "call_mcp_tool": + self.mcp_tool_call_payloads.append(payload) + print(f"Captured standard_logging_payload: {payload}") def _set_authorized_user(server_ids): @@ -138,7 +146,11 @@ async def test_mcp_cost_tracking(): # wait 1-2 seconds for logging to be processed await asyncio.sleep(2) - logged_standard_logging_payload = test_logger.standard_logging_payload + logged_standard_logging_payload = ( + test_logger.mcp_tool_call_payloads[-1] + if test_logger.mcp_tool_call_payloads + else None + ) print("logged_standard_logging_payload", logged_standard_logging_payload) # Add assertions @@ -277,7 +289,11 @@ async def test_mcp_cost_tracking_per_tool(): # wait for logging to be processed await asyncio.sleep(2) - logged_standard_logging_payload_1 = test_logger.standard_logging_payload + logged_standard_logging_payload_1 = ( + test_logger.mcp_tool_call_payloads[-1] + if test_logger.mcp_tool_call_payloads + else None + ) print( "logged_standard_logging_payload_1", logged_standard_logging_payload_1 ) @@ -290,6 +306,7 @@ async def test_mcp_cost_tracking_per_tool(): # Reset logger for second test test_logger.standard_logging_payload = None + test_logger.mcp_tool_call_payloads.clear() # Test 2: Call cheap_tool - should cost 0.1 response2 = await mcp_server_tool_call( @@ -300,7 +317,11 @@ async def test_mcp_cost_tracking_per_tool(): # wait for logging to be processed await asyncio.sleep(2) - logged_standard_logging_payload_2 = test_logger.standard_logging_payload + logged_standard_logging_payload_2 = ( + test_logger.mcp_tool_call_payloads[-1] + if test_logger.mcp_tool_call_payloads + else None + ) print( "logged_standard_logging_payload_2", logged_standard_logging_payload_2 ) @@ -329,16 +350,7 @@ async def test_mcp_cost_tracking_per_tool(): assert mock_client.call_tool.call_count == 2 -class MCPLoggerHook(CustomLogger): - def __init__(self): - self.standard_logging_payload = None - super().__init__() - - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - print("success event") - self.standard_logging_payload = kwargs.get("standard_logging_object", None) - print(f"Captured standard_logging_payload: {self.standard_logging_payload}") - +class MCPLoggerHook(TestMCPLogger): async def async_post_mcp_tool_call_hook( self, kwargs, response_obj: MCPPostCallResponseObject, start_time, end_time ) -> Optional[MCPPostCallResponseObject]: @@ -436,7 +448,11 @@ async def test_mcp_tool_call_hook(): await asyncio.sleep(2) # check logged standard logging payload - logged_standard_logging_payload = test_logger.standard_logging_payload + logged_standard_logging_payload = ( + test_logger.mcp_tool_call_payloads[-1] + if test_logger.mcp_tool_call_payloads + else None + ) print("logged_standard_logging_payload", logged_standard_logging_payload) assert ( logged_standard_logging_payload is not None diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index f7f89cd1d8d..16d57437043 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -56,9 +56,7 @@ async def test_async_post_streaming_status_error_should_not_wait_forever_for_bod litellm_handler = AsyncHTTPHandler() await litellm_handler.client.aclose() - litellm_handler.client = httpx.AsyncClient( - transport=httpx.MockTransport(mock_handler) - ) + litellm_handler.client = httpx.AsyncClient(transport=httpx.MockTransport(mock_handler)) try: with pytest.raises(MaskedHTTPStatusError) as exc_info: await asyncio.wait_for( @@ -202,9 +200,7 @@ async def test_ssl_verification_with_aiohttp_transport(monkeypatch: pytest.Monke transport_connector = transport._get_valid_client_session().connector assert isinstance(transport_connector, TCPConnector) - aiohttp_session = aiohttp.ClientSession( - connector=aiohttp.TCPConnector(ssl=False) - ) + aiohttp_session = aiohttp.ClientSession(connector=aiohttp.TCPConnector(ssl=False)) try: aiohttp_connector = aiohttp_session.connector assert isinstance(aiohttp_connector, aiohttp.TCPConnector) @@ -378,7 +374,8 @@ async def test_get_async_httpx_client_with_shared_session(): # Test with shared session client = get_async_httpx_client( - llm_provider=LlmProviders.ANTHROPIC, shared_session=mock_session # type: ignore + llm_provider=LlmProviders.ANTHROPIC, + shared_session=mock_session, # type: ignore ) # Verify the client was created successfully @@ -397,9 +394,7 @@ async def test_get_async_httpx_client_without_shared_session(): from litellm.types.utils import LlmProviders # Test without shared session - client = get_async_httpx_client( - llm_provider=LlmProviders.ANTHROPIC, shared_session=None - ) + client = get_async_httpx_client(llm_provider=LlmProviders.ANTHROPIC, shared_session=None) # Verify the client was created successfully assert client is not None @@ -476,11 +471,13 @@ async def test_session_reuse_integration(): # Create two clients with the same session client1 = get_async_httpx_client( - llm_provider=LlmProviders.ANTHROPIC, shared_session=mock_session # type: ignore + llm_provider=LlmProviders.ANTHROPIC, + shared_session=mock_session, # type: ignore ) client2 = get_async_httpx_client( - llm_provider=LlmProviders.OPENAI, shared_session=mock_session # type: ignore + llm_provider=LlmProviders.OPENAI, + shared_session=mock_session, # type: ignore ) # Both clients should be created successfully @@ -512,9 +509,7 @@ async def test_session_reuse_integration(): (None, None, None, False), # None value - skip configuration ], ) -def test_ssl_ecdh_curve( - env_curve, litellm_curve, expected_curve, should_call, monkeypatch -): +def test_ssl_ecdh_curve(env_curve, litellm_curve, expected_curve, should_call, monkeypatch): """Test SSL ECDH curve configuration with valid curves and precedence""" from litellm.llms.custom_httpx.http_handler import _ssl_context_cache @@ -717,9 +712,7 @@ class TestDefaultCachedClientTimeoutHonorsRequestTimeout: _default_cached_client_timeout, ) - monkeypatch.setattr( - litellm, "request_timeout", litellm.constants.DEFAULT_REQUEST_TIMEOUT_SECONDS - ) + monkeypatch.setattr(litellm, "request_timeout", litellm.constants.DEFAULT_REQUEST_TIMEOUT_SECONDS) monkeypatch.setattr(litellm, "request_timeout_explicitly_set", False) assert _default_cached_client_timeout() is _DEFAULT_TIMEOUT @@ -734,9 +727,7 @@ class TestDefaultCachedClientTimeoutHonorsRequestTimeout: assert resolved.read == 300.0 assert resolved.connect == 5.0 - def test_cached_async_client_built_with_explicit_request_timeout( - self, monkeypatch: pytest.MonkeyPatch - ): + def test_cached_async_client_built_with_explicit_request_timeout(self, monkeypatch: pytest.MonkeyPatch): from litellm.caching.llm_caching_handler import LLMClientCache from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.types.utils import LlmProviders @@ -1195,3 +1186,131 @@ async def test_aiohttp_session_never_replays_one_upstreams_cookie_to_another(): assert len(jar) == 0 assert dict(jar.filter_cookies(URL("https://upstream-a.example.com"))) == {} await session.close() + + +def _mint_session_on_dead_loop(handler: AsyncHTTPHandler) -> ClientSession: + """Create the transport's real ClientSession on a loop that then closes. + + This is the lifecycle of every client minted for a short-lived event loop + (the loop-id-keyed LLM client cache creates one handler per loop): the + session outlives its loop and can only ever be disposed loop-lessly. + """ + transport = handler.client._transport + assert isinstance(transport, LiteLLMAiohttpTransport) + loop = asyncio.new_event_loop() + + async def _create() -> ClientSession: + return transport._get_valid_client_session() + + session = loop.run_until_complete(_create()) + loop.close() + return session + + +def test_finalizer_without_running_loop_closes_dead_loop_session(): + """A handler finalized with no running event loop must still dispose its + aiohttp session. + + The async close can never run in that context; without the synchronous + fallback the session and its connector are abandoned to GC and emit + "Unclosed client session" / "Unclosed connector" warnings.""" + handler = AsyncHTTPHandler(timeout=61.0) + session = _mint_session_on_dead_loop(handler) + assert not session.closed + + del handler + gc.collect() + + assert session.closed + + +@pytest.mark.asyncio +async def test_finalizer_with_running_loop_schedules_close_and_holds_task_ref(): + """With a running loop, finalization schedules an async close and must keep + a strong reference to the task until it completes — a bare create_task() + result may be collected before it runs, leaving the session unclosed.""" + handler = AsyncHTTPHandler(timeout=61.0) + transport = handler.client._transport + assert isinstance(transport, LiteLLMAiohttpTransport) + session = transport._get_valid_client_session() + assert not session.closed + del transport + + baseline_tasks = set(AsyncHTTPHandler._finalizer_close_tasks) + del handler + gc.collect() + + scheduled = AsyncHTTPHandler._finalizer_close_tasks - baseline_tasks + assert len(scheduled) == 1 + + await asyncio.gather(*scheduled) + assert session.closed + assert not (AsyncHTTPHandler._finalizer_close_tasks & scheduled) + + +@pytest.mark.asyncio +async def test_sync_close_helper_respects_session_ownership(): + """The loop-less fallback closes only sessions the transport owns; a + shared session (e.g. the proxy's) must never be closed by a handler.""" + owned_handler = AsyncHTTPHandler(timeout=61.0) + owned_transport = owned_handler.client._transport + assert isinstance(owned_transport, LiteLLMAiohttpTransport) + owned_session = owned_transport._get_valid_client_session() + + baseline = set(LiteLLMAiohttpTransport._background_close_tasks) + owned_handler._dispose_wrapped_aiohttp_session() + scheduled = LiteLLMAiohttpTransport._background_close_tasks - baseline + await asyncio.gather(*scheduled) + assert owned_session.closed + + shared_session = ClientSession() + shared_handler = AsyncHTTPHandler(timeout=61.0, shared_session=shared_session) + shared_transport = shared_handler.client._transport + assert isinstance(shared_transport, LiteLLMAiohttpTransport) + assert shared_transport._owns_session is False + + shared_handler._dispose_wrapped_aiohttp_session() + assert not shared_session.closed + + await shared_session.close() + await shared_handler.close() + await owned_handler.close() + + +@pytest.mark.asyncio +async def test_finalizer_close_done_consumes_exception(): + """A failing finalizer close must have its exception retrieved by the done + callback, or asyncio emits "Task exception was never retrieved" at GC — + the same log noise the finalizer path exists to eliminate.""" + + async def failing_close() -> None: + raise RuntimeError("close failed") + + task = asyncio.get_running_loop().create_task(failing_close()) + AsyncHTTPHandler._finalizer_close_tasks.add(task) + await asyncio.sleep(0) + + AsyncHTTPHandler._on_finalizer_close_done(task) + assert task not in AsyncHTTPHandler._finalizer_close_tasks + + cancelled = asyncio.get_running_loop().create_task(asyncio.sleep(30)) + cancelled.cancel() + await asyncio.sleep(0) + AsyncHTTPHandler._on_finalizer_close_done(cancelled) + + +@pytest.mark.asyncio +async def test_finalizer_on_live_loop_disposes_foreign_loop_session_without_scheduling(): + """GC on a live loop (e.g. the app's) of a handler whose session belongs to + another, dead loop must not schedule aclose() here — that is the cross-loop + path the transport refuses — and must still dispose the session.""" + handler = AsyncHTTPHandler(timeout=61.0) + session = await asyncio.to_thread(_mint_session_on_dead_loop, handler) + assert not session.closed + + baseline_tasks = set(AsyncHTTPHandler._finalizer_close_tasks) + del handler + gc.collect() + + assert AsyncHTTPHandler._finalizer_close_tasks == baseline_tasks + assert session.closed From fc810484f790f283e2f0a897fa972783e707431c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 09:28:55 -0700 Subject: [PATCH 70/93] Revert "ci: raise three unit shard job timeouts to satisfy the startup safety gate" This reverts commit b71b574af410f436954f9e6fcb28b44eff6c1d34. --- .github/workflows/test-unit.yml | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 2dfca3d308f..a7c67f2b35d 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -211,7 +211,7 @@ jobs: workers: 2 reruns: 2 timeout-minutes: 20 - job-timeout-minutes: 60 + job-timeout-minutes: 55 - shard: proxy-extras artifact-name: proxy-extras @@ -219,7 +219,7 @@ jobs: workers: 2 reruns: 2 timeout-minutes: 20 - job-timeout-minutes: 60 + job-timeout-minutes: 55 - shard: enterprise-package artifact-name: enterprise-package @@ -227,7 +227,7 @@ jobs: workers: 4 reruns: 2 timeout-minutes: 20 - job-timeout-minutes: 60 + job-timeout-minutes: 55 - shard: responses-caching-types artifact-name: responses-caching-types From a73f11ae9c736059299c2078ee602bc05e43559d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 09:40:23 -0700 Subject: [PATCH 71/93] fix(completion_extras): forward reasoning_effort=max through the Responses API bridge --- .../transformation.py | 22 +++------- litellm/types/llms/openai.py | 2 +- ...responses_transformation_transformation.py | 42 ++++++++++++++++--- .../response_api_endpoints/test_endpoints.py | 2 + 4 files changed, 46 insertions(+), 22 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index b94e91b3034..17815976b4a 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -5,7 +5,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req import json import os from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast +from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast, get_args from openai.types.responses.custom_tool_param import CustomToolParam from openai.types.responses.response_input_param import ( @@ -35,6 +35,7 @@ from litellm.responses.sse_output_recovery import ( ) from litellm.responses.utils import normalize_responses_api_stream_options from litellm.types.llms.openai import ( + REASONING_EFFORT, ChatCompletionAnnotation, ChatCompletionReasoningItem, ChatCompletionToolCallChunk, @@ -1113,22 +1114,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): litellm.reasoning_auto_summary or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true" ) - # If string is passed, map with optional summary based on flag/env var - if reasoning_effort == "none": - return Reasoning(effort="none", summary="detailed") if auto_summary_enabled else Reasoning(effort="none") - elif reasoning_effort == "high": - return Reasoning(effort="high", summary="detailed") if auto_summary_enabled else Reasoning(effort="high") - elif reasoning_effort == "xhigh": - return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh") - elif reasoning_effort == "medium": + if reasoning_effort in get_args(REASONING_EFFORT): return ( - Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium") - ) - elif reasoning_effort == "low": - return Reasoning(effort="low", summary="detailed") if auto_summary_enabled else Reasoning(effort="low") - elif reasoning_effort == "minimal": - return ( - Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal") + Reasoning(effort=reasoning_effort, summary="detailed") + if auto_summary_enabled + else Reasoning(effort=reasoning_effort) ) return None diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index e7a3f825455..4a6c4a5bbb5 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1840,7 +1840,7 @@ ResponsesAPIStreamingResponse = Annotated[ ] -REASONING_EFFORT = Literal["none", "minimal", "low", "medium", "high", "xhigh"] +REASONING_EFFORT = Literal["none", "minimal", "low", "medium", "high", "xhigh", "max"] class OpenAIRealtimeStreamSession(TypedDict, total=False): 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 1fb74b2b7bf..6ca48ce63b8 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 @@ -2,7 +2,7 @@ import datetime import json import os import unittest -from typing import TYPE_CHECKING, List, Literal, Optional, Tuple +from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple from unittest.mock import ANY, MagicMock, Mock, patch import httpx @@ -1585,10 +1585,16 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch): assert result_dict["summary"] == "custom_summary" print("✓ Dict input is passed through without modification") - # Test 5: None/unknown values return None - result_unknown = handler._map_reasoning_effort("unknown_value") - assert result_unknown is None - print("✓ Unknown reasoning_effort values return None") + # Test 5: every REASONING_EFFORT level reaches the provider, and anything else (a typo, an + # unshipped level, "default") is dropped so the request still succeeds at the provider default + from litellm.types.llms.openai import Reasoning + + for effort in ("max", "xhigh", "none"): + result_passthrough = handler._map_reasoning_effort(effort) + assert result_passthrough == Reasoning(effort=effort) + for dropped in ("ultra", "hgih", "unknown_value", "", "default"): + assert handler._map_reasoning_effort(dropped) is None + print("✓ Enumerated levels pass through and unknown ones are dropped") print( "✓ All reasoning_effort behaviors work correctly with flag/env var control" @@ -2438,6 +2444,32 @@ def test_map_optional_params_preserves_reasoning_summary(): assert responses_api_request["reasoning"]["summary"] == "detailed" +@pytest.mark.parametrize("reasoning_effort", ["max", "high"]) +def test_transform_request_bedrock_mantle_tools_keeps_reasoning_effort(monkeypatch, reasoning_effort): + """Regression for reasoning_effort=max being dropped on the chat -> Responses bridge (issue #38084).""" + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + monkeypatch.setattr(litellm, "reasoning_auto_summary", False) + monkeypatch.delenv("LITELLM_REASONING_AUTO_SUMMARY", raising=False) + handler: Final = LiteLLMResponsesTransformationHandler() + + result: Final = handler.transform_request( + model="openai.gpt-5.6-sol", + messages=[{"role": "user", "content": "Say pong"}], + optional_params={ + "reasoning_effort": reasoning_effort, + "tools": [{"type": "function", "function": {"name": "get_weather", "parameters": {"type": "object"}}}], + }, + litellm_params={"custom_llm_provider": "bedrock_mantle"}, + headers={}, + litellm_logging_obj=Mock(), + ) + + assert result["reasoning"] == {"effort": reasoning_effort} + + def test_map_optional_params_tool_choice_chat_nested_to_responses_api(): """Chat tool_choice must become Responses ToolChoiceFunction (top-level name).""" from litellm.completion_extras.litellm_responses_transformation.transformation import ( diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index 9177944df2d..791d64c6428 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -1353,6 +1353,8 @@ class TestParseCursorModelVariant: ("claude-opus-5-fast", "claude-opus-5", None), ("gpt-5.6-sol", "gpt-5.6-sol", None), ("foo-thinking-ultra-fast", "foo-thinking-ultra", None), + ("gpt-5.6-thinking-max", "gpt-5.6", "max"), + ("foo-thinking-mega-fast", "foo-thinking-mega", None), ("-thinking-high", "-thinking-high", None), ], ) From 1d695a714b41d2f4ebc0cb87ea560be2d620b0f4 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Tue, 25 Aug 2026 09:50:09 -0700 Subject: [PATCH 72/93] fix(proxy): reset a stuck team member's budget (#37971) * fix(proxy): reset a stuck team member's budget A per-team-member budget check reads a cross-pod spend counter that nothing ever invalidates. Once a member exceeds their per-member budget, resetting the key's spend, raising the user's or the team's own budget, or issuing a new key all leave the member stuck, because none of them touch this counter or its cached membership object. Add POST /team/{team_id}/member/{user_id}/reset_spend to reset a member's tracked spend, and invalidate the same cached state from /team/member_update when it raises a member's own budget, so that path also takes effect immediately instead of waiting on the membership cache's TTL. Name the entity in the check's error message so a stuck member is diagnosable from the 429 body alone. * fix(proxy): close reset-vs-floor-read race and surface double Redis write failure on member spend reset Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): broadcast spend reset as a SET so the handler's self-delivered message cannot erase the reset guard Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): omit null fields from the invalidation message so plain evictions keep the old wire format Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 11 + litellm/proxy/auth/auth_checks.py | 114 +++++- litellm/proxy/auth/user_api_key_auth.py | 50 ++- .../auth_cache_invalidation_pubsub.py | 57 ++- .../proxy/common_utils/user_api_key_cache.py | 15 + .../management_endpoints/team_endpoints.py | 131 +++++- litellm/proxy/proxy_server.py | 7 + .../spend_tracking/budget_reservation.py | 10 +- .../test_team_member_reset_spend.py | 152 +++++++ .../proxy/auth/test_auth_checks.py | 305 ++++++++++++++ .../test_auth_cache_invalidation_pubsub.py | 32 ++ .../test_team_endpoints.py | 380 ++++++++++++++++++ tests/test_litellm/proxy/test_proxy_server.py | 35 ++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 63 +++ 14 files changed, 1324 insertions(+), 38 deletions(-) create mode 100644 tests/proxy_behavior/management/test_team_member_reset_spend.py diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0840d37ffa1..628e569e1b8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -815,6 +815,7 @@ class LiteLLMRoutes(enum.Enum): "/team/member_add", "/team/member_delete", "/team/member_update", + "/team/{team_id}/member/{user_id}/reset_spend", "/team/permissions_list", "/team/permissions_update", "/team/daily/activity", @@ -1287,6 +1288,16 @@ class RegenerateKeyRequest(GenerateKeyRequest): class ResetSpendRequest(LiteLLMPydanticObjectBase): reset_to: float + @field_validator("reset_to", mode="before") + @classmethod + def reject_bool_reset_to(cls, v): + # bool is a subclass of int, so pydantic silently coerces True/False into + # 1.0/0.0 for a `float` field: a caller who accidentally sends a boolean + # would otherwise get an unintended spend reset instead of a 422. + if isinstance(v, bool): + raise ValueError("reset_to must be a number, not a boolean") # noqa: TRY004 # pydantic needs ValueError + return v + class KeyRequest(LiteLLMPydanticObjectBase): keys: list[str] | None = None diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index e7b98b3cc7f..4af942a357e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -87,6 +87,8 @@ from litellm.proxy.common_utils.user_api_key_cache import ( object_permission_cache_key, tag_cache_key, tag_registry_cache_key, + team_membership_auth_cache_key, + team_membership_reservation_cache_key, ) from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.guardrails.tool_name_extraction import ( @@ -1967,7 +1969,7 @@ async def get_team_membership( if user_id is None or team_id is None: return None - _key: Final = f"team_membership:{user_id}:{team_id}" + _key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id) # check if in cache cached_membership_obj: Final = await user_api_key_cache.async_get_cache( @@ -2402,6 +2404,116 @@ async def _cache_team_object( ) +async def invalidate_team_member_spend_state( + user_id: str, + team_id: str, + user_api_key_cache: UserApiKeyCache, + new_spend: float | None = None, +) -> None: + """ + Clear every cached read path for one team member's budget so a spend + reset or a raised cap takes effect on the next request instead of + waiting on the membership cache's TTL. + + Two independently-keyed cache entries hold the same LiteLLM_TeamMembership + row: user_api_key_auth.py's admission check writes ``{team_id}_{user_id}``, + while budget_reservation.py's pre-call reservation and auth_checks.py's own + get_team_membership() (used by _check_team_member_budget) both write + ``team_membership:{user_id}:{team_id}``. Both formats must be invalidated + explicitly; writing one does not refresh the other. All keys are also + broadcast (LIT-3803): each worker's own in-memory copy (membership object, + spend counter, or the counter's own short-TTL DB-floor marker) survives + eviction elsewhere until its TTL, so the handling worker alone clearing its + copy leaves every other worker still enforcing the pre-reset budget. + + ``new_spend`` is only passed by reset_team_member_spend_fn, which knows the + exact post-reset value: it is SET everywhere (matching /key/{key}/reset_spend's + own precedent) rather than deleted, so a worker's next read reflects it + directly instead of re-deriving it through a DB reseed. team_member_update + only changes the budget cap, not the tracked spend, so it passes no + new_spend; the live spend counter is untouched in that case (deleting it + would force a reseed from the DB's own spend column, which lags the live + counter via periodic batch writes, briefly under-enforcing the raised cap + against a spend value lower than what was actually tracked) and only the + membership caches carrying the new cap are invalidated. + + The floor marker (``spend_db_floor:``, proxy_server.py's + _authoritative_floor_spend) caches the pre-reset DB spend for + SPEND_DB_FLOOR_CACHE_TTL_SECONDS; left stale after a real reset, a request + landing on the pod that cached it can read that higher floor and raise the + counter right back above the just-reset spend. It is overwritten here with + the post-reset floor (not merely deleted) and _authoritative_floor_spend + re-checks the marker after its DB read, so a floor read already in flight + on this pod when the reset commits cannot clobber it with the pre-reset + value. Both keys are broadcast as SETs carrying new_spend, not deletes: + every subscriber (remote pods AND this pod's own, which receives its own + message) writes the post-reset value, so the self-delivered message cannot + erase the guard just written here. + + Raises HTTPException(503) if Redis still holds the stale pre-reset counter + after both the SET and the fallback DELETE fail: budget checks read Redis + first, so returning success would leave the old value authoritative for + every worker despite the DB write having committed. + """ + from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( + evict_and_broadcast, + publish_auth_cache_invalidation, + ) + + if new_spend is not None: + from litellm.proxy.proxy_server import SPEND_DB_FLOOR_CACHE_TTL_SECONDS, spend_counter_cache + + spend_counter_key: Final = f"spend:team_member:{user_id}:{team_id}" + spend_db_floor_key: Final = f"spend_db_floor:{spend_counter_key}" + + spend_counter_cache.in_memory_cache.set_cache(key=spend_counter_key, value=new_spend, ttl=60) + if spend_counter_cache.redis_cache is not None: + try: + await spend_counter_cache.redis_cache.async_set_cache(key=spend_counter_key, value=new_spend, ttl=60) + except Exception as e: # noqa: BLE001 # fall back to deleting the stale entry before giving up + verbose_proxy_logger.warning( + "Failed to set spend counter %s in Redis after reset: %s; deleting it instead so the next " + "read reseeds from the DB rather than keeping the stale pre-reset value authoritative", + spend_counter_key, + e, + ) + try: + await spend_counter_cache.redis_cache.async_delete_cache(key=spend_counter_key) + except Exception: # noqa: BLE001 # stale value now authoritative in Redis; surface instead of reporting success + verbose_proxy_logger.warning( + "Failed to delete stale spend counter %s in Redis after a failed reset write", + spend_counter_key, + exc_info=True, + ) + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail={ # mutable-ok: HTTPException.detail takes a dict + "error": "Spend was reset in the database, but Redis is unreachable and still " + "holds the pre-reset counter. Retry once Redis is reachable." + }, + ) from e + + spend_counter_cache.in_memory_cache.set_cache( + key=spend_db_floor_key, + value=new_spend, + ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS, + ) + await publish_auth_cache_invalidation(cache_key=spend_counter_key, new_value=new_spend, ttl=60) + await publish_auth_cache_invalidation( + cache_key=spend_db_floor_key, + new_value=new_spend, + ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS, + ) + + await evict_and_broadcast( + cache_keys=( + team_membership_auth_cache_key(team_id=team_id, user_id=user_id), + team_membership_reservation_cache_key(user_id=user_id, team_id=team_id), + ), + user_api_key_cache=user_api_key_cache, + ) + + async def delete_cache_team_object( team_id: str, team_alias: str | None, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 658d176f6a7..28d76e6799c 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -87,7 +87,10 @@ from litellm.proxy.common_utils.http_parsing_utils import ( populate_request_with_path_params, ) from litellm.proxy.common_utils.realtime_utils import _realtime_request_body -from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.common_utils.user_api_key_cache import ( + UserApiKeyCache, + team_membership_auth_cache_key, +) from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.utils import ( @@ -1970,8 +1973,10 @@ async def _user_api_key_auth_builder( # Check 3. Check if user is in their team budget if not skip_budget_checks and valid_token.team_member_spend is not None: - if prisma_client is not None: - _cache_key: Final = f"{valid_token.team_id}_{valid_token.user_id}" + _user_id: Final = valid_token.user_id + _team_id: Final = valid_token.team_id + if prisma_client is not None and _user_id is not None and _team_id is not None: + _cache_key: Final = team_membership_auth_cache_key(team_id=_team_id, user_id=_user_id) team_member_info = await user_api_key_cache.async_get_cache( key=_cache_key, @@ -1979,25 +1984,21 @@ async def _user_api_key_auth_builder( ) if team_member_info is None: # read from DB - _user_id: Final = valid_token.user_id - _team_id: Final = valid_token.team_id - - if _user_id is not None and _team_id is not None: - _db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first( - where={ - "user_id": _user_id, - "team_id": _team_id, - }, - include={"litellm_budget_table": True}, + _db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first( + where={ + "user_id": _user_id, + "team_id": _team_id, + }, + include={"litellm_budget_table": True}, + ) + if _db_member is not None: + team_member_info = LiteLLM_TeamMembership(**_db_member.dict()) + await user_api_key_cache.async_set_cache( + key=_cache_key, + value=team_member_info, + model_type=LiteLLM_TeamMembership, + ttl=5, ) - if _db_member is not None: - team_member_info = LiteLLM_TeamMembership(**_db_member.dict()) - await user_api_key_cache.async_set_cache( - key=_cache_key, - value=team_member_info, - model_type=LiteLLM_TeamMembership, - ttl=5, - ) if team_member_info is not None and team_member_info.litellm_budget_table is not None: team_member_budget: Final = team_member_info.litellm_budget_table.max_budget @@ -2013,11 +2014,16 @@ async def _user_api_key_auth_builder( max_budget=team_member_budget, ) if team_member_spend > team_member_budget: + _entity_id: Final = f"{valid_token.user_id}:{valid_token.team_id}" raise litellm.BudgetExceededError( current_cost=team_member_spend, max_budget=team_member_budget, + message=( + f"Budget has been exceeded! TeamMember={_entity_id} " + f"Current cost: {team_member_spend}, Max budget: {team_member_budget}" + ), entity_type=Litellm_EntityType.TEAM_MEMBER.value, - entity_id=f"{valid_token.user_id}:{valid_token.team_id}", + entity_id=_entity_id, ) # Check 3. If token is expired diff --git a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py index acdc9728390..fb2ca6372c0 100644 --- a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py +++ b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py @@ -12,6 +12,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import ( ) if TYPE_CHECKING: + from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import RedisCache from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -30,15 +31,24 @@ def auth_cache_invalidation_channel(redis_cache: "RedisCache") -> str: @dataclass(frozen=True, slots=True) class _CacheInvalidationMessage: cache_key: str + new_value: float | None = None + ttl: float | None = None -def _cache_invalidation_message_json(cache_key: str) -> str: - return json.dumps(asdict(_CacheInvalidationMessage(cache_key=cache_key))) +def _cache_invalidation_message_json(cache_key: str, new_value: float | None = None, ttl: float | None = None) -> str: + message: Final = asdict(_CacheInvalidationMessage(cache_key=cache_key, new_value=new_value, ttl=ttl)) + return json.dumps({field: value for field, value in message.items() if value is not None}) -def _cache_key_from_message_data(data: object) -> str | None: +def _finite_number_or_none(value: object) -> float | None: + if isinstance(value, bool) or not isinstance(value, (int, float)): + return None + return float(value) + + +def _message_from_data(data: object) -> _CacheInvalidationMessage | None: if isinstance(data, bytes): - data = data.decode("utf-8", errors="replace") + data = data.decode("utf-8", errors="replace") # rebind-ok: normalizing the wire payload to str if not isinstance(data, str): return None try: @@ -48,14 +58,28 @@ def _cache_key_from_message_data(data: object) -> str | None: if not isinstance(parsed, dict): return None cache_key: Final = parsed.get("cache_key") - return cache_key if isinstance(cache_key, str) else None + if not isinstance(cache_key, str): + return None + return _CacheInvalidationMessage( + cache_key=cache_key, + new_value=_finite_number_or_none(parsed.get("new_value")), + ttl=_finite_number_or_none(parsed.get("ttl")), + ) -async def publish_auth_cache_invalidation(cache_key: str) -> None: +async def publish_auth_cache_invalidation( + cache_key: str, new_value: float | None = None, ttl: float | None = None +) -> None: """ Best-effort broadcast so every worker drops its local in-memory copy of a mutated management object; without this, only the handling worker and Redis are evicted and other workers keep serving the stale object until its TTL. + + Passing ``new_value`` broadcasts a SET instead of a delete: every subscriber + (including the publishing worker's own, which receives its own message) + writes the value into its additional in-memory caches rather than deleting + the key. A spend reset uses this so the handler's self-delivered message + cannot erase the freshly-written post-reset counter or floor marker. """ redis_cache: Final = coordination_redis_cache() if redis_cache is None: @@ -68,7 +92,10 @@ async def publish_auth_cache_invalidation(cache_key: str) -> None: cache_key, ) return - await client.publish(auth_cache_invalidation_channel(redis_cache), _cache_invalidation_message_json(cache_key)) + await client.publish( + auth_cache_invalidation_channel(redis_cache), + _cache_invalidation_message_json(cache_key, new_value=new_value, ttl=ttl), + ) except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors verbose_proxy_logger.warning("auth cache invalidation publish for %s failed: %s", cache_key, e) @@ -95,15 +122,17 @@ async def evict_and_broadcast(cache_keys: Sequence[str], user_api_key_cache: "Us class AuthCacheInvalidationSubscriber: - __slots__ = ("_redis_cache", "_task", "_user_api_key_cache") + __slots__ = ("_additional_in_memory_caches", "_redis_cache", "_task", "_user_api_key_cache") def __init__( self, redis_cache: "RedisCache", user_api_key_cache: "UserApiKeyCache", + additional_in_memory_caches: Sequence["InMemoryCache"] = (), ) -> None: self._redis_cache = redis_cache self._user_api_key_cache = user_api_key_cache + self._additional_in_memory_caches = tuple(additional_in_memory_caches) self._task: asyncio.Task[None] | None = None def start(self) -> None: @@ -160,12 +189,18 @@ class AuthCacheInvalidationSubscriber: def _apply_message(self, message: object) -> None: data: Final = message.get("data") if isinstance(message, dict) else None - cache_key: Final = _cache_key_from_message_data(data) - if cache_key is None: + parsed: Final = _message_from_data(data) + if parsed is None: + return + if parsed.new_value is not None: + for additional_cache in self._additional_in_memory_caches: + additional_cache.set_cache(parsed.cache_key, parsed.new_value, ttl=parsed.ttl) return in_memory_cache: Final = self._user_api_key_cache.in_memory_cache if in_memory_cache is not None: - in_memory_cache.delete_cache(cache_key) + in_memory_cache.delete_cache(parsed.cache_key) + for additional_cache in self._additional_in_memory_caches: + additional_cache.delete_cache(parsed.cache_key) @staticmethod async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None: diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index 93d51bdd461..b8df0105b7b 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -200,6 +200,21 @@ def end_user_restricted_registry_cache_key() -> str: return "end_user_restricted_registry" +def team_membership_auth_cache_key(team_id: str, user_id: str) -> str: + """Cache key one team member's ``LiteLLM_TeamMembership`` row is stored under for the admission check.""" + return f"{team_id}_{user_id}" + + +def team_membership_reservation_cache_key(user_id: str, team_id: str) -> str: + """Cache key the pre-call budget reservation stores the same ``LiteLLM_TeamMembership`` row under. + + Deliberately not unified with ``team_membership_auth_cache_key``: the two readers wrote independent + keys before this file existed, so a fix that invalidates one must invalidate both explicitly rather + than assume a single write is visible to both. + """ + return f"team_membership:{user_id}:{team_id}" + + def get_management_object_ttl(cache: DualCache) -> float: """ In-memory TTL for management-object cache writes (keys, teams, users, budgets, ...). diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 01254d5c064..49461d7841d 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -16,7 +16,7 @@ import traceback from collections.abc import Mapping, Sequence from datetime import datetime, timezone from types import MappingProxyType -from typing import Annotated, Final, NamedTuple, Protocol, TypedDict, TypeVar, cast +from typing import Annotated, Final, NamedTuple, NoReturn, Protocol, TypedDict, TypeVar, cast import fastapi from fastapi import APIRouter, Depends, Header, HTTPException, Request, status @@ -56,6 +56,7 @@ from litellm.proxy._types import ( PatchTeamRequest, ProxyErrorTypes, ProxyException, + ResetSpendRequest, SpecialManagementEndpointEnums, SpecialModelNames, SpecialProxyStrings, @@ -84,6 +85,7 @@ from litellm.proxy.auth.auth_checks import ( get_team_membership, get_team_object, get_user_object, + invalidate_team_member_spend_state, ) from litellm.proxy.auth.auth_utils import ( enforce_batch_enqueued_token_limit_is_admin_only, @@ -3392,7 +3394,7 @@ async def team_member_update( Update team member budgets and team member role """ - from litellm.proxy.proxy_server import premium_user, prisma_client + from litellm.proxy.proxy_server import premium_user, prisma_client, user_api_key_cache if prisma_client is None: raise HTTPException(status_code=500, detail={"error": "No db connected"}) @@ -3491,6 +3493,12 @@ async def team_member_update( budget_patch=budget_patch, team_default_budget_id=team_default_budget_id, ) + if budget_patch: + await invalidate_team_member_spend_state( + user_id=received_user_id, + team_id=data.team_id, + user_api_key_cache=user_api_key_cache, + ) ### update team member role if data.role is not None: @@ -3527,6 +3535,125 @@ async def team_member_update( ) +def _check_not_resetting_own_spend(user_id: str, user_api_key_dict: UserAPIKeyAuth) -> None: + """ + _verify_team_access authorizes a team admin (or org admin) over their own + team, with no check that the target user_id differs from the caller. Left + unchecked, that admin could target their own LiteLLM_TeamMembership row and + repeatedly reset it to 0 right before it crosses their per-member cap, + consuming the shared team budget without the configured limit ever binding. + Only a proxy admin may reset an admin's own spend. + """ + if user_id == user_api_key_dict.user_id and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + _raise_reset_spend_error(status.HTTP_403_FORBIDDEN, "Cannot reset your own spend. Ask a proxy admin.") + + +def _raise_reset_spend_error(status_code: int, message: str) -> NoReturn: + detail: Final = {"error": message} # mutable-ok: HTTPException.detail takes a dict + raise HTTPException(status_code=status_code, detail=detail) + + +def _validate_team_member_reset_spend_value( + reset_to: object, + membership: LiteLLM_TeamMembership, +) -> float: + if not isinstance(reset_to, (int, float)): + _raise_reset_spend_error(status.HTTP_400_BAD_REQUEST, "reset_to must be a float") + + reset_to_float: Final = float(reset_to) + if not math.isfinite(reset_to_float) or reset_to_float < 0: + _raise_reset_spend_error(status.HTTP_400_BAD_REQUEST, "reset_to must be a finite number >= 0") + + current_spend: Final = membership.spend or 0.0 + if reset_to_float > current_spend: + _raise_reset_spend_error( + status.HTTP_400_BAD_REQUEST, + f"reset_to ({reset_to_float}) must be <= current spend ({current_spend})", + ) + + max_budget: Final = membership.litellm_budget_table.max_budget if membership.litellm_budget_table else None + if max_budget is not None and reset_to_float > max_budget: + _raise_reset_spend_error( + status.HTTP_400_BAD_REQUEST, + f"reset_to ({reset_to_float}) must be <= budget ({max_budget})", + ) + + return reset_to_float + + +@router.post( + "/team/{team_id}/member/{user_id}/reset_spend", + tags=["team management"], # mutable-ok: FastAPI's `tags` param is typed as list[str], not Sequence + dependencies=(Depends(user_api_key_auth),), +) +@management_endpoint_wrapper +async def reset_team_member_spend_fn( + team_id: str, + user_id: str, + data: ResetSpendRequest, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +): + """ + Reset a team member's tracked spend against their per-member budget. + + A member's spend is tracked separately from both their own personal + budget and the team's own budget (LiteLLM_TeamMembership.spend), so + neither /user/update nor /team/update can clear it: this is the only + endpoint that does. The cross-pod spend counter and cached membership + reads are invalidated so the reset takes effect on the member's next + request rather than waiting on the membership cache's TTL. + """ + from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache + + if prisma_client is None: + _raise_reset_spend_error(status.HTTP_500_INTERNAL_SERVER_ERROR, "DB not connected. prisma_client is None") + + team_obj: Final = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + check_db_only=True, + ) + await _verify_team_access(team_obj=team_obj, user_api_key_dict=user_api_key_dict) + _check_not_resetting_own_spend(user_id=user_id, user_api_key_dict=user_api_key_dict) + + membership_where: Final = { # mutable-ok: prisma client requires a plain dict where= argument + "user_id_team_id": {"user_id": user_id, "team_id": team_id} # mutable-ok: same prisma where= argument + } + _membership_row: Final = await _team_membership_db(prisma_client).find_unique( + where=membership_where, + include={"litellm_budget_table": True}, # mutable-ok: prisma client requires a plain dict include= argument + ) + if _membership_row is None: + _raise_reset_spend_error(status.HTTP_404_NOT_FOUND, f"User {user_id} is not a member of team {team_id}.") + membership: Final = LiteLLM_TeamMembership.model_validate(_membership_row.model_dump()) + + current_spend: Final = membership.spend or 0.0 + reset_to: Final = _validate_team_member_reset_spend_value(data.reset_to, membership) + + await _team_membership_db(prisma_client).update( + where=membership_where, + data={"spend": reset_to}, # mutable-ok: prisma client requires a plain dict data= argument + ) + + await invalidate_team_member_spend_state( + user_id=user_id, + team_id=team_id, + user_api_key_cache=user_api_key_cache, + new_spend=reset_to, + ) + + return { # mutable-ok: matches this router's established untyped-response-dict convention + "team_id": team_id, + "user_id": user_id, + "spend": reset_to, + "previous_spend": current_spend, + "max_budget": membership.litellm_budget_table.max_budget if membership.litellm_budget_table else None, + } + + def _create_results_from_response( members: list[Member], response: TeamAddMemberResponse, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7dced4e26b6..0abcdeaf3f6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2555,6 +2555,12 @@ async def _authoritative_floor_spend( if db_spend is None: return None + # a spend reset that committed during the DB read above wrote the post-reset + # floor to the marker; keep it over this read's now-stale pre-commit value + rechecked: Final = spend_counter_cache.in_memory_cache.get_cache(key=marker_key) + if rechecked is not None: + return float(rechecked) + spend_counter_cache.in_memory_cache.set_cache( key=marker_key, value=db_spend, @@ -6798,6 +6804,7 @@ class ProxyConfig: subscriber: Final = AuthCacheInvalidationSubscriber( redis_cache=redis_cache, user_api_key_cache=user_api_key_cache, + additional_in_memory_caches=(spend_counter_cache.in_memory_cache,), ) self.auth_cache_invalidation_subscriber = subscriber subscriber.start() diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index ce6c9330620..149f9b960a1 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -25,7 +25,11 @@ from litellm.proxy._types import ( from litellm.proxy.auth.auth_utils import get_model_from_request from litellm.proxy.auth.budget_throttle import should_throttle_budget_exceeded from litellm.proxy.auth.route_checks import RouteChecks -from litellm.proxy.common_utils.user_api_key_cache import end_user_cache_key, tag_cache_key +from litellm.proxy.common_utils.user_api_key_cache import ( + end_user_cache_key, + tag_cache_key, + team_membership_reservation_cache_key, +) from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.router import Router @@ -546,7 +550,9 @@ async def _get_team_member_budget_counter( if team_object is None or team_object.team_id is None or user_object is None or valid_token.user_id is None: return None - membership_cache_key: Final = f"team_membership:{valid_token.user_id}:{team_object.team_id}" + membership_cache_key: Final = team_membership_reservation_cache_key( + user_id=valid_token.user_id, team_id=team_object.team_id + ) cached_team_membership: Final = await user_api_key_cache.async_get_cache(key=membership_cache_key) team_membership: LiteLLM_TeamMembership | None = None if isinstance(cached_team_membership, LiteLLM_TeamMembership): diff --git a/tests/proxy_behavior/management/test_team_member_reset_spend.py b/tests/proxy_behavior/management/test_team_member_reset_spend.py new file mode 100644 index 00000000000..ec2c78139fe --- /dev/null +++ b/tests/proxy_behavior/management/test_team_member_reset_spend.py @@ -0,0 +1,152 @@ +import uuid + +import pytest + +from .actors import Actor +from .conftest import create_scratch_team + +pytestmark = pytest.mark.asyncio(loop_scope="session") + +_SEED_SPEND = 5.0 +_RESET_TO = 2.0 + + +# POST /team/{team_id}/member/{user_id}/reset_spend. The handler gate is +# _verify_team_access (proxy admin / team admin of this team / org admin of +# the team's org) — the same gate /team/member_update uses, so this mirrors +# that file's matrix exactly. +_MATRIX = [ + ("alpha/proxy_admin", Actor.PROXY_ADMIN, "alpha", 200), + ("alpha/org_admin", Actor.ORG_ADMIN, "alpha", 200), + ("alpha/team_admin", Actor.TEAM_ADMIN, "alpha", 200), + ("alpha/internal_user", Actor.INTERNAL_USER, "alpha", 403), + ("alpha/owner", Actor.OWNER, "alpha", 403), + ("alpha/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "alpha", 403), + ("alpha/cross_org_user", Actor.CROSS_ORG_USER, "alpha", 403), + ("alpha/service_account", Actor.SERVICE_ACCOUNT, "alpha", 403), + ("alpha/org_b_admin", Actor.ORG_B_ADMIN, "alpha", 403), + ("beta/proxy_admin", Actor.PROXY_ADMIN, "beta", 200), + ("beta/org_admin", Actor.ORG_ADMIN, "beta", 403), + ("beta/team_admin", Actor.TEAM_ADMIN, "beta", 403), + ("beta/org_b_admin", Actor.ORG_B_ADMIN, "beta", 200), +] + + +async def _seed_target(prisma, world, shape: str, team_id: str, member_id: str) -> None: + if shape == "alpha": + await create_scratch_team( + prisma, + team_id, + organization_id=world.org_a_id, + admin_user_ids=[world.keys[Actor.TEAM_ADMIN].user_id], + ) + elif shape == "beta": + await create_scratch_team(prisma, team_id, organization_id=world.org_b_id) + else: # pragma: no cover - guard + pytest.fail(f"unknown shape={shape}") + await prisma.db.litellm_teammembership.create( + data={"user_id": member_id, "team_id": team_id, "spend": _SEED_SPEND} + ) + + +@pytest.mark.parametrize( + "actor,shape,expected_status", + [(a, sh, s) for (_id, a, sh, s) in _MATRIX], + ids=[s[0] for s in _MATRIX], +) +async def test_team_member_reset_spend_authz_matrix( + actor: Actor, + shape: str, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + member_id = scratch.tag("member") + await _seed_target(prisma, world, shape, scratch.prefix, member_id) + caller = world.keys[actor] + + resp = await proxy_client.post( + f"/team/{scratch.prefix}/member/{member_id}/reset_spend", + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"reset_to": _RESET_TO}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value} {shape}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_teammembership.find_unique( + where={"user_id_team_id": {"user_id": member_id, "team_id": scratch.prefix}} + ) + assert row is not None + if expected_status == 200: + assert row.spend == _RESET_TO + else: + assert row.spend == _SEED_SPEND, "denied but spend reset" + + +async def test_team_member_reset_spend_missing_team_is_404(proxy_client, world): + resp = await proxy_client.post( + f"/team/behavior-pin-no-such-team/member/{uuid.uuid4().hex}/reset_spend", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"reset_to": 0.0}, + ) + assert resp.status_code == 404, resp.text + + +async def test_team_member_reset_spend_missing_membership_is_404( + proxy_client, prisma, scratch, world +): + """A well-formed team but a user_id with no LiteLLM_TeamMembership row is 404.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + resp = await proxy_client.post( + f"/team/{scratch.prefix}/member/{uuid.uuid4().hex}/reset_spend", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"reset_to": 0.0}, + ) + assert resp.status_code == 404, resp.text + + +async def test_team_member_reset_spend_above_current_spend_is_400( + proxy_client, prisma, scratch, world +): + member_id = scratch.tag("member") + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + await prisma.db.litellm_teammembership.create( + data={"user_id": member_id, "team_id": scratch.prefix, "spend": 1.0} + ) + resp = await proxy_client.post( + f"/team/{scratch.prefix}/member/{member_id}/reset_spend", + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"reset_to": 5.0}, + ) + assert resp.status_code == 400, resp.text + + +async def test_team_member_reset_spend_team_admin_cannot_reset_own_spend( + proxy_client, prisma, scratch, world +): + """A team admin targeting their own LiteLLM_TeamMembership row is 403: unchecked, an + admin could repeatedly zero their own spend right before it crosses their per-member + cap, consuming the shared team budget without the configured limit ever binding.""" + team_admin = world.keys[Actor.TEAM_ADMIN] + await create_scratch_team( + prisma, + scratch.prefix, + organization_id=world.org_a_id, + admin_user_ids=[team_admin.user_id], + ) + await prisma.db.litellm_teammembership.create( + data={"user_id": team_admin.user_id, "team_id": scratch.prefix, "spend": _SEED_SPEND} + ) + resp = await proxy_client.post( + f"/team/{scratch.prefix}/member/{team_admin.user_id}/reset_spend", + headers={"Authorization": f"Bearer {team_admin.cleartext}"}, + json={"reset_to": 0.0}, + ) + assert resp.status_code == 403, resp.text + row = await prisma.db.litellm_teammembership.find_unique( + where={"user_id_team_id": {"user_id": team_admin.user_id, "team_id": scratch.prefix}} + ) + assert row is not None and row.spend == _SEED_SPEND, "denied but spend reset" diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 04f38b5e2ed..abe73f7d05c 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -47,6 +47,7 @@ from litellm.proxy.auth.auth_checks import ( _virtual_key_soft_budget_check, get_key_object, get_user_object, + invalidate_team_member_spend_state, vector_store_access_check, ) from litellm.caching.in_memory_cache import InMemoryCache @@ -6939,3 +6940,307 @@ def test_model_has_no_cost_mapping_alias_to_a_group_priced_through_model_info_is router = _router_with_a_group_priced_through_model_info() assert model_has_no_cost_mapping(model="model-info-priced-alias", llm_router=router) is False + + +@pytest.mark.asyncio +async def test_invalidate_team_member_spend_state_sets_the_spend_counter_and_clears_both_membership_cache_keys(): + """A team-member budget reset (new_spend passed) must SET the spend counter to the reset + value, clear its DB-floor marker, AND invalidate both independently-keyed membership caches + (user_api_key_auth.py's admission check writes one key format, budget_reservation.py and + auth_checks.py's own get_team_membership() write the other) or a stale read keeps 429ing + after the reset. Asserted against real cache reads, not mock call args, so a change that + keeps the call but drops its effect still fails.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + real_cache = UserApiKeyCache() + await real_cache.async_set_cache(key="team-1_user-1", value="stale-membership") + await real_cache.async_set_cache(key="team_membership:user-1:team-1", value="stale-membership") + + real_spend_counter_cache = DualCache() + real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:user-1:team-1", value=999.0) + real_spend_counter_cache.in_memory_cache.set_cache( + key="spend_db_floor:spend:team_member:user-1:team-1", value=999.0 + ) + + with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock + "litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache + ): + await invalidate_team_member_spend_state( + user_id="user-1", + team_id="team-1", + user_api_key_cache=real_cache, + new_spend=0.0, + ) + + assert await real_cache.async_get_cache(key="team-1_user-1") is None + assert await real_cache.async_get_cache(key="team_membership:user-1:team-1") is None + assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == 0.0 + assert ( + real_spend_counter_cache.in_memory_cache.get_cache(key="spend_db_floor:spend:team_member:user-1:team-1") + == 0.0 + ), "the DB-floor marker kept the pre-reset value; a stale-floor read can raise the counter right back up" + + +@pytest.mark.asyncio +async def test_invalidate_team_member_spend_state_leaves_the_live_spend_counter_alone_without_new_spend(): + """team_member_update only changes the budget cap, not the tracked spend, so it calls + invalidate_team_member_spend_state with no new_spend. Deleting the live spend counter in that + case would force the next read to reseed from the DB's own spend column, which lags the live + counter via periodic batch writes, briefly UNDER-enforcing the raised cap against a spend + value lower than what was actually tracked (regression: PR #37971 Bugbot finding).""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + real_cache = UserApiKeyCache() + await real_cache.async_set_cache(key="team-1_user-1", value="stale-membership") + + real_spend_counter_cache = DualCache() + real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:user-1:team-1", value=999.0) + + with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock + "litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache + ): + await invalidate_team_member_spend_state( + user_id="user-1", + team_id="team-1", + user_api_key_cache=real_cache, + ) + + assert await real_cache.async_get_cache(key="team-1_user-1") is None + assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == 999.0 + + +@pytest.mark.asyncio +async def test_invalidate_team_member_spend_state_sets_new_spend_instead_of_deleting(): + """/key/{key}/reset_spend SETs its counter to the reset value rather than deleting it, so a + worker's next read reflects it directly instead of falling back through a DB reseed. A reset + caller passing new_spend must match that precedent, not merely delete the counter.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + real_cache = UserApiKeyCache() + real_spend_counter_cache = DualCache() + fake_redis_cache = MagicMock() + fake_redis_cache.async_set_cache = AsyncMock() + real_spend_counter_cache.redis_cache = fake_redis_cache + + with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock + "litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache + ): + await invalidate_team_member_spend_state( + user_id="user-1", + team_id="team-1", + user_api_key_cache=real_cache, + new_spend=2.5, + ) + + assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == 2.5 + fake_redis_cache.async_set_cache.assert_awaited_once_with(key="spend:team_member:user-1:team-1", value=2.5, ttl=60) + + +@pytest.mark.asyncio +async def test_invalidate_team_member_spend_state_deletes_redis_counter_when_set_fails(): # test-quality-ok: only observable effect is the fallback call on the same fake client + """Redis reads take priority over the local in-memory copy (get_current_spend reads Redis + first), so a failed Redis SET would otherwise leave the OLD pre-reset value authoritative + for every worker even though the reset reported success. On a failed SET, the stale Redis + entry must be deleted instead, so the next read clean-misses and reseeds from the DB.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + real_cache = UserApiKeyCache() + real_spend_counter_cache = DualCache() + fake_redis_cache = MagicMock() + fake_redis_cache.async_set_cache = AsyncMock(side_effect=ConnectionError("redis down")) + fake_redis_cache.async_delete_cache = AsyncMock() + real_spend_counter_cache.redis_cache = fake_redis_cache + + with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock + "litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache + ): + await invalidate_team_member_spend_state( + user_id="user-1", + team_id="team-1", + user_api_key_cache=real_cache, + new_spend=2.5, + ) + + fake_redis_cache.async_delete_cache.assert_awaited_once_with(key="spend:team_member:user-1:team-1") + + +@pytest.mark.asyncio +async def test_invalidate_team_member_spend_state_raises_503_when_both_redis_writes_fail(): + """If the Redis SET fails AND the fallback DELETE fails, the stale pre-reset counter is still + authoritative in Redis for every worker. Reporting success would silently keep 429ing the + member, so the reset must surface a 503 instead (regression: PR #37971 Greptile finding).""" + from fastapi import HTTPException + + from litellm.caching.dual_cache import DualCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + real_cache = UserApiKeyCache() + real_spend_counter_cache = DualCache() + fake_redis_cache = MagicMock() + fake_redis_cache.async_set_cache = AsyncMock(side_effect=ConnectionError("redis down")) + fake_redis_cache.async_delete_cache = AsyncMock(side_effect=ConnectionError("redis still down")) + real_spend_counter_cache.redis_cache = fake_redis_cache + + with ( + patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock + "litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache + ), + pytest.raises(HTTPException) as exc_info, + ): + await invalidate_team_member_spend_state( + user_id="user-1", + team_id="team-1", + user_api_key_cache=real_cache, + new_spend=2.5, + ) + + assert exc_info.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_invalidate_team_member_spend_state_broadcasts_the_spend_counter_to_remote_workers(): + """The test above only proves the handling worker's own spend counter is + cleared. A remote worker's spend counter is a separate DualCache instance; + if the reset never reaches it, that worker keeps enforcing the pre-reset + spend the moment its own Redis read for the counter fails and it falls + back to its own (now-stale) in-memory copy. Drives the actual message + published onto the invalidation channel through a second, independent + AuthCacheInvalidationSubscriber standing in for that remote worker, rather + than asserting on the publish call args.""" + from redis.asyncio import Redis + + from litellm.caching.dual_cache import DualCache + from litellm.caching.in_memory_cache import InMemoryCache + from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + published: list[tuple[str, str]] = [] + + class _RecordingRedisClient(Redis): + def __init__(self) -> None: + pass + + async def publish(self, channel: str, message: str) -> int: + published.append((channel, message)) + return 1 + + class _FakeRedisCache: + def __init__(self) -> None: + self.namespace = None + + def init_async_client(self) -> object: + return _RecordingRedisClient() + + local_spend_counter_cache = DualCache() + + remote_user_api_key_cache = UserApiKeyCache() + remote_spend_counter_in_memory_cache = InMemoryCache() + remote_spend_counter_in_memory_cache.set_cache("spend:team_member:user-1:team-1", 999.0) + remote_spend_counter_in_memory_cache.set_cache("spend_db_floor:spend:team_member:user-1:team-1", 999.0) + + with ( + patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock + "litellm.proxy.proxy_server.spend_counter_cache", local_spend_counter_cache + ), + patch( # test-quality-ok: injects a fake pub/sub-capable redis cache; no live redis in this unit test + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache", + return_value=_FakeRedisCache(), + ), + ): + await invalidate_team_member_spend_state( + user_id="user-1", + team_id="team-1", + user_api_key_cache=UserApiKeyCache(), + new_spend=0.0, + ) + + def _published_message_for(cache_key: str) -> str: + matches = [message for _, message in published if json.loads(message)["cache_key"] == cache_key] + assert matches, f"{cache_key} never reached the cross-worker invalidation channel" + return matches[-1] + + remote_subscriber = AuthCacheInvalidationSubscriber( + redis_cache=_FakeRedisCache(), + user_api_key_cache=remote_user_api_key_cache, + additional_in_memory_caches=(remote_spend_counter_in_memory_cache,), + ) + for cache_key in ("spend:team_member:user-1:team-1", "spend_db_floor:spend:team_member:user-1:team-1"): + remote_subscriber._apply_message( # pyright: ignore[reportPrivateUsage] # exercising the real cross-worker message handler, not a public API + {"type": "message", "data": _published_message_for(cache_key)} + ) + + assert remote_spend_counter_in_memory_cache.get_cache("spend:team_member:user-1:team-1") == 0.0 + assert ( + remote_spend_counter_in_memory_cache.get_cache("spend_db_floor:spend:team_member:user-1:team-1") == 0.0 + ), "the DB-floor marker was not broadcast; a remote worker can re-raise the counter off its stale floor" + + +@pytest.mark.asyncio +async def test_invalidate_team_member_spend_state_self_delivered_broadcast_does_not_erase_the_reset(): + """The handling worker subscribes to the same invalidation channel it publishes on, so it + receives its own reset message. A delete-style broadcast would erase the post-reset counter + and floor marker the handler just wrote, reopening the stale-floor race the reset closed + (regression: PR #37971 Greptile finding). The broadcast carries the reset value as a SET, so + applying the self-delivered message must leave both keys at the post-reset value.""" + from redis.asyncio import Redis + + from litellm.caching.dual_cache import DualCache + from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + published: list[tuple[str, str]] = [] + + class _RecordingRedisClient(Redis): + def __init__(self) -> None: + pass + + async def publish(self, channel: str, message: str) -> int: + published.append((channel, message)) + return 1 + + class _FakeRedisCache: + def __init__(self) -> None: + self.namespace = None + + def init_async_client(self) -> object: + return _RecordingRedisClient() + + local_spend_counter_cache = DualCache() + local_user_api_key_cache = UserApiKeyCache() + + with ( + patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock + "litellm.proxy.proxy_server.spend_counter_cache", local_spend_counter_cache + ), + patch( # test-quality-ok: injects a fake pub/sub-capable redis cache; no live redis in this unit test + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache", + return_value=_FakeRedisCache(), + ), + ): + await invalidate_team_member_spend_state( + user_id="user-1", + team_id="team-1", + user_api_key_cache=local_user_api_key_cache, + new_spend=0.0, + ) + + own_subscriber = AuthCacheInvalidationSubscriber( + redis_cache=_FakeRedisCache(), + user_api_key_cache=local_user_api_key_cache, + additional_in_memory_caches=(local_spend_counter_cache.in_memory_cache,), + ) + for _, message in published: + own_subscriber._apply_message( # pyright: ignore[reportPrivateUsage] # exercising the real cross-worker message handler, not a public API + {"type": "message", "data": message} + ) + + assert local_spend_counter_cache.in_memory_cache.get_cache("spend:team_member:user-1:team-1") == 0.0, ( + "the handler's self-delivered broadcast erased the post-reset spend counter" + ) + assert ( + local_spend_counter_cache.in_memory_cache.get_cache("spend_db_floor:spend:team_member:user-1:team-1") == 0.0 + ), "the handler's self-delivered broadcast erased the post-reset floor marker, reopening the stale-floor race" diff --git a/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py b/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py index 468e8aabae8..7d5fc1a3544 100644 --- a/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py +++ b/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py @@ -6,6 +6,7 @@ from unittest.mock import patch import pytest from redis.asyncio import Redis +from litellm.caching.in_memory_cache import InMemoryCache from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( AUTH_CACHE_INVALIDATION_CHANNEL, AuthCacheInvalidationSubscriber, @@ -144,6 +145,37 @@ async def test_subscriber_deletes_local_cache_entry_on_message() -> None: assert pubsub.subscribed_channels == [AUTH_CACHE_INVALIDATION_CHANNEL] +@pytest.mark.asyncio +async def test_subscriber_deletes_additional_in_memory_cache_entry_on_message() -> None: + """ + The spend-counter half of the same cross-worker gap: a remote worker's own + spend counter can hold a stale value (its fallback path when that worker's + own Redis read for the counter fails), and only clearing user_api_key_cache + on message would leave that separate DualCache's in-memory copy untouched. + """ + cache = UserApiKeyCache() + spend_counter_in_memory_cache = InMemoryCache() + spend_counter_in_memory_cache.set_cache("spend:team_member:u-1:t-1", 999.0) + assert spend_counter_in_memory_cache.get_cache("spend:team_member:u-1:t-1") is not None + + pubsub = _QueuePubSub(initial_messages=[_invalidation_message("spend:team_member:u-1:t-1")]) + subscriber = AuthCacheInvalidationSubscriber( + redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[pubsub])), + user_api_key_cache=cache, + additional_in_memory_caches=(spend_counter_in_memory_cache,), + ) + subscriber.start() + try: + for _ in range(200): + if spend_counter_in_memory_cache.get_cache("spend:team_member:u-1:t-1") is None: + break + await asyncio.sleep(0.01) + finally: + await subscriber.stop() + + assert spend_counter_in_memory_cache.get_cache("spend:team_member:u-1:t-1") is None + + @pytest.mark.asyncio async def test_subscriber_ignores_malformed_messages() -> None: cache = UserApiKeyCache() 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 f6d74a189bc..7f5d3eb0a14 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -9,11 +9,13 @@ from unittest.mock import AsyncMock, MagicMock, call, patch import pytest from fastapi import HTTPException from fastapi.testclient import TestClient +from pydantic import ValidationError from litellm._uuid import uuid from litellm.proxy._types import UserAPIKeyAuth # Import UserAPIKeyAuth from litellm.proxy._types import ( + LiteLLM_BudgetTable, LiteLLM_BudgetTableFull, LiteLLM_ModelTable, LiteLLM_OrganizationMembershipTable, @@ -27,7 +29,9 @@ from litellm.proxy._types import ( Member, ProxyErrorTypes, ProxyException, + ResetSpendRequest, TeamMemberAddRequest, + TeamMemberUpdateRequest, UpdateTeamRequest, ) from litellm.proxy.management_endpoints.team_endpoints import ( @@ -42,12 +46,15 @@ from litellm.proxy.management_endpoints.team_endpoints import ( _transform_teams_to_deleted_records, _update_model_table, _validate_and_populate_member_user_info, + _validate_team_member_reset_spend_value, _verify_team_access, delete_team, list_available_teams, + reset_team_member_spend_fn, router, team_member_add_duplication_check, team_member_delete, + team_member_update, update_team, validate_team_org_change, ) @@ -12603,3 +12610,376 @@ async def test_invalidate_access_group_cache_deletes_the_cached_object(): "user_api_key_cache": cache, "proxy_logging_obj": logging_obj, } + + +def test_validate_team_member_reset_spend_value_rejects_non_numeric(): + with pytest.raises(HTTPException) as exc: + _validate_team_member_reset_spend_value( + reset_to="not-a-number", + membership=LiteLLM_TeamMembership(user_id="u1", team_id="t1", spend=10.0), + ) + assert exc.value.status_code == 400 + + +def test_validate_team_member_reset_spend_value_rejects_negative(): + with pytest.raises(HTTPException) as exc: + _validate_team_member_reset_spend_value( + reset_to=-1.0, + membership=LiteLLM_TeamMembership(user_id="u1", team_id="t1", spend=10.0), + ) + assert exc.value.status_code == 400 + + +@pytest.mark.parametrize("reset_to", [float("nan"), float("inf"), float("-inf")]) +def test_validate_team_member_reset_spend_value_rejects_non_finite(reset_to): + """NaN and +/-inf are instances of float and compare False against every bound + below (`nan < 0`, `nan > current_spend` are both False), so an isinstance-and-range + check alone lets them through to persist as the member's spend and silently + disable every later budget comparison against it.""" + with pytest.raises(HTTPException) as exc: + _validate_team_member_reset_spend_value( + reset_to=reset_to, + membership=LiteLLM_TeamMembership(user_id="u1", team_id="t1", spend=10.0), + ) + assert exc.value.status_code == 400 + + +@pytest.mark.parametrize("reset_to", [True, False]) +def test_reset_spend_request_rejects_bool_reset_to(reset_to): + """bool is a subclass of int, so pydantic silently coerces True/False into 1.0/0.0 for a + ``float`` field: {"reset_to": true} would otherwise reach _validate_team_member_reset_spend_value + as an indistinguishable 1.0 and reset the member's spend instead of failing the request.""" + with pytest.raises(ValidationError): + ResetSpendRequest(reset_to=reset_to) + + +def test_validate_team_member_reset_spend_value_rejects_above_current_spend(): + with pytest.raises(HTTPException) as exc: + _validate_team_member_reset_spend_value( + reset_to=20.0, + membership=LiteLLM_TeamMembership(user_id="u1", team_id="t1", spend=10.0), + ) + assert exc.value.status_code == 400 + + +def test_validate_team_member_reset_spend_value_rejects_above_max_budget(): + with pytest.raises(HTTPException) as exc: + _validate_team_member_reset_spend_value( + reset_to=10.0, + membership=LiteLLM_TeamMembership( + user_id="u1", + team_id="t1", + spend=10.0, + litellm_budget_table=LiteLLM_BudgetTable(budget_id="b1", max_budget=5.0), + ), + ) + assert exc.value.status_code == 400 + + +def test_validate_team_member_reset_spend_value_accepts_valid_reset(): + result = _validate_team_member_reset_spend_value( + reset_to=0.0, + membership=LiteLLM_TeamMembership(user_id="u1", team_id="t1", spend=10.0), + ) + assert result == 0.0 + + +@pytest.mark.asyncio +async def test_reset_team_member_spend_fn_success(monkeypatch): + """A proxy admin resetting a stuck team member's spend must write the DB + row to reset_to AND invalidate the cached spend/membership state, or the + 429 the endpoint exists to clear keeps firing off the stale cache. + Asserted against real cache reads, not mock call args, so a change that + keeps the call but drops its effect still fails.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + mock_prisma_client = MagicMock() + mock_proxy_logging_obj = MagicMock() + real_cache = UserApiKeyCache() + await real_cache.async_set_cache(key="team-1_member-1", value="stale-membership") + await real_cache.async_set_cache(key="team_membership:member-1:team-1", value="stale-membership") + real_spend_counter_cache = DualCache() + real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:member-1:team-1", value=999.0) + + membership_row = LiteLLM_TeamMembership( + user_id="member-1", + team_id="team-1", + spend=10.0, + litellm_budget_table=LiteLLM_BudgetTable(budget_id="b1", max_budget=50.0), + ) + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row) + mock_prisma_client.db.litellm_teammembership.update = AsyncMock(return_value=membership_row) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", real_cache) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj) + monkeypatch.setattr("litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache) + + with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints.get_team_object", + AsyncMock(return_value=LiteLLM_TeamTable(team_id="team-1")), + ): + response = await reset_team_member_spend_fn( + team_id="team-1", + user_id="member-1", + data=ResetSpendRequest(reset_to=0.0), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + + assert response["spend"] == 0.0 + assert response["previous_spend"] == 10.0 + assert response["max_budget"] == 50.0 + mock_prisma_client.db.litellm_teammembership.update.assert_awaited_once_with( + where={"user_id_team_id": {"user_id": "member-1", "team_id": "team-1"}}, + data={"spend": 0.0}, + ) + assert await real_cache.async_get_cache(key="team-1_member-1") is None + assert await real_cache.async_get_cache(key="team_membership:member-1:team-1") is None + assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:member-1:team-1") == 0.0 + + +@pytest.mark.asyncio +async def test_reset_team_member_spend_fn_membership_not_found(monkeypatch): + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + + with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints.get_team_object", + AsyncMock(return_value=LiteLLM_TeamTable(team_id="team-1")), + ): + with pytest.raises(HTTPException) as exc: + await reset_team_member_spend_fn( + team_id="team-1", + user_id="ghost-user", + data=ResetSpendRequest(reset_to=0.0), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + assert exc.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_reset_team_member_spend_fn_team_not_found(monkeypatch): + mock_prisma_client = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + + with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints.get_team_object", + AsyncMock(side_effect=HTTPException(status_code=404, detail={"error": "Team doesn't exist in db."})), + ): + with pytest.raises(HTTPException) as exc: + await reset_team_member_spend_fn( + team_id="ghost-team", + user_id="member-1", + data=ResetSpendRequest(reset_to=0.0), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + assert exc.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_reset_team_member_spend_fn_forbidden_for_non_admin(monkeypatch): + """A caller who is neither proxy admin, org admin, nor this team's admin must be refused, + matching every other team-mutating endpoint's authorization.""" + mock_prisma_client = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + + with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints.get_team_object", + AsyncMock(return_value=LiteLLM_TeamTable(team_id="team-1", members_with_roles=[])), + ): + with pytest.raises(HTTPException) as exc: + await reset_team_member_spend_fn( + team_id="team-1", + user_id="member-1", + data=ResetSpendRequest(reset_to=0.0), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-user", user_id="plain-user" + ), + ) + assert exc.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_reset_team_member_spend_fn_team_admin_cannot_reset_own_spend(monkeypatch): + """_verify_team_access authorizes a team admin over their own team with no check that the + target differs from the caller. Unchecked, that admin could target their own membership row + and repeatedly zero it right before it crosses their per-member cap, consuming the shared + team budget without the configured limit ever binding (Veria finding on PR #37971).""" + mock_prisma_client = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + + team_admin = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-admin", user_id="team-admin-1") + with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints.get_team_object", + AsyncMock( + return_value=LiteLLM_TeamTable( + team_id="team-1", + members_with_roles=[Member(user_id="team-admin-1", role="admin")], + ) + ), + ): + with pytest.raises(HTTPException) as exc: + await reset_team_member_spend_fn( + team_id="team-1", + user_id="team-admin-1", + data=ResetSpendRequest(reset_to=0.0), + user_api_key_dict=team_admin, + ) + assert exc.value.status_code == 403 + mock_prisma_client.db.litellm_teammembership.update.assert_not_called() + + +@pytest.mark.asyncio +async def test_reset_team_member_spend_fn_proxy_admin_can_reset_own_spend(monkeypatch): + """The self-reset guard is scoped to non-proxy-admin roles: a proxy admin resetting their + own membership spend is the platform-wide trust boundary, not a team-scoped one.""" + mock_prisma_client = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + + membership_row = LiteLLM_TeamMembership(user_id="admin-user", team_id="team-1", spend=10.0) + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row) + mock_prisma_client.db.litellm_teammembership.update = AsyncMock(return_value=membership_row) + + with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints.get_team_object", + AsyncMock(return_value=LiteLLM_TeamTable(team_id="team-1")), + ): + response = await reset_team_member_spend_fn( + team_id="team-1", + user_id="admin-user", + data=ResetSpendRequest(reset_to=0.0), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + assert response["spend"] == 0.0 + + +@pytest.mark.asyncio +async def test_team_member_update_invalidates_team_member_spend_state_when_budget_patch_applied(monkeypatch): + """Raising a stuck member's max_budget_in_team via the documented /team/member_update + endpoint must invalidate the cached membership state, or the raised cap never reaches the + admission check and the member stays 429ing. The live spend counter itself must be left + untouched: only the cap changed, and deleting the counter would force a reseed from the + DB's own spend column, which lags the live counter via periodic batch writes, briefly + UNDER-enforcing the raised cap against a spend value lower than what was actually tracked. + Asserted against real cache reads, not mock call args, so a change that keeps the call but + drops its effect still fails.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + mock_prisma_client = MagicMock() + real_cache = UserApiKeyCache() + await real_cache.async_set_cache(key="team-1_member-1", value="stale-membership") + await real_cache.async_set_cache(key="team_membership:member-1:team-1", value="stale-membership") + real_spend_counter_cache = DualCache() + real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:member-1:team-1", value=999.0) + + team_row = LiteLLM_TeamTable(team_id="team-1", metadata={}, members_with_roles=[]) + team_info_response = { + "team_info": team_row, + "team_memberships": [LiteLLM_TeamMembership(user_id="member-1", team_id="team-1", budget_id=None)], + } + + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", real_cache) + monkeypatch.setattr("litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache) + + mock_tx = AsyncMock() + mock_prisma_client.tx.return_value.__aenter__ = AsyncMock(return_value=mock_tx) + mock_prisma_client.tx.return_value.__aexit__ = AsyncMock(return_value=None) + + with ( + patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints.team_info", + AsyncMock(return_value=team_info_response), + ), + patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + AsyncMock(), + ), + ): + await team_member_update( + data=TeamMemberUpdateRequest(team_id="team-1", user_id="member-1", max_budget_in_team=999999.0), + http_request=MagicMock(), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + + assert await real_cache.async_get_cache(key="team-1_member-1") is None + assert await real_cache.async_get_cache(key="team_membership:member-1:team-1") is None + assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:member-1:team-1") == 999.0 + + +@pytest.mark.asyncio +async def test_team_member_update_skips_invalidation_when_no_budget_fields_sent(monkeypatch): + """A role-only update carries an empty budget_patch and touches no budget state, + so the member's cached spend/membership state must be left untouched.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + mock_prisma_client = MagicMock() + real_cache = UserApiKeyCache() + await real_cache.async_set_cache(key="team-1_member-1", value="still-fresh-membership") + real_spend_counter_cache = DualCache() + real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:member-1:team-1", value=1.5) + + team_row = LiteLLM_TeamTable(team_id="team-1", metadata={}, members_with_roles=[]) + team_info_response = { + "team_info": team_row, + "team_memberships": [LiteLLM_TeamMembership(user_id="member-1", team_id="team-1", budget_id=None)], + } + + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", real_cache) + monkeypatch.setattr("litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache) + + mock_tx = AsyncMock() + mock_prisma_client.tx.return_value.__aenter__ = AsyncMock(return_value=mock_tx) + mock_prisma_client.tx.return_value.__aexit__ = AsyncMock(return_value=None) + + with ( + patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints.team_info", + AsyncMock(return_value=team_info_response), + ), + patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + AsyncMock(), + ), + ): + await team_member_update( + data=TeamMemberUpdateRequest(team_id="team-1", user_id="member-1"), + http_request=MagicMock(), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + + assert await real_cache.async_get_cache(key="team-1_member-1") == "still-fresh-membership" + assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:member-1:team-1") == 1.5 diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 31d2a6cef98..3383527e932 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -11347,3 +11347,38 @@ class TestRouterModelNameOnStreamingChunks: assert len(frames) >= 3 assert '"router_model_name":"deep-model"' in frames[0] assert all('"router_model_name":"backup-tier"' in frame for frame in frames[1:]) + + +@pytest.mark.asyncio +async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the_db_read(): + """A team-member spend reset writes the post-reset floor to the spend_db_floor marker + (auth_checks.invalidate_team_member_spend_state). A floor read already in flight when the + reset commits would otherwise cache its stale pre-reset DB value over the fresh marker, + letting a budget check raise the counter right back above the just-reset spend + (regression: PR #37971 Greptile finding).""" + from litellm.proxy.proxy_server import _authoritative_floor_spend + + real_spend_counter_cache = DualCache() + counter_key = "spend:team_member:user-1:team-1" + marker_key = f"spend_db_floor:{counter_key}" + + async def db_read_racing_with_a_reset(prisma_client, counter_key): + real_spend_counter_cache.in_memory_cache.set_cache(key=marker_key, value=0.0) + return 999.0 + + with ( + patch.object( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock + proxy_server_module, "spend_counter_cache", real_spend_counter_cache + ), + patch.object( # test-quality-ok: the DB read must race the reset; no injectable seam for module-global prisma reads + proxy_server_module.SpendCounterReseed, + "from_db", + AsyncMock(side_effect=db_read_racing_with_a_reset), + ), + ): + result = await _authoritative_floor_spend(counter_key=counter_key) + + assert result == 0.0 + assert real_spend_counter_cache.in_memory_cache.get_cache(key=marker_key) == 0.0, ( + "the in-flight DB read clobbered the post-reset floor marker with the stale pre-reset value" + ) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 78329a1e53c..d51bff27784 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -14936,6 +14936,33 @@ export interface paths { patch?: never; trace?: never; }; + "/team/{team_id}/member/{user_id}/reset_spend": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** + * Reset Team Member Spend Fn + * @description Reset a team member's tracked spend against their per-member budget. + * + * A member's spend is tracked separately from both their own personal + * budget and the team's own budget (LiteLLM_TeamMembership.spend), so + * neither /user/update nor /team/update can clear it: this is the only + * endpoint that does. The cross-pod spend counter and cached membership + * reads are invalidated so the reset takes effect on the member's next + * request rather than waiting on the membership cache's TTL. + */ + post: operations["reset_team_member_spend_fn_team__team_id__member__user_id__reset_spend_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/team/{team_id}/members/me": { parameters: { query?: never; @@ -55083,6 +55110,42 @@ export interface operations { }; }; }; + reset_team_member_spend_fn_team__team_id__member__user_id__reset_spend_post: { + parameters: { + query?: never; + header?: never; + path: { + team_id: string; + user_id: string; + }; + cookie?: never; + }; + requestBody: { + content: { + "application/json": components["schemas"]["ResetSpendRequest"]; + }; + }; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; team_member_me_team__team_id__members_me_get: { parameters: { query?: never; From 9224b2ce5d3e0327bcc0b02e6f111794574716d6 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 09:50:23 -0700 Subject: [PATCH 73/93] fix(router): freeze reasoning effort flag mappings --- .../router_utils/reasoning_effort_capability.py | 16 ++++++++++------ 1 file changed, 10 insertions(+), 6 deletions(-) diff --git a/litellm/router_utils/reasoning_effort_capability.py b/litellm/router_utils/reasoning_effort_capability.py index d20b1151803..3e4478e5e21 100644 --- a/litellm/router_utils/reasoning_effort_capability.py +++ b/litellm/router_utils/reasoning_effort_capability.py @@ -25,12 +25,14 @@ it above. """ from collections.abc import Mapping, Sequence +from types import MappingProxyType from typing import Final, get_args import litellm from litellm.types.llms.openai import REASONING_EFFORT REASONING_EFFORT_ADVERTISEMENT_ORDER: Final = get_args(REASONING_EFFORT) +_EMPTY_ENTRY: Final[Mapping[str, object]] = MappingProxyType({}) _EFFORT_FLAGS: Final = ( ("none", "supports_none_reasoning_effort"), @@ -52,17 +54,19 @@ def _bare_model_entry(model_info: Mapping[str, object]) -> Mapping[str, object]: key: Final = model_info.get("key") provider: Final = model_info.get("litellm_provider") if not isinstance(key, str) or not isinstance(provider, str) or not key.startswith(f"{provider}/"): - return {} + return _EMPTY_ENTRY entry: Final[Mapping[str, object] | None] = litellm.model_cost.get(key.removeprefix(f"{provider}/")) - return entry if entry is not None else {} + return entry if entry is not None else _EMPTY_ENTRY def _declared_effort_flags(model_info: Mapping[str, object]) -> Mapping[str, object]: bare: Final = _bare_model_entry(model_info) - return { - effort: model_info.get(flag) if model_info.get(flag) is not None else bare.get(flag) - for effort, flag in _EFFORT_FLAGS - } + return MappingProxyType( + { + effort: model_info.get(flag) if model_info.get(flag) is not None else bare.get(flag) + for effort, flag in _EFFORT_FLAGS + } + ) def _supports_none_reasoning_effort(model_info: Mapping[str, object], flag: object) -> bool: From f583151a5b8928361237e715abe75b305fc4b3a5 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 09:53:19 -0700 Subject: [PATCH 74/93] fix(model_prices): raise bedrock_mantle gpt-5.6 max_input_tokens to Mantle's enforced 1050000 --- ...odel_prices_and_context_window_backup.json | 9 ++-- model_prices_and_context_window.json | 9 ++-- .../llm_cost_calc/test_llm_cost_calc_utils.py | 4 +- ...bedrock_mantle_responses_transformation.py | 44 ++++++++++++++++++- 4 files changed, 56 insertions(+), 10 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 9f953e11df1..1aa4c7cd060 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -49016,12 +49016,13 @@ "output_cost_per_token": 3.3e-05, "output_cost_per_token_above_272k_tokens": 4.95e-05, "litellm_provider": "bedrock_mantle", - "max_input_tokens": 1000000, + "max_input_tokens": 1050000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", "use_openai_responses_path": true, "supported_endpoints": [ + "/v1/chat/completions", "/v1/responses" ], "supported_modalities": [ @@ -49048,12 +49049,13 @@ "output_cost_per_token": 1.32e-05, "output_cost_per_token_above_272k_tokens": 1.98e-05, "litellm_provider": "bedrock_mantle", - "max_input_tokens": 1000000, + "max_input_tokens": 1050000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", "use_openai_responses_path": true, "supported_endpoints": [ + "/v1/chat/completions", "/v1/responses" ], "supported_modalities": [ @@ -49080,12 +49082,13 @@ "output_cost_per_token": 1.32e-06, "output_cost_per_token_above_272k_tokens": 1.98e-06, "litellm_provider": "bedrock_mantle", - "max_input_tokens": 1000000, + "max_input_tokens": 1050000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", "use_openai_responses_path": true, "supported_endpoints": [ + "/v1/chat/completions", "/v1/responses" ], "supported_modalities": [ diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 9f953e11df1..1aa4c7cd060 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -49016,12 +49016,13 @@ "output_cost_per_token": 3.3e-05, "output_cost_per_token_above_272k_tokens": 4.95e-05, "litellm_provider": "bedrock_mantle", - "max_input_tokens": 1000000, + "max_input_tokens": 1050000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", "use_openai_responses_path": true, "supported_endpoints": [ + "/v1/chat/completions", "/v1/responses" ], "supported_modalities": [ @@ -49048,12 +49049,13 @@ "output_cost_per_token": 1.32e-05, "output_cost_per_token_above_272k_tokens": 1.98e-05, "litellm_provider": "bedrock_mantle", - "max_input_tokens": 1000000, + "max_input_tokens": 1050000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", "use_openai_responses_path": true, "supported_endpoints": [ + "/v1/chat/completions", "/v1/responses" ], "supported_modalities": [ @@ -49080,12 +49082,13 @@ "output_cost_per_token": 1.32e-06, "output_cost_per_token_above_272k_tokens": 1.98e-06, "litellm_provider": "bedrock_mantle", - "max_input_tokens": 1000000, + "max_input_tokens": 1050000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", "use_openai_responses_path": true, "supported_endpoints": [ + "/v1/chat/completions", "/v1/responses" ], "supported_modalities": [ diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index c8c36032793..6f513ce1bd4 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -478,10 +478,10 @@ def test_generic_cost_per_token_minimax_m3_above_512k_tokens(_local_model_cost_m ], ) def test_generic_cost_per_token_bedrock_mantle_gpt56_long_context(_local_model_cost_map, model): - """Bedrock GPT-5.6 supports a 1M context window, billed at the long-context rates above 272K.""" + """Bedrock GPT-5.6 enforces a 1,050,000-token context window, billed at the long-context rates above 272K.""" model_cost_map = litellm.model_cost[model] - assert model_cost_map["max_input_tokens"] == 1000000 + assert model_cost_map["max_input_tokens"] == 1050000 cached_tokens = 100000 completion_tokens = 1000 diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py index 9e05d48a18f..fd279a2bc1f 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py @@ -8,7 +8,8 @@ gate, the URL construction for both paths, and the shared Bearer auth. """ import copy - +import json +from pathlib import Path import pytest from botocore.exceptions import ( @@ -1523,7 +1524,7 @@ class TestBedrockMantleResponsesPricing: assert info["cache_creation_input_token_cost"] == pytest.approx(cache_creation_cost) assert info["cache_read_input_token_cost"] == pytest.approx(cache_read_cost) assert info["output_cost_per_token"] == pytest.approx(output_cost) - assert info["max_input_tokens"] == 1000000 + assert info["max_input_tokens"] == 1050000 assert info["input_cost_per_token_above_272k_tokens"] == pytest.approx(input_cost * 2) assert info["cache_creation_input_token_cost_above_272k_tokens"] == pytest.approx(cache_creation_cost * 2) assert info["cache_read_input_token_cost_above_272k_tokens"] == pytest.approx(cache_read_cost * 2) @@ -1565,3 +1566,42 @@ class TestBedrockMantleResponsesPricing: def test_models_registered(self, local_cost_map): assert "bedrock_mantle/openai.gpt-5.5" in litellm.bedrock_mantle_models assert "bedrock_mantle/openai.gpt-5.4" in litellm.bedrock_mantle_models + + +def _repo_cost_map(map_name: str) -> dict: + repo_root = Path(__file__).resolve().parents[4] + paths = { + "root": repo_root / "model_prices_and_context_window.json", + "bundled_backup": repo_root / "litellm" / "model_prices_and_context_window_backup.json", + } + return json.loads(paths[map_name].read_text()) + + +class TestGpt56MantleRegistryEntries: + """Locks the gpt-5.6 frontier entries to Bedrock Mantle's live behavior. + + Mantle enforces a 1,050,000-token prompt maximum for gpt-5.6 sol/terra/luna + (oversize requests 400 with "prompt tokens (N) exceed model maximum + (1050000)", and a 1,030,590-token request completes), matching the OpenAI + Bedrock guide. mode must stay "responses": Mantle's native + /v1/chat/completions rejects function tools unless reasoning_effort is + "none", so chat traffic has to keep bridging to the Responses API + (see the responses_api_bridge tests above). + """ + + @pytest.mark.parametrize("map_name", ("root", "bundled_backup")) + @pytest.mark.parametrize( + "key", + ( + "bedrock_mantle/openai.gpt-5.6-sol", + "bedrock_mantle/openai.gpt-5.6-terra", + "bedrock_mantle/openai.gpt-5.6-luna", + ), + ) + def test_entry_matches_mantle_enforced_limits(self, map_name, key): + entry = _repo_cost_map(map_name)[key] + assert entry["max_input_tokens"] == 1050000 + assert entry["max_output_tokens"] == 128000 + assert entry["mode"] == "responses" + assert entry["use_openai_responses_path"] is True + assert entry["supported_endpoints"] == ["/v1/chat/completions", "/v1/responses"] From 530dab32b9f5d308fa628586b35bc97eae95a670 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 09:55:13 -0700 Subject: [PATCH 75/93] feat(vertex_ai): add native Vertex AI Interactions API support --- litellm/__init__.py | 3 + litellm/_lazy_imports_registry.py | 5 + litellm/interactions/utils.py | 7 + .../llms/vertex_ai/interactions/__init__.py | 0 .../vertex_ai/interactions/transformation.py | 149 +++++++++++ ...t_vertex_ai_interactions_transformation.py | 231 ++++++++++++++++++ 6 files changed, 395 insertions(+) create mode 100644 litellm/llms/vertex_ai/interactions/__init__.py create mode 100644 litellm/llms/vertex_ai/interactions/transformation.py create mode 100644 tests/test_litellm/llms/vertex_ai/interactions/test_vertex_ai_interactions_transformation.py diff --git a/litellm/__init__.py b/litellm/__init__.py index e95b553c5d4..ee2c551481c 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1801,6 +1801,9 @@ if TYPE_CHECKING: from .llms.gemini.interactions.transformation import ( GoogleAIStudioInteractionsConfig as GoogleAIStudioInteractionsConfig, ) + from .llms.vertex_ai.interactions.transformation import ( + VertexAIInteractionsConfig as VertexAIInteractionsConfig, + ) from .llms.openai.chat.o_series_transformation import ( OpenAIOSeriesConfig as OpenAIOSeriesConfig, OpenAIOSeriesConfig as OpenAIO1Config, diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 89c72acc06d..c34c9eefe85 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -242,6 +242,7 @@ LLM_CONFIG_NAMES: Final = ( "OpenRouterResponsesAPIConfig", "BedrockMantleResponsesAPIConfig", "GoogleAIStudioInteractionsConfig", + "VertexAIInteractionsConfig", "OpenAIOSeriesConfig", "AnthropicSkillsConfig", "BaseSkillsAPIConfig", @@ -977,6 +978,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = { ".llms.gemini.interactions.transformation", "GoogleAIStudioInteractionsConfig", ), + "VertexAIInteractionsConfig": ( + ".llms.vertex_ai.interactions.transformation", + "VertexAIInteractionsConfig", + ), "OpenAIOSeriesConfig": ( ".llms.openai.chat.o_series_transformation", "OpenAIOSeriesConfig", diff --git a/litellm/interactions/utils.py b/litellm/interactions/utils.py index 8a1e8836894..3895a85061d 100644 --- a/litellm/interactions/utils.py +++ b/litellm/interactions/utils.py @@ -47,6 +47,13 @@ def get_provider_interactions_api_config( return GoogleAIStudioInteractionsConfig() + if provider in (LlmProviders.VERTEX_AI.value, LlmProviders.VERTEX_AI_BETA.value): + from litellm.llms.vertex_ai.interactions.transformation import ( + VertexAIInteractionsConfig, + ) + + return VertexAIInteractionsConfig() + return None diff --git a/litellm/llms/vertex_ai/interactions/__init__.py b/litellm/llms/vertex_ai/interactions/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/vertex_ai/interactions/transformation.py b/litellm/llms/vertex_ai/interactions/transformation.py new file mode 100644 index 00000000000..0764a8bea62 --- /dev/null +++ b/litellm/llms/vertex_ai/interactions/transformation.py @@ -0,0 +1,149 @@ +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from typing import Final + +from litellm.litellm_core_utils.url_utils import encode_url_path_segment +from litellm.llms.gemini.interactions.transformation import GoogleAIStudioInteractionsConfig +from litellm.llms.vertex_ai.common_utils import validate_vertex_location +from litellm.llms.vertex_ai.vertex_llm_base import VertexBase +from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders + +VERTEX_INTERACTIONS_API_VERSION: Final = "v1beta1" +VERTEX_INTERACTIONS_DEFAULT_LOCATION: Final = "global" + + +@dataclass(frozen=True, slots=True) +class VertexInteractionsTarget: + base_url: str + project_id: str + location: str + + @property + def collection_url(self) -> str: + return ( + f"{self.base_url}/{VERTEX_INTERACTIONS_API_VERSION}" + f"/projects/{self.project_id}/locations/{self.location}/interactions" + ) + + def interaction_url(self, interaction_id: str) -> str: + encoded_interaction_id: Final = encode_url_path_segment(interaction_id, field_name="interaction_id") + return f"{self.collection_url}/{encoded_interaction_id}" + + +class VertexAIInteractionsConfig(VertexBase, GoogleAIStudioInteractionsConfig): + def __init__( + self, + mint_access_token: Callable[[VERTEX_CREDENTIALS_TYPES | None, str | None], tuple[str, str]] | None = None, + ) -> None: + super().__init__() + self._mint_access_token: Final[Callable[[VERTEX_CREDENTIALS_TYPES | None, str | None], tuple[str, str]]] = ( + mint_access_token or self._mint_access_token_with_vertex_base + ) + + def _mint_access_token_with_vertex_base( + self, + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + ) -> tuple[str, str]: + return self._ensure_access_token( + credentials=credentials, project_id=project_id, custom_llm_provider="vertex_ai" + ) + + @property + def custom_llm_provider(self) -> LlmProviders: + return LlmProviders.VERTEX_AI + + @property + def api_version(self) -> str: + return VERTEX_INTERACTIONS_API_VERSION + + def get_default_vertex_location(self) -> str: + return VERTEX_INTERACTIONS_DEFAULT_LOCATION + + def _mint(self, litellm_params: GenericLiteLLMParams) -> tuple[str, str]: + raw_params: Final = litellm_params.model_dump() + return self._mint_access_token( + self.safe_get_vertex_ai_credentials(raw_params), + self.safe_get_vertex_ai_project(raw_params), + ) + + def _target(self, api_base: str | None, litellm_params: GenericLiteLLMParams) -> VertexInteractionsTarget: + _, project_id = self._mint(litellm_params) + if not project_id: + raise ValueError( + "Vertex AI project is required. Set vertex_project, litellm.vertex_project, or VERTEXAI_PROJECT" + ) + location: Final = validate_vertex_location( + self.explicit_vertex_ai_location(litellm_params.model_dump()) or VERTEX_INTERACTIONS_DEFAULT_LOCATION + ) + return VertexInteractionsTarget( + base_url=self.get_api_base(api_base or None, location), + project_id=project_id, + location=location, + ) + + def validate_environment( + self, + headers: Mapping[str, str], + model: str, + litellm_params: GenericLiteLLMParams | None, + ) -> dict: # mutable-ok: BaseInteractionsAPIConfig declares plain-dict headers + access_token, _ = self._mint(litellm_params or GenericLiteLLMParams()) + return { # mutable-ok: BaseInteractionsAPIConfig declares plain-dict headers + "Content-Type": "application/json", + "Authorization": f"Bearer {access_token}", + **headers, + } + + def get_complete_url( + self, + api_base: str | None, + model: str | None, + agent: str | None = None, + litellm_params: Mapping[str, object] | None = None, + stream: bool | None = None, + ) -> str: + params: Final = ( + GenericLiteLLMParams.model_validate(litellm_params) if litellm_params else GenericLiteLLMParams() + ) + collection_url: Final = self._target(api_base, params).collection_url + return f"{collection_url}?alt=sse" if stream else collection_url + + def _interaction_by_id_request( + self, + interaction_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + url_suffix: str = "", + ) -> tuple[str, dict]: # mutable-ok: BaseInteractionsAPIConfig declares a plain-dict request body + target: Final = self._target(api_base or None, litellm_params) + return f"{target.interaction_url(interaction_id)}{url_suffix}", {} # mutable-ok: same base contract + + def transform_get_interaction_request( + self, + interaction_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: Mapping[str, str], + ) -> tuple[str, dict]: # mutable-ok: BaseInteractionsAPIConfig declares a plain-dict request body + return self._interaction_by_id_request(interaction_id, api_base, litellm_params) + + def transform_delete_interaction_request( + self, + interaction_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: Mapping[str, str], + ) -> tuple[str, dict]: # mutable-ok: BaseInteractionsAPIConfig declares a plain-dict request body + return self._interaction_by_id_request(interaction_id, api_base, litellm_params) + + def transform_cancel_interaction_request( + self, + interaction_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: Mapping[str, str], + ) -> tuple[str, dict]: # mutable-ok: BaseInteractionsAPIConfig declares a plain-dict request body + return self._interaction_by_id_request(interaction_id, api_base, litellm_params, url_suffix=":cancel") diff --git a/tests/test_litellm/llms/vertex_ai/interactions/test_vertex_ai_interactions_transformation.py b/tests/test_litellm/llms/vertex_ai/interactions/test_vertex_ai_interactions_transformation.py new file mode 100644 index 00000000000..3364b1b3872 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/interactions/test_vertex_ai_interactions_transformation.py @@ -0,0 +1,231 @@ +import pytest + +import litellm +from litellm.interactions.utils import get_provider_interactions_api_config +from litellm.llms.gemini.interactions.transformation import ( + GoogleAIStudioInteractionsConfig, +) +from litellm.llms.vertex_ai.interactions.transformation import ( + VertexAIInteractionsConfig, +) +from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders + +GLOBAL_BASE = "https://aiplatform.googleapis.com/v1beta1/projects/test-proj/locations/global/interactions" + + +class MinterRecorder: + def __init__(self, resolved_project: str = "creds-proj") -> None: + self.calls: list[tuple[VERTEX_CREDENTIALS_TYPES | None, str | None]] = [] + self.resolved_project = resolved_project + + def __call__( + self, + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + ) -> tuple[str, str]: + self.calls.append((credentials, project_id)) + return "test-token", project_id or self.resolved_project + + +@pytest.fixture +def minter(): + return MinterRecorder() + + +@pytest.fixture +def config(minter): + return VertexAIInteractionsConfig(mint_access_token=minter) + + +@pytest.fixture +def litellm_params(): + return GenericLiteLLMParams(vertex_project="test-proj", vertex_credentials="creds.json") + + +class TestRegistration: + def test_vertex_ai_returns_vertex_config(self): + assert isinstance(get_provider_interactions_api_config("vertex_ai"), VertexAIInteractionsConfig) + + def test_vertex_ai_beta_returns_vertex_config(self): + assert isinstance(get_provider_interactions_api_config("vertex_ai_beta"), VertexAIInteractionsConfig) + + def test_gemini_still_returns_google_ai_studio_config(self): + gemini_config = get_provider_interactions_api_config("gemini") + assert isinstance(gemini_config, GoogleAIStudioInteractionsConfig) + assert not isinstance(gemini_config, VertexAIInteractionsConfig) + + def test_lazy_import_resolves(self): + assert litellm.VertexAIInteractionsConfig is VertexAIInteractionsConfig + + def test_custom_llm_provider_is_vertex_ai(self, config): + assert config.custom_llm_provider == LlmProviders.VERTEX_AI + + +class TestValidateEnvironment: + def test_sets_bearer_auth_without_gemini_headers(self, config, minter, litellm_params): + headers = config.validate_environment( + headers={}, + model="gemini-omni-flash-preview", + litellm_params=litellm_params, + ) + + assert headers["Authorization"] == "Bearer test-token" + assert headers["Content-Type"] == "application/json" + assert "x-goog-api-key" not in headers + assert "Api-Revision" not in headers + assert minter.calls == [("creds.json", "test-proj")] + + def test_caller_authorization_wins(self, config, litellm_params): + headers = config.validate_environment( + headers={"Authorization": "Bearer caller-token"}, + model="gemini-omni-flash-preview", + litellm_params=litellm_params, + ) + + assert headers["Authorization"] == "Bearer caller-token" + + +class TestGetCompleteUrl: + def test_defaults_to_global_v1beta1(self, config, litellm_params): + url = config.get_complete_url( + api_base=None, + model="gemini-omni-flash-preview", + litellm_params=dict(litellm_params), + ) + + assert url == GLOBAL_BASE + + def test_stream_appends_alt_sse(self, config, litellm_params): + url = config.get_complete_url( + api_base=None, + model="gemini-omni-flash-preview", + litellm_params=dict(litellm_params), + stream=True, + ) + + assert url == f"{GLOBAL_BASE}?alt=sse" + + def test_multi_region_location_uses_rep_host(self, config): + url = config.get_complete_url( + api_base=None, + model="gemini-omni-flash-preview", + litellm_params={"vertex_project": "test-proj", "vertex_location": "us"}, + ) + + assert url == "https://aiplatform.us.rep.googleapis.com/v1beta1/projects/test-proj/locations/us/interactions" + + def test_regional_location_uses_regional_host(self, config): + url = config.get_complete_url( + api_base=None, + model="gemini-omni-flash-preview", + litellm_params={"vertex_project": "test-proj", "vertex_location": "us-central1"}, + ) + + assert url == ( + "https://us-central1-aiplatform.googleapis.com" + "/v1beta1/projects/test-proj/locations/us-central1/interactions" + ) + + def test_location_env_fallback_is_ignored(self, config, monkeypatch): + monkeypatch.setenv("VERTEXAI_LOCATION", "us-east5") + + url = config.get_complete_url( + api_base=None, + model="gemini-omni-flash-preview", + litellm_params={"vertex_project": "test-proj"}, + ) + + assert url == GLOBAL_BASE + + def test_api_base_override(self, config, litellm_params): + url = config.get_complete_url( + api_base="https://proxy.example.test", + model="gemini-omni-flash-preview", + litellm_params=dict(litellm_params), + ) + + assert url == "https://proxy.example.test/v1beta1/projects/test-proj/locations/global/interactions" + + def test_project_resolved_from_credentials_when_not_passed(self, config, monkeypatch): + monkeypatch.delenv("VERTEXAI_PROJECT", raising=False) + monkeypatch.setattr(litellm, "vertex_project", None) + + url = config.get_complete_url( + api_base=None, + model="gemini-omni-flash-preview", + litellm_params={"vertex_credentials": "creds.json"}, + ) + + assert url == "https://aiplatform.googleapis.com/v1beta1/projects/creds-proj/locations/global/interactions" + + def test_invalid_location_rejected(self, config): + with pytest.raises(ValueError, match="Invalid vertex_location"): + config.get_complete_url( + api_base=None, + model="gemini-omni-flash-preview", + litellm_params={"vertex_project": "test-proj", "vertex_location": "evil.com#"}, + ) + + def test_missing_project_rejected(self, monkeypatch): + monkeypatch.delenv("VERTEXAI_PROJECT", raising=False) + monkeypatch.setattr(litellm, "vertex_project", None) + + def unresolved_minter( + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + ) -> tuple[str, str]: + return "test-token", "" + + with pytest.raises(ValueError, match="Vertex AI project is required"): + VertexAIInteractionsConfig(mint_access_token=unresolved_minter).get_complete_url( + api_base=None, + model="gemini-omni-flash-preview", + litellm_params={}, + ) + + +class TestInteractionByIdRequests: + def test_get_url(self, config, litellm_params): + url, request_body = config.transform_get_interaction_request( + interaction_id="abc123", + api_base="", + litellm_params=litellm_params, + headers={}, + ) + + assert url == f"{GLOBAL_BASE}/abc123" + assert request_body == {} + + def test_get_url_encodes_interaction_id(self, config, litellm_params): + url, _ = config.transform_get_interaction_request( + interaction_id="id/with space", + api_base="", + litellm_params=litellm_params, + headers={}, + ) + + assert url == f"{GLOBAL_BASE}/id%2Fwith%20space" + + def test_delete_url(self, config, litellm_params): + url, request_body = config.transform_delete_interaction_request( + interaction_id="abc123", + api_base="", + litellm_params=litellm_params, + headers={}, + ) + + assert url == f"{GLOBAL_BASE}/abc123" + assert request_body == {} + + def test_cancel_url(self, config, litellm_params): + url, request_body = config.transform_cancel_interaction_request( + interaction_id="abc123", + api_base="", + litellm_params=litellm_params, + headers={}, + ) + + assert url == f"{GLOBAL_BASE}/abc123:cancel" + assert request_body == {} From ed28581d791e289c2a5b20e0f91678232d5d17ee Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 09:57:25 -0700 Subject: [PATCH 76/93] fix(bedrock_mantle): normalize Codex input item types Mantle rejects Mantle 400s ("Invalid 'input': value did not match any expected variant") on the Codex history item types agent_message, context_compaction, and local_shell_call, killing every Codex multi-agent session on the first sub-agent turn. Rewrite agent_message into an assistant output_text message (preserving encrypted_content slot payloads, which carry the plaintext task through Mantle), context_compaction into Mantle's supported compaction spelling, and local_shell_call into the function_call its recorded function_call_output already pairs with. --- .../responses/transformation.py | 118 +++++++++++- ...bedrock_mantle_responses_transformation.py | 176 ++++++++++++++++++ 2 files changed, 293 insertions(+), 1 deletion(-) diff --git a/litellm/llms/bedrock_mantle/responses/transformation.py b/litellm/llms/bedrock_mantle/responses/transformation.py index 92da5835b2d..9068a641940 100644 --- a/litellm/llms/bedrock_mantle/responses/transformation.py +++ b/litellm/llms/bedrock_mantle/responses/transformation.py @@ -15,8 +15,12 @@ role / access key / profile / web identity), signed via the shared BaseAWSLLM._sign_request after the request body is finalized. """ +import json +from collections.abc import Mapping from typing import Any, Final +from typing_extensions import ReadOnly, TypedDict + import litellm from litellm._logging import verbose_logger from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM @@ -50,6 +54,33 @@ _BEDROCK_MANTLE_SUPPORTED_SERVICE_TIERS: Final = frozenset({"auto", "default"}) _CODEX_ADDITIONAL_TOOLS_INPUT_ITEM_TYPE: Final = "additional_tools" +_CODEX_AGENT_MESSAGE_INPUT_ITEM_TYPE: Final = "agent_message" +_CODEX_CONTEXT_COMPACTION_INPUT_ITEM_TYPE: Final = "context_compaction" +_CODEX_LOCAL_SHELL_CALL_INPUT_ITEM_TYPE: Final = "local_shell_call" + + +class _RewrittenOutputTextBlock(TypedDict): + type: ReadOnly[str] + text: ReadOnly[str] + + +class _RewrittenAssistantMessageItem(TypedDict): + type: ReadOnly[str] + role: ReadOnly[str] + content: ReadOnly[tuple[_RewrittenOutputTextBlock, ...]] + + +class _RewrittenCompactionItem(TypedDict): + type: ReadOnly[str] + encrypted_content: ReadOnly[str] + + +class _RewrittenFunctionCallItem(TypedDict): + type: ReadOnly[str] + call_id: ReadOnly[str] + name: ReadOnly[str] + arguments: ReadOnly[str] + class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPIConfig): def __init__( @@ -155,6 +186,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI headers: dict, ) -> dict: remaining_input, hoisted_tools = self._hoist_codex_additional_tools(input) + normalized_input: Final = self._normalize_codex_input_items(remaining_input) request_params: Final = ( { **response_api_optional_request_params, @@ -168,7 +200,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI ) return super().transform_responses_api_request( model=model, - input=remaining_input, + input=normalized_input, response_api_optional_request_params=request_params, litellm_params=litellm_params, headers=headers, @@ -210,6 +242,90 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI ) return remaining_input, cls._filter_unsupported_tools(hoisted_tools) + @staticmethod + def _agent_message_text(item: "Mapping[str, Any]") -> str: + content: Final = item.get("content") + if not isinstance(content, list): + return "" + return "".join( + str(block.get("text") or block.get("encrypted_content") or "") + for block in content + if isinstance(block, dict) + ) + + @classmethod + def _normalize_agent_message_item(cls, item: "Mapping[str, Any]") -> "_RewrittenAssistantMessageItem | None": + text: Final = cls._agent_message_text(item) + if not text: + return None + rewritten: Final[_RewrittenAssistantMessageItem] = { + "type": "message", + "role": "assistant", + "content": ({"type": "output_text", "text": text},), + } + return rewritten + + @staticmethod + def _normalize_context_compaction_item(item: "Mapping[str, Any]") -> "_RewrittenCompactionItem | None": + encrypted_content: Final = item.get("encrypted_content") + if not isinstance(encrypted_content, str) or not encrypted_content: + return None + rewritten: Final[_RewrittenCompactionItem] = {"type": "compaction", "encrypted_content": encrypted_content} + return rewritten + + @staticmethod + def _normalize_local_shell_call_item(item: "Mapping[str, Any]") -> "_RewrittenFunctionCallItem | None": + call_id: Final = item.get("call_id") + if not isinstance(call_id, str) or not call_id: + return None + action: Final = item.get("action") + rewritten: Final[_RewrittenFunctionCallItem] = { + "type": "function_call", + "call_id": call_id, + "name": "local_shell", + "arguments": json.dumps(action) if isinstance(action, dict) else "{}", + } + return rewritten + + @classmethod + def _normalize_codex_input_item(cls, item: object) -> "tuple[Any, str | None]": + """Returns (normalized item or None to drop it, original type when rewritten).""" + if not isinstance(item, dict): + return item, None + item_type: Final = item.get("type") + if item_type == _CODEX_AGENT_MESSAGE_INPUT_ITEM_TYPE: + return cls._normalize_agent_message_item(item), item_type + if item_type == _CODEX_CONTEXT_COMPACTION_INPUT_ITEM_TYPE: + return cls._normalize_context_compaction_item(item), item_type + if item_type == _CODEX_LOCAL_SHELL_CALL_INPUT_ITEM_TYPE: + return cls._normalize_local_shell_call_item(item), item_type + return item, None + + @classmethod + def _normalize_codex_input_items( + cls, + input: "str | ResponseInputParam", + ) -> "str | ResponseInputParam": + """Rewrite Codex history item types Mantle rejects with 400 "Invalid + 'input': value did not match any expected variant" into supported + equivalents. `agent_message` (Codex multi-agent traffic; its + encrypted_content slot carries the plaintext payload when the model + never issued encrypted args) becomes an assistant message, + `context_compaction` becomes the `compaction` spelling Mantle accepts, + and `local_shell_call` becomes the function_call its recorded + function_call_output already pairs with. + """ + if not isinstance(input, list): + return input + normalized: Final = tuple(cls._normalize_codex_input_item(item) for item in input) + rewritten_types: Final = sorted(frozenset(item_type for _, item_type in normalized if item_type is not None)) + if rewritten_types: + verbose_logger.warning( + "Bedrock Mantle Responses API: rewrote Codex input item type(s) %s that Mantle rejects.", + rewritten_types, + ) + return [item for item, _ in normalized if item is not None] # mutable-ok: ResponseInputParam is a list + def map_openai_params( self, response_api_optional_params: ResponsesAPIOptionalRequestParams, diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py index 9e05d48a18f..984ff997292 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py @@ -8,6 +8,7 @@ gate, the URL construction for both paths, and the shared Bearer auth. """ import copy +import logging import pytest @@ -623,6 +624,181 @@ class TestBedrockMantleCodexAdditionalTools: assert "additional_tools" in str(mock_debug.call_args) +class TestBedrockMantleCodexInputItemNormalization: + """Mantle 400s ("Invalid 'input': value did not match any expected variant") + on the Codex history item types agent_message, context_compaction, and + local_shell_call (verified against bedrock-mantle.us-east-1.api.aws with + openai.gpt-5.6-sol), so the config must rewrite them into supported + equivalents. agent_message is what every Codex multi-agent v2 session sends, + and its encrypted_content slot carries the verbatim plaintext payload when + the upstream model never issued encrypted args, so that slot must be + preserved, not dropped. Mantle also rejects assistant messages with + input_text content, so the rewrite must use output_text.""" + + _USER_MESSAGE = { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "Continue."}], + } + + def _transform(self, input): + cfg = BedrockMantleResponsesAPIConfig() + return cfg.transform_responses_api_request( + model="openai.gpt-5.6-sol", + input=input, + response_api_optional_request_params={}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + def test_plaintext_agent_message_becomes_assistant_output_text_message(self): + body = self._transform( + input=[ + self._USER_MESSAGE, + { + "type": "agent_message", + "id": "amsg_1", + "author": "/root/arithmetic", + "recipient": "/root", + "content": [{"type": "input_text", "text": "Message Type: FINAL_ANSWER\nPayload:\n2+2 is 4."}], + }, + ] + ) + assert body["input"] == [ + self._USER_MESSAGE, + { + "type": "message", + "role": "assistant", + "content": ({"type": "output_text", "text": "Message Type: FINAL_ANSWER\nPayload:\n2+2 is 4."},), + }, + ] + + def test_agent_message_encrypted_content_payload_is_preserved(self): + body = self._transform( + input=[ + { + "type": "agent_message", + "author": "/root", + "recipient": "/root/arithmetic", + "content": [ + {"type": "input_text", "text": "Message Type: NEW_TASK\nPayload:\n"}, + {"type": "encrypted_content", "encrypted_content": "Answer the question 'what is 2+2'."}, + ], + }, + self._USER_MESSAGE, + ] + ) + assert body["input"][0] == { + "type": "message", + "role": "assistant", + "content": ( + { + "type": "output_text", + "text": "Message Type: NEW_TASK\nPayload:\nAnswer the question 'what is 2+2'.", + }, + ), + } + + def test_agent_message_without_any_text_is_dropped(self): + body = self._transform( + input=[ + {"type": "agent_message", "author": "/root", "recipient": "/root/a", "content": []}, + self._USER_MESSAGE, + ] + ) + assert body["input"] == [self._USER_MESSAGE] + + def test_context_compaction_becomes_compaction_with_same_ciphertext(self): + body = self._transform( + input=[ + {"type": "context_compaction", "id": "cc_1", "encrypted_content": "smry_abc123"}, + self._USER_MESSAGE, + ] + ) + assert body["input"] == [ + {"type": "compaction", "encrypted_content": "smry_abc123"}, + self._USER_MESSAGE, + ] + + def test_context_compaction_without_ciphertext_is_dropped(self): + body = self._transform( + input=[ + {"type": "context_compaction", "id": "cc_1"}, + self._USER_MESSAGE, + ] + ) + assert body["input"] == [self._USER_MESSAGE] + + def test_local_shell_call_becomes_function_call_keeping_call_id_pairing(self): + body = self._transform( + input=[ + { + "type": "local_shell_call", + "id": "lsh_1", + "call_id": "call_1", + "status": "completed", + "action": {"type": "exec", "command": ["echo", "hi"]}, + }, + {"type": "function_call_output", "call_id": "call_1", "output": "hi\n"}, + self._USER_MESSAGE, + ] + ) + assert body["input"] == [ + { + "type": "function_call", + "call_id": "call_1", + "name": "local_shell", + "arguments": '{"type": "exec", "command": ["echo", "hi"]}', + }, + {"type": "function_call_output", "call_id": "call_1", "output": "hi\n"}, + self._USER_MESSAGE, + ] + + def test_local_shell_call_without_call_id_is_dropped(self): + body = self._transform( + input=[ + {"type": "local_shell_call", "status": "completed", "action": {"type": "exec", "command": ["ls"]}}, + self._USER_MESSAGE, + ] + ) + assert body["input"] == [self._USER_MESSAGE] + + def test_mantle_supported_item_types_pass_through_untouched(self): + supported_items = [ + self._USER_MESSAGE, + {"type": "compaction", "encrypted_content": "smry_abc123"}, + {"type": "function_call", "name": "shell", "arguments": "{}", "call_id": "call_2"}, + {"type": "function_call_output", "call_id": "call_2", "output": "ok"}, + {"type": "tool_search_call", "call_id": "call_3", "execution": "server", "arguments": {"query": "x"}}, + {"type": "tool_search_output", "call_id": "call_3", "status": "completed", "execution": "server", "tools": []}, + {"type": "compaction_trigger"}, + ] + body = self._transform(input=copy.deepcopy(supported_items)) + assert body["input"] == supported_items + + def test_string_input_passes_through(self): + body = self._transform(input="Say hi.") + assert body["input"] == "Say hi." + + def test_rewrite_is_logged_as_warning_naming_the_types(self, caplog): + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + body = self._transform( + input=[ + {"type": "agent_message", "author": "a", "recipient": "b", "content": [{"type": "input_text", "text": "hi"}]}, + self._USER_MESSAGE, + ] + ) + assert body["input"][0]["role"] == "assistant" + rewrite_warnings = [ + record.getMessage() + for record in caplog.records + if record.levelno == logging.WARNING and "rewrote Codex input item type" in record.getMessage() + ] + assert rewrite_warnings == [ + "Bedrock Mantle Responses API: rewrote Codex input item type(s) ['agent_message'] that Mantle rejects." + ] + + class TestBedrockMantleResponsesRegistry: def test_registry_returns_config_for_gpt_5_5(self, local_cost_map): # gpt-5.x advertises /v1/responses in supported_endpoints (capability) From 68ad575fc21077af8d0e038c92826110ccc0d7c7 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:10:38 -0700 Subject: [PATCH 77/93] fix(bedrock_mantle): register a Bedrock runtime passthrough config so /bedrock/model//invoke works --- .../bedrock/passthrough/transformation.py | 5 + litellm/llms/bedrock_mantle/common_utils.py | 42 +++-- .../passthrough/transformation.py | 44 +++++ litellm/passthrough/main.py | 2 +- litellm/utils.py | 6 + ...drock_mantle_passthrough_transformation.py | 152 ++++++++++++++++++ 6 files changed, 234 insertions(+), 17 deletions(-) create mode 100644 litellm/llms/bedrock_mantle/passthrough/transformation.py create mode 100644 tests/test_litellm/llms/bedrock_mantle/passthrough/test_bedrock_mantle_passthrough_transformation.py diff --git a/litellm/llms/bedrock/passthrough/transformation.py b/litellm/llms/bedrock/passthrough/transformation.py index 0ce2e6f60d3..d0a3c37ffb3 100644 --- a/litellm/llms/bedrock/passthrough/transformation.py +++ b/litellm/llms/bedrock/passthrough/transformation.py @@ -1,4 +1,5 @@ import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Final, Optional, cast from httpx import Response @@ -93,6 +94,9 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD endpoint_url, ) + def get_bedrock_bearer_token(self, litellm_params: Mapping[str, object]) -> str | None: + return None + def sign_request( self, headers: dict, @@ -109,6 +113,7 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD request_data=request_data or {}, api_base=api_base, model=model, + api_key=self.get_bedrock_bearer_token(optional_params), ) def logging_non_streaming_response( diff --git a/litellm/llms/bedrock_mantle/common_utils.py b/litellm/llms/bedrock_mantle/common_utils.py index 889361cd808..d877fbb4e09 100644 --- a/litellm/llms/bedrock_mantle/common_utils.py +++ b/litellm/llms/bedrock_mantle/common_utils.py @@ -13,6 +13,7 @@ global state. """ import re +from collections.abc import Mapping from typing import Final from botocore.exceptions import ( @@ -31,30 +32,39 @@ BEDROCK_MANTLE_DEFAULT_REGION: Final = "us-east-1" MANTLE_HOST_RE: Final = re.compile(r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE) +def resolve_mantle_bearer_token(api_key: str | None) -> str | None: + return api_key or get_secret_str("BEDROCK_MANTLE_API_KEY") or get_secret_str("AWS_BEARER_TOKEN_BEDROCK") + + +def resolve_mantle_region(params: Mapping[str, object]) -> str: + region: Final = params.get("aws_region_name") + if isinstance(region, str) and region: + BaseAWSLLM._validate_aws_region_name(region) + return region + api_base: Final = params.get("api_base") + base: Final = (api_base if isinstance(api_base, str) else None) or get_secret_str("BEDROCK_MANTLE_API_BASE") + if base: + match: Final = MANTLE_HOST_RE.match(base.rstrip("/")) + if match: + return match.group(1) + return ( + get_secret_str("BEDROCK_MANTLE_REGION") + or get_secret_str("AWS_REGION_NAME") + or get_secret_str("AWS_REGION") + or BEDROCK_MANTLE_DEFAULT_REGION + ) + + class BedrockMantleAuthMixin: _aws_signer: BaseAWSLLM @staticmethod def _resolve_bearer_token(api_key: str | None) -> str | None: - return api_key or get_secret_str("BEDROCK_MANTLE_API_KEY") or get_secret_str("AWS_BEARER_TOKEN_BEDROCK") + return resolve_mantle_bearer_token(api_key) @staticmethod def _resolve_region(params: dict) -> str: - region: Final = params.get("aws_region_name") - if region: - BaseAWSLLM._validate_aws_region_name(region) - return region - base: Final = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE") - if base: - match: Final = MANTLE_HOST_RE.match(base.rstrip("/")) - if match: - return match.group(1) - return ( - get_secret_str("BEDROCK_MANTLE_REGION") - or get_secret_str("AWS_REGION_NAME") - or get_secret_str("AWS_REGION") - or BEDROCK_MANTLE_DEFAULT_REGION - ) + return resolve_mantle_region(params) def sign_request( self, diff --git a/litellm/llms/bedrock_mantle/passthrough/transformation.py b/litellm/llms/bedrock_mantle/passthrough/transformation.py new file mode 100644 index 00000000000..1393ac7c6e7 --- /dev/null +++ b/litellm/llms/bedrock_mantle/passthrough/transformation.py @@ -0,0 +1,44 @@ +from collections.abc import Mapping +from typing import Final, Literal + +from litellm.llms.bedrock.passthrough.transformation import BedrockPassthroughConfig +from litellm.llms.bedrock_mantle.common_utils import ( + MANTLE_HOST_RE, + resolve_mantle_bearer_token, + resolve_mantle_region, +) + + +class BedrockMantlePassthroughConfig(BedrockPassthroughConfig): + """Native Bedrock runtime passthrough (InvokeModel, Converse) for deployments declared as bedrock_mantle. + + The Mantle host only serves the OpenAI-compatible surface, so a Mantle api_base lends its region and the + request itself goes to bedrock-runtime, signed with the deployment's Bearer token or SigV4 credentials. + """ + + def _get_aws_region_name( + self, + optional_params: Mapping[str, object], + model: str | None = None, + model_id: str | None = None, + ) -> str: + return resolve_mantle_region(optional_params) + + def get_runtime_endpoint( + self, + api_base: str | None, + aws_bedrock_runtime_endpoint: str | None, + aws_region_name: str, + endpoint_type: Literal["runtime", "agent", "agentcore"] | None = "runtime", + ) -> tuple[str, str]: + is_mantle_host: Final = api_base is not None and MANTLE_HOST_RE.match(api_base.rstrip("/")) is not None + return super().get_runtime_endpoint( + api_base=None if is_mantle_host else api_base, + aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, + aws_region_name=aws_region_name, + endpoint_type=endpoint_type, + ) + + def get_bedrock_bearer_token(self, litellm_params: Mapping[str, object]) -> str | None: + api_key: Final = litellm_params.get("api_key") + return resolve_mantle_bearer_token(api_key if isinstance(api_key, str) else None) diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index 8a2ee2a3af8..4b30afb2f98 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -199,7 +199,7 @@ def llm_passthrough_route( api_key=api_key, ) - litellm_params_dict: Final = get_litellm_params(**kwargs) + litellm_params_dict: Final = get_litellm_params(api_key=api_key, api_base=api_base, **kwargs) if client is None: from litellm.llms.custom_httpx.http_handler import ( diff --git a/litellm/utils.py b/litellm/utils.py index 012e8785321..5cbd0519032 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8610,6 +8610,12 @@ class ProviderConfigManager: ) return BedrockPassthroughConfig() + elif LlmProviders.BEDROCK_MANTLE == provider: + from litellm.llms.bedrock_mantle.passthrough.transformation import ( + BedrockMantlePassthroughConfig, + ) + + return BedrockMantlePassthroughConfig() elif LlmProviders.VLLM == provider or LlmProviders.HOSTED_VLLM == provider: from litellm.llms.vllm.passthrough.transformation import ( VLLMPassthroughConfig, diff --git a/tests/test_litellm/llms/bedrock_mantle/passthrough/test_bedrock_mantle_passthrough_transformation.py b/tests/test_litellm/llms/bedrock_mantle/passthrough/test_bedrock_mantle_passthrough_transformation.py new file mode 100644 index 00000000000..b7f9e492e14 --- /dev/null +++ b/tests/test_litellm/llms/bedrock_mantle/passthrough/test_bedrock_mantle_passthrough_transformation.py @@ -0,0 +1,152 @@ +import json +from unittest.mock import MagicMock, patch + +import httpx +import pytest +from botocore.credentials import Credentials + +from litellm.llms.bedrock.passthrough.transformation import BedrockPassthroughConfig +from litellm.llms.bedrock_mantle.passthrough.transformation import BedrockMantlePassthroughConfig +from litellm.llms.custom_httpx.http_handler import HTTPHandler +from litellm.passthrough.main import llm_passthrough_route +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager + +MANTLE_API_BASE = "https://bedrock-mantle.us-east-2.api.aws" +INVOKE_ENDPOINT = "model/us.openai.gpt-5.6-sol/invoke" +REQUEST_BODY = {"messages": [{"role": "user", "content": "say pong"}], "max_completion_tokens": 64} + + +@pytest.fixture +def no_ambient_aws(monkeypatch): + for name in ( + "AWS_BEARER_TOKEN_BEDROCK", + "BEDROCK_MANTLE_API_KEY", + "BEDROCK_MANTLE_API_BASE", + "BEDROCK_MANTLE_REGION", + "AWS_BEDROCK_RUNTIME_ENDPOINT", + "AWS_REGION_NAME", + "AWS_REGION", + "AWS_DEFAULT_REGION", + ): + monkeypatch.delenv(name, raising=False) + + +def test_bedrock_mantle_registers_its_own_bedrock_passthrough_config(): + config = ProviderConfigManager.get_provider_passthrough_config( + model="us.openai.gpt-5.6-sol", provider=LlmProviders.BEDROCK_MANTLE + ) + assert isinstance(config, BedrockMantlePassthroughConfig) + assert isinstance(config, BedrockPassthroughConfig) + + +def test_mantle_api_base_only_lends_its_region_to_the_runtime_url(no_ambient_aws): + url, base_url = BedrockMantlePassthroughConfig().get_complete_url( + api_base=MANTLE_API_BASE, + api_key=None, + model="us.openai.gpt-5.6-sol", + endpoint=INVOKE_ENDPOINT, + request_query_params=None, + litellm_params={"api_base": MANTLE_API_BASE}, + ) + assert str(url) == f"https://bedrock-runtime.us-east-2.amazonaws.com/{INVOKE_ENDPOINT}" + assert base_url == "https://bedrock-runtime.us-east-2.amazonaws.com" + + +def test_explicit_region_and_non_mantle_api_base_are_kept(no_ambient_aws): + vpc_endpoint = "https://vpce-0123.bedrock-runtime.us-east-1.vpce.amazonaws.com" + url, base_url = BedrockMantlePassthroughConfig().get_complete_url( + api_base=vpc_endpoint, + api_key=None, + model="us.openai.gpt-5.6-sol", + endpoint=INVOKE_ENDPOINT, + request_query_params=None, + litellm_params={"api_base": vpc_endpoint, "aws_region_name": "us-east-1"}, + ) + assert str(url) == f"{vpc_endpoint}/{INVOKE_ENDPOINT}" + assert base_url == vpc_endpoint + + +def test_region_falls_back_to_the_mantle_default_without_any_hint(no_ambient_aws): + url, _ = BedrockMantlePassthroughConfig().get_complete_url( + api_base=None, + api_key=None, + model="us.openai.gpt-5.6-sol", + endpoint=INVOKE_ENDPOINT, + request_query_params=None, + litellm_params={}, + ) + assert str(url) == f"https://bedrock-runtime.us-east-1.amazonaws.com/{INVOKE_ENDPOINT}" + + +@pytest.mark.parametrize( + ("litellm_params", "env", "expected_bearer"), + [ + ({"api_key": "deployment-bedrock-api-key"}, {}, "deployment-bedrock-api-key"), + ({}, {"BEDROCK_MANTLE_API_KEY": "mantle-env-key"}, "mantle-env-key"), + ({}, {"AWS_BEARER_TOKEN_BEDROCK": "aws-env-key"}, "aws-env-key"), + ], +) +def test_sign_request_uses_the_deployment_bearer_token(no_ambient_aws, monkeypatch, litellm_params, env, expected_bearer): + for name, value in env.items(): + monkeypatch.setenv(name, value) + headers, body = BedrockMantlePassthroughConfig().sign_request( + headers={}, + litellm_params=litellm_params, + request_data=REQUEST_BODY, + api_base=f"https://bedrock-runtime.us-east-1.amazonaws.com/{INVOKE_ENDPOINT}", + model="us.openai.gpt-5.6-sol", + ) + assert headers["Authorization"] == f"Bearer {expected_bearer}" + assert body is not None + assert json.loads(body) == REQUEST_BODY + + +def test_sign_request_falls_back_to_sigv4_scoped_to_the_mantle_region(no_ambient_aws): + config = BedrockMantlePassthroughConfig() + with patch.object(config, "get_credentials", return_value=Credentials("AKIA", "secret")): + headers, body = config.sign_request( + headers={}, + litellm_params={"api_base": MANTLE_API_BASE}, + request_data=REQUEST_BODY, + api_base=f"https://bedrock-runtime.us-east-2.amazonaws.com/{INVOKE_ENDPOINT}", + model="us.openai.gpt-5.6-sol", + ) + assert headers["Authorization"].startswith("AWS4-HMAC-SHA256 Credential=AKIA/") + assert "/us-east-2/bedrock/aws4_request" in headers["Authorization"] + assert body is not None + assert json.loads(body) == REQUEST_BODY + + +@pytest.mark.parametrize( + ("route_kwargs", "env", "expected_bearer"), + [ + ({"api_key": "deployment-bedrock-api-key"}, {}, "deployment-bedrock-api-key"), + ({}, {"BEDROCK_MANTLE_API_KEY": "mantle-env-key"}, "mantle-env-key"), + ], +) +def test_invoke_passthrough_route_reaches_bedrock_runtime_for_a_mantle_deployment( + no_ambient_aws, monkeypatch, route_kwargs, env, expected_bearer +): + for name, value in env.items(): + monkeypatch.setenv(name, value) + client = HTTPHandler() + with ( + patch.object(client.client, "send", return_value=MagicMock(status_code=200)), + patch.object(client.client, "build_request", wraps=client.client.build_request) as build_request, + ): + response = llm_passthrough_route( + model="bedrock_mantle/us.openai.gpt-5.6-sol", + endpoint=INVOKE_ENDPOINT, + method="POST", + api_base=MANTLE_API_BASE, + json=dict(REQUEST_BODY), + client=client, + litellm_logging_obj=MagicMock(), + **route_kwargs, + ) + assert response.status_code == 200 + sent = build_request.call_args.kwargs + assert str(sent["url"]) == f"https://bedrock-runtime.us-east-2.amazonaws.com/{INVOKE_ENDPOINT}" + assert sent["headers"]["Authorization"] == f"Bearer {expected_bearer}" + assert json.loads(sent["content"]) == REQUEST_BODY From b46f17faf5d56ce853fff94d0daa5ffa0d2fb428 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:18:55 -0700 Subject: [PATCH 78/93] fix(together_ai): default endpoints to api.together.ai instead of api.together.xyz Together AI moved its canonical API host from api.together.xyz to api.together.ai. Default the provider api_base and the rerank handler to the new host, make rerank honor api_base and TOGETHER_AI_API_BASE like chat already does, map both hosts to together_ai when passed as api_base, and delete the dead models/info fetch in factory.py. --- basedpyright-code-budget.json | 8 +-- litellm/constants.py | 1 + .../get_llm_provider_logic.py | 10 ++- .../prompt_templates/factory.py | 43 ------------ litellm/llms/together_ai/rerank/handler.py | 12 +++- litellm/rerank_api/main.py | 3 + ruff-strict-budget.json | 6 +- .../test_get_llm_provider_endpoint_match.py | 39 +++++++++++ tests/test_litellm/rerank_api/test_main.py | 67 +++++++++++++++++++ type-discipline-budget.json | 2 +- 10 files changed, 136 insertions(+), 55 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 664e1669834..f4d4e25859a 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,6 +1,6 @@ { "reportAny": { - "limit": 19955 + "limit": 19949 }, "reportArgumentType": { "limit": 2566 @@ -54,7 +54,7 @@ "limit": 0 }, "reportMissingParameterType": { - "limit": 5663 + "limit": 5661 }, "reportMissingTypeArgument": { "limit": 15555 @@ -105,10 +105,10 @@ "limit": 109 }, "reportUnknownMemberType": { - "limit": 39011 + "limit": 39009 }, "reportUnknownParameterType": { - "limit": 19885 + "limit": 19883 }, "reportUnknownVariableType": { "limit": 30569 diff --git a/litellm/constants.py b/litellm/constants.py index 0a1ada3bab2..78aba30f9c0 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -750,6 +750,7 @@ openai_compatible_endpoints: Final[list] = [ "api.groq.com/openai/v1", "https://integrate.api.nvidia.com/v1", "api.deepseek.com/v1", + "api.together.ai/v1", "api.together.xyz/v1", "app.empower.dev/api/v1", "https://api.friendli.ai/serverless/v1", diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index e674fc37673..d2d82064c47 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -272,6 +272,14 @@ def get_llm_provider( elif endpoint == "api.deepseek.com/v1": custom_llm_provider = "deepseek" dynamic_api_key = get_secret_str("DEEPSEEK_API_KEY") + elif endpoint == "api.together.ai/v1" or endpoint == "api.together.xyz/v1": + custom_llm_provider = "together_ai" + dynamic_api_key = ( + get_secret_str("TOGETHER_API_KEY") + or get_secret_str("TOGETHER_AI_API_KEY") + or get_secret_str("TOGETHERAI_API_KEY") + or get_secret_str("TOGETHER_AI_TOKEN") + ) elif endpoint == "ollama.com": custom_llm_provider = "ollama" dynamic_api_key = get_secret_str("OLLAMA_API_KEY") @@ -707,7 +715,7 @@ def _get_openai_compatible_provider_info( dynamic_api_key, ) = litellm.ZAIChatConfig()._get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "together_ai": - api_base = api_base or get_secret_str("TOGETHER_AI_API_BASE") or "https://api.together.xyz/v1" + api_base = api_base or get_secret_str("TOGETHER_AI_API_BASE") or "https://api.together.ai/v1" dynamic_api_key = api_key or ( get_secret_str("TOGETHER_API_KEY") or get_secret_str("TOGETHER_AI_API_KEY") diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 826a890eca9..86cfbf70255 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -643,49 +643,6 @@ def claude_2_1_pt( return prompt -### TOGETHER AI - - -def get_model_info(token, model): - try: - headers: Final = {"Authorization": f"Bearer {token}"} - client: Final = HTTPHandler(concurrent_limit=1) - response: Final = client.get("https://api.together.xyz/models/info", headers=headers) - if response.status_code == 200: - model_info: Final = response.json() - for m in model_info: - if m["name"].lower().strip() == model.strip(): - return m["config"].get("prompt_format", None), m["config"].get("chat_template", None) - return None, None - else: - return None, None - except Exception: # safely fail a prompt template request - return None, None - - -## OLD TOGETHER AI FLOW -# def format_prompt_togetherai(messages, prompt_format, chat_template): -# if prompt_format is None: -# return default_pt(messages) - -# human_prompt, assistant_prompt = prompt_format.split("{prompt}") - -# if chat_template is not None: -# prompt = hf_chat_template( -# model=None, messages=messages, chat_template=chat_template -# ) -# elif prompt_format is not None: -# prompt = custom_prompt( -# role_dict={}, -# messages=messages, -# initial_prompt_value=human_prompt, -# final_prompt_value=assistant_prompt, -# ) -# else: -# prompt = default_pt(messages) -# return prompt - - ### IBM Granite diff --git a/litellm/llms/together_ai/rerank/handler.py b/litellm/llms/together_ai/rerank/handler.py index 10246451a9d..8407018b898 100644 --- a/litellm/llms/together_ai/rerank/handler.py +++ b/litellm/llms/together_ai/rerank/handler.py @@ -16,11 +16,16 @@ from litellm.llms.together_ai.rerank.transformation import TogetherAIRerankConfi from litellm.types.rerank import RerankRequest, RerankResponse +def _rerank_url(api_base: str) -> str: + return f"{api_base.rstrip('/')}/rerank" + + class TogetherAIRerank(BaseLLM): def rerank( self, model: str, api_key: str, + api_base: str, query: str, documents: list[str | dict[str, Any]], top_n: int | None = None, @@ -46,10 +51,10 @@ class TogetherAIRerank(BaseLLM): raise ValueError("TogetherAI does not support max_chunks_per_doc") if _is_async: - return self.async_rerank(request_data_dict, api_key) # Call async method + return self.async_rerank(request_data_dict, api_key, api_base) # Call async method response: Final = client.post( - "https://api.together.xyz/v1/rerank", + _rerank_url(api_base), headers={ "accept": "application/json", "content-type": "application/json", @@ -69,11 +74,12 @@ class TogetherAIRerank(BaseLLM): self, request_data_dict: dict[str, Any], api_key: str, + api_base: str, ) -> RerankResponse: client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.TOGETHER_AI) # Use async client response: Final = await client.post( - "https://api.together.xyz/v1/rerank", + _rerank_url(api_base), headers={ "accept": "application/json", "content-type": "application/json", diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 15a6f18a6bb..c8f7842aebf 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -277,6 +277,8 @@ def rerank( if api_key is None: raise ValueError("TogetherAI API key is required, please set 'TOGETHERAI_API_KEY' in your environment") + api_base = dynamic_api_base or optional_params.api_base or litellm.api_base or "https://api.together.ai/v1" + response = together_rerank.rerank( model=model, query=query, @@ -286,6 +288,7 @@ def rerank( return_documents=return_documents, max_chunks_per_doc=max_chunks_per_doc, api_key=api_key, + api_base=api_base, _is_async=_is_async, ) elif _custom_llm_provider == litellm.LlmProviders.JINA_AI: diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 03318718fb5..1ca152985f9 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -1,6 +1,6 @@ { "ANN001": { - "limit": 3018 + "limit": 3016 }, "ANN002": { "limit": 71 @@ -9,7 +9,7 @@ "limit": 827 }, "ANN201": { - "limit": 2016 + "limit": 2015 }, "ANN202": { "limit": 852 @@ -57,7 +57,7 @@ "limit": 3 }, "BLE001": { - "limit": 2919 + "limit": 2918 }, "C401": { "limit": 8 diff --git a/tests/test_litellm/litellm_core_utils/test_get_llm_provider_endpoint_match.py b/tests/test_litellm/litellm_core_utils/test_get_llm_provider_endpoint_match.py index bda7ab4afc6..5c20284282a 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_llm_provider_endpoint_match.py +++ b/tests/test_litellm/litellm_core_utils/test_get_llm_provider_endpoint_match.py @@ -133,3 +133,42 @@ class TestGetLlmProviderRejectsAttackerSmuggledApiBase: assert provider == "groq" assert dynamic_api_key == "server-real-groq-key" + + +class TestTogetherApiBaseResolvesProvider: + """ + Regression for the Together host migration: both the current + ``api.together.ai`` host and the legacy ``api.together.xyz`` host must + resolve to ``together_ai`` when passed as ``api_base``. Before the fix + the endpoint list carried the legacy host but the provider-mapping + chain had no branch for it, so the match fell through with a None + provider and the deployment failed with "LLM Provider NOT provided". + """ + + @pytest.mark.parametrize( + "api_base", + [ + "https://api.together.ai/v1", + "https://api.together.xyz/v1", + ], + ) + def test_together_api_base_resolves_to_together_ai(self, api_base, monkeypatch): + monkeypatch.setenv("TOGETHER_API_KEY", "together-key-from-env") + + model, provider, dynamic_api_key, returned_api_base = get_llm_provider( + model="some-model", + api_base=api_base, + ) + + assert provider == "together_ai" + assert dynamic_api_key == "together-key-from-env" + assert returned_api_base == api_base + assert model == "some-model" + + def test_together_default_api_base_is_together_ai(self, monkeypatch): + monkeypatch.delenv("TOGETHER_AI_API_BASE", raising=False) + + _, provider, _, api_base = get_llm_provider(model="together_ai/some-model") + + assert provider == "together_ai" + assert api_base == "https://api.together.ai/v1" diff --git a/tests/test_litellm/rerank_api/test_main.py b/tests/test_litellm/rerank_api/test_main.py index 85777afe81c..587be59c550 100644 --- a/tests/test_litellm/rerank_api/test_main.py +++ b/tests/test_litellm/rerank_api/test_main.py @@ -1,6 +1,10 @@ import logging from unittest.mock import MagicMock, patch +import httpx +import pytest +import respx + import litellm @@ -62,3 +66,66 @@ def test_rerank_does_not_log_request_content_at_info(caplog): assert all( r.levelno == logging.DEBUG for r in optional_params_logs ), "optional_rerank_params must be logged at DEBUG, not INFO" + + +TOGETHER_RERANK_BODY = { + "id": "rerank-mock-id", + "results": [{"index": 0, "relevance_score": 0.95}], + "usage": {"prompt_tokens": 10, "total_tokens": 10}, +} + + +def test_together_rerank_defaults_to_together_ai_host(respx_mock: respx.MockRouter, monkeypatch): + """Regression for the Together host migration: rerank used to hardcode + https://api.together.xyz/v1/rerank. The default must now be api.together.ai.""" + monkeypatch.delenv("TOGETHER_AI_API_BASE", raising=False) + + mock_route = respx_mock.post("https://api.together.ai/v1/rerank") + mock_route.return_value = httpx.Response(200, json=TOGETHER_RERANK_BODY) + + response = litellm.rerank( + model="together_ai/mixedbread-ai/mxbai-rerank-large-v2", + query=MARKER_QUERY, + documents=[MARKER_DOC], + api_key="fake-together-key", + ) + + assert mock_route.called + assert response.results[0]["relevance_score"] == 0.95 + + +def test_together_rerank_honors_api_base(respx_mock: respx.MockRouter): + """Regression: a custom api_base was silently ignored by the Together rerank handler.""" + mock_route = respx_mock.post("https://custom-together.example/v1/rerank") + mock_route.return_value = httpx.Response(200, json=TOGETHER_RERANK_BODY) + + litellm.rerank( + model="together_ai/mixedbread-ai/mxbai-rerank-large-v2", + query=MARKER_QUERY, + documents=[MARKER_DOC], + api_key="fake-together-key", + api_base="https://custom-together.example/v1", + ) + + assert mock_route.called + assert mock_route.calls[0].request.headers["authorization"] == "Bearer fake-together-key" + + +@pytest.mark.asyncio +async def test_together_rerank_async_honors_env_api_base(respx_mock: respx.MockRouter, monkeypatch): + """Regression: TOGETHER_AI_API_BASE was honored by chat but ignored by rerank.""" + monkeypatch.setenv("TOGETHER_AI_API_BASE", "https://env-together.example/v1") + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + + mock_route = respx_mock.post("https://env-together.example/v1/rerank") + mock_route.return_value = httpx.Response(200, json=TOGETHER_RERANK_BODY) + + response = await litellm.arerank( + model="together_ai/mixedbread-ai/mxbai-rerank-large-v2", + query=MARKER_QUERY, + documents=[MARKER_DOC], + api_key="fake-together-key", + ) + + assert mock_route.called + assert response.results[0]["relevance_score"] == 0.95 diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 627811a7f1d..f9fe3042f1e 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -3,7 +3,7 @@ "limit": 22805 }, "LIT002": { - "limit": 26873 + "limit": 26872 }, "LIT003": { "limit": 269 From fd1dca05deb0fd4d50153b42412f716e341236c3 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:19:56 -0700 Subject: [PATCH 79/93] fix(bedrock_mantle): type the codex item dispatcher without Any --- litellm/llms/bedrock_mantle/responses/transformation.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/litellm/llms/bedrock_mantle/responses/transformation.py b/litellm/llms/bedrock_mantle/responses/transformation.py index 9068a641940..3e5dd4ff87d 100644 --- a/litellm/llms/bedrock_mantle/responses/transformation.py +++ b/litellm/llms/bedrock_mantle/responses/transformation.py @@ -288,7 +288,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI return rewritten @classmethod - def _normalize_codex_input_item(cls, item: object) -> "tuple[Any, str | None]": + def _normalize_codex_input_item(cls, item: object) -> "tuple[object, str | None]": """Returns (normalized item or None to drop it, original type when rewritten).""" if not isinstance(item, dict): return item, None @@ -324,7 +324,8 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI "Bedrock Mantle Responses API: rewrote Codex input item type(s) %s that Mantle rejects.", rewritten_types, ) - return [item for item, _ in normalized if item is not None] # mutable-ok: ResponseInputParam is a list + kept: Final = [item for item, _ in normalized if item is not None] # mutable-ok: ResponseInputParam is a list + return kept # pyright: ignore[reportReturnType] # Codex passthrough items sit outside the OpenAI input union def map_openai_params( self, From 6be000f1f35091cbbdfedb7df4e0dd8d494c0eaf Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:23:12 -0700 Subject: [PATCH 80/93] test(bedrock_mantle): type _repo_cost_map return instead of bare dict --- .../test_bedrock_mantle_responses_transformation.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py index fd279a2bc1f..28c6060e5cc 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py @@ -1568,7 +1568,7 @@ class TestBedrockMantleResponsesPricing: assert "bedrock_mantle/openai.gpt-5.4" in litellm.bedrock_mantle_models -def _repo_cost_map(map_name: str) -> dict: +def _repo_cost_map(map_name: str) -> dict[str, dict[str, object]]: repo_root = Path(__file__).resolve().parents[4] paths = { "root": repo_root / "model_prices_and_context_window.json", From 0bd4d323da3a5969a7a3940faef03ecd239d2a98 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:33:40 -0700 Subject: [PATCH 81/93] fix(router): resolve provider from api_base in deployment validation and acompletion Router._add_deployment called get_llm_provider without the deployment's api_base, so a config entry with a bare model plus a known OpenAI-compatible endpoint failed startup validation with LLM Provider NOT provided and the proxy returned 400 no healthy deployments for that model group. acompletion had the same gap at request time: it forwarded only base_url into its get_llm_provider call, dropping the api_base kwarg the router passes. Both now forward api_base so endpoint matching resolves the provider the same way sync completion already does --- litellm/main.py | 2 +- litellm/router.py | 1 + tests/test_litellm/test_main.py | 13 +++++++ tests/test_litellm/test_router.py | 62 +++++++++++++++++++++++++++++++ 4 files changed, 77 insertions(+), 1 deletion(-) diff --git a/litellm/main.py b/litellm/main.py index d3967473f99..6dfd8c2d675 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -602,7 +602,7 @@ async def acompletion( _, custom_llm_provider, _, _ = get_llm_provider( model=model, custom_llm_provider=custom_llm_provider, - api_base=base_url, + api_base=kwargs.get("api_base") or base_url, ) fallbacks = fallbacks or litellm.model_fallbacks diff --git a/litellm/router.py b/litellm/router.py index 6ee474730c9..33da39e2677 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -8341,6 +8341,7 @@ class Router: ) = litellm.get_llm_provider( model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.get("custom_llm_provider", None), + api_base=deployment.litellm_params.api_base, ) # done reading model["litellm_params"] # Check if provider is supported: either in enum or JSON-configured diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 99b1cc826aa..8f2b06be4b3 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -2944,3 +2944,16 @@ def test_a_stream_that_reported_no_usage_is_still_billed(local_cost_map): assert cost == pytest.approx( _priced_at(rebuilt.usage.prompt_tokens, rebuilt.usage.completion_tokens) ) + + +@pytest.mark.asyncio +async def test_acompletion_resolves_provider_from_api_base(): + response = await litellm.acompletion( + model="deepseek-chat", + api_base="https://api.deepseek.com/v1", + api_key="fake-key", + messages=[{"role": "user", "content": "hi"}], + mock_response="resolved", + ) + + assert response.choices[0].message.content == "resolved" diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 910b874c2ac..df6754cd7ca 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -8921,3 +8921,65 @@ class TestAzureBaseModelFallbackLogging: deployment=None, received_model_name="my-group", id="azure-base-model-test-id" ) assert model_info["max_input_tokens"] == litellm.model_cost["azure/gpt-4o-mini"]["max_input_tokens"] + + +class TestAddDeploymentApiBaseProviderResolution: + def test_bare_model_with_known_api_base_initializes(self): + router = litellm.Router( + model_list=[ + { + "model_name": "groq-pinned", + "litellm_params": { + "model": "llama-3.3-70b-versatile", + "api_base": "https://api.groq.com/openai/v1", + "api_key": "fake-key", + }, + }, + { + "model_name": "deepseek-pinned", + "litellm_params": { + "model": "deepseek-chat", + "api_base": "https://api.deepseek.com/v1", + "api_key": "fake-key", + }, + }, + ] + ) + + model_list = router.get_model_list() + assert model_list is not None + assert {m["model_name"] for m in model_list} == {"groq-pinned", "deepseek-pinned"} + + def test_bare_model_with_unknown_api_base_still_raises(self): + with pytest.raises(litellm.BadRequestError, match="LLM Provider NOT provided"): + litellm.Router( + model_list=[ + { + "model_name": "mystery", + "litellm_params": { + "model": "some-unknown-model", + "api_base": "https://llm.internal.example.com/v1", + "api_key": "fake-key", + }, + } + ] + ) + + def test_explicit_custom_llm_provider_beats_api_base_endpoint_match(self): + router = litellm.Router( + model_list=[ + { + "model_name": "openai-via-gateway", + "litellm_params": { + "model": "gpt-3.5-turbo", + "custom_llm_provider": "openai", + "api_base": "https://api.groq.com/openai/v1", + "api_key": "fake-key", + }, + } + ] + ) + + deployment = router.get_deployment_by_model_group_name("openai-via-gateway") + assert deployment is not None + assert deployment.litellm_params.custom_llm_provider == "openai" From 367a6e5dc5fc1e23d31b7395d96fa2a5332fd6d2 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:37:48 -0700 Subject: [PATCH 82/93] test(router): pin the guard that keeps a junk-typed operator effort value out of model group info --- tests/test_litellm/test_router.py | 34 +++++++++++++++++++++++++++++++ 1 file changed, 34 insertions(+) diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 5f1941574eb..e8830d05065 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -9092,6 +9092,40 @@ def test_model_group_info_reasoning_efforts_ignore_a_value_declared_in_model_inf assert result.supported_reasoning_efforts == ("none", "minimal", "low", "medium", "high") +def test_model_group_info_survives_a_junk_typed_operator_effort_value(): + """A deployment's registered model_info reads back with whatever the operator wrote under any + key, so a wrong-typed supported_reasoning_efforts must not fail the group's info. Only the + constructor's trailing override keeps the junk away from ModelGroupInfo validation.""" + router = litellm.Router( + model_list=[ + { + "model_name": "junk-declared-group", + "litellm_params": {"model": "openai/lone-reasoner"}, + "model_info": {"id": "junk-deployment"}, + }, + ] + ) + + def _model_info(model_id: str, model_name: str): + return { + "key": model_name, + "litellm_provider": "openai", + "mode": "chat", + "supports_reasoning": True, + "supports_none_reasoning_effort": True, + "supported_reasoning_efforts": "high", + } + + with patch.object(router, "get_deployment_model_info", side_effect=_model_info): + result = router._set_model_group_info( + model_group="junk-declared-group", + user_facing_model_group_name="junk-declared-group", + ) + + assert result is not None + assert result.supported_reasoning_efforts == ("none", "minimal", "low", "medium", "high") + + def test_model_group_info_reasoning_efforts_ignore_a_mode_the_operator_declared(): """A deployment is registered in the cost map under its own id with whatever model_info the operator wrote, so a mode they set themselves reads back exactly like one the map supplied. Only From 5e6b6c6281f8d26bda113cdd54be0fe71afa4d77 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:39:10 -0700 Subject: [PATCH 83/93] fix(together_ai): let an explicit api_key beat the Together env key on api_base match --- litellm/litellm_core_utils/get_llm_provider_logic.py | 2 +- litellm/llms/together_ai/rerank/handler.py | 2 +- .../test_get_llm_provider_endpoint_match.py | 12 ++++++++++++ 3 files changed, 14 insertions(+), 2 deletions(-) diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index d2d82064c47..005e94ebe82 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -274,7 +274,7 @@ def get_llm_provider( dynamic_api_key = get_secret_str("DEEPSEEK_API_KEY") elif endpoint == "api.together.ai/v1" or endpoint == "api.together.xyz/v1": custom_llm_provider = "together_ai" - dynamic_api_key = ( + dynamic_api_key = api_key or ( get_secret_str("TOGETHER_API_KEY") or get_secret_str("TOGETHER_AI_API_KEY") or get_secret_str("TOGETHERAI_API_KEY") diff --git a/litellm/llms/together_ai/rerank/handler.py b/litellm/llms/together_ai/rerank/handler.py index 8407018b898..b8079e52c97 100644 --- a/litellm/llms/together_ai/rerank/handler.py +++ b/litellm/llms/together_ai/rerank/handler.py @@ -51,7 +51,7 @@ class TogetherAIRerank(BaseLLM): raise ValueError("TogetherAI does not support max_chunks_per_doc") if _is_async: - return self.async_rerank(request_data_dict, api_key, api_base) # Call async method + return self.async_rerank(request_data_dict, api_key, api_base) response: Final = client.post( _rerank_url(api_base), diff --git a/tests/test_litellm/litellm_core_utils/test_get_llm_provider_endpoint_match.py b/tests/test_litellm/litellm_core_utils/test_get_llm_provider_endpoint_match.py index 5c20284282a..6cacd119030 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_llm_provider_endpoint_match.py +++ b/tests/test_litellm/litellm_core_utils/test_get_llm_provider_endpoint_match.py @@ -165,6 +165,18 @@ class TestTogetherApiBaseResolvesProvider: assert returned_api_base == api_base assert model == "some-model" + def test_explicit_api_key_beats_together_env_key(self, monkeypatch): + monkeypatch.setenv("TOGETHER_API_KEY", "together-key-from-env") + + _, provider, dynamic_api_key, _ = get_llm_provider( + model="some-model", + api_base="https://api.together.ai/v1", + api_key="explicit-caller-key", + ) + + assert provider == "together_ai" + assert dynamic_api_key == "explicit-caller-key" + def test_together_default_api_base_is_together_ai(self, monkeypatch): monkeypatch.delenv("TOGETHER_AI_API_BASE", raising=False) From 0ea6f5e159350aa75e9118ba027b7286074cb0fe Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Tue, 25 Aug 2026 10:50:11 -0700 Subject: [PATCH 84/93] fix(azure_ai): stamp the model router's selected model instead of matching on the model name The model Azure Model Router served was recovered by checking whether the text "model_router" or "model-router" appeared in a model string. Spend logs applied that check to the litellm model path, where the route prefix guarantees a match, but the proxy applied it to the client's model group alias, which carries no prefix. A model group named anything else therefore lost the selected model in both the response and the spend row. AzureModelRouterConfig now stamps the served model onto _hidden_params, and the spend log payload and the proxy's response restamping read that stamp. The name heuristic survives as a fallback for callers with no response in hand, routed through get_azure_ai_route so it lives in one place. --- litellm/litellm_core_utils/litellm_logging.py | 12 +- .../azure_model_router/transformation.py | 23 +- litellm/llms/azure_ai/common_utils.py | 36 ++ litellm/proxy/common_request_processing.py | 21 +- .../test_litellm_logging.py | 472 +++++++----------- .../chat/test_azure_ai_transformation.py | 97 +++- .../proxy/test_common_request_processing.py | 443 +++++++--------- 7 files changed, 538 insertions(+), 566 deletions(-) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 9b7707eabe1..01802c3bbb3 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -5762,11 +5762,15 @@ def get_standard_logging_object_payload( response_model_name = final_response_obj.get("model") # For Azure Model Router, preserve the actual model in the top-level standard - # logging payload only when the user has opted in. + # logging payload. + from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo + requested_model: Final = kwargs.get("model") - if ( - isinstance(requested_model, str) - and ("model_router" in requested_model.lower() or "model-router" in requested_model.lower()) + stamped_selected_model: Final = AzureFoundryModelInfo.get_model_router_selected_model(hidden_params) + if stamped_selected_model is not None: + model_name = stamped_selected_model + elif ( + AzureFoundryModelInfo.is_model_router_call(model=requested_model, hidden_params=hidden_params) and isinstance(response_model_name, str) and response_model_name ): diff --git a/litellm/llms/azure_ai/azure_model_router/transformation.py b/litellm/llms/azure_ai/azure_model_router/transformation.py index 1a924088390..d33564c0f9a 100644 --- a/litellm/llms/azure_ai/azure_model_router/transformation.py +++ b/litellm/llms/azure_ai/azure_model_router/transformation.py @@ -65,15 +65,24 @@ class AzureModelRouterConfig(AzureAIStudioConfig): Extracts the actual model used from the Azure response (e.g., gpt-5-nano-2025-08-07) and returns it with the azure_ai/ prefix for proper display and cost tracking. + + Also stamps that model onto ``_hidden_params`` so downstream consumers (spend logs, + response restamping) can read it instead of guessing the route from the model string. """ - from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo + from litellm.llms.azure_ai.common_utils import ( + AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY, + AzureFoundryModelInfo, + ) + from litellm.router_utils.add_retry_fallback_headers import ( + get_hidden_params_dict, + ) # Get base model for the parent call (strips routing prefixes for API compatibility) base_model: Final[str] = AzureFoundryModelInfo.get_base_model(model) # Call parent transform_response first - this will extract the actual model # from the raw response (e.g., "gpt-5-nano-2025-08-07") - model_response = super().transform_response( + transformed_response: Final = super().transform_response( model=base_model, raw_response=raw_response, model_response=model_response, @@ -86,7 +95,15 @@ class AzureModelRouterConfig(AzureAIStudioConfig): api_key=api_key, json_mode=json_mode, ) - return model_response + selected_model: Final = transformed_response.model + if selected_model: + # Rebuilt rather than mutated in place: ModelResponseBase declares _hidden_params as a + # class-level dict, so an in-place write can bleed into unrelated responses. + transformed_response._hidden_params = { # pyright: ignore[reportPrivateUsage] # ModelResponse exposes no public hidden-params setter + **get_hidden_params_dict(transformed_response), + AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY: selected_model, + } + return transformed_response def calculate_additional_costs(self, model: str, prompt_tokens: int, completion_tokens: int) -> dict | None: """ diff --git a/litellm/llms/azure_ai/common_utils.py b/litellm/llms/azure_ai/common_utils.py index d25a8fd6561..9d37d8d1185 100644 --- a/litellm/llms/azure_ai/common_utils.py +++ b/litellm/llms/azure_ai/common_utils.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from typing import Final, Literal import litellm @@ -5,6 +6,8 @@ from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues +AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY: Final = "azure_model_router_selected_model" + class AzureFoundryModelInfo(BaseLLMModelInfo): """Model info for Azure AI / Azure Foundry models.""" @@ -37,6 +40,39 @@ class AzureFoundryModelInfo(BaseLLMModelInfo): return "model_router" return "default" + @staticmethod + def get_model_router_selected_model(hidden_params: Mapping[str, object] | None) -> str | None: + """The model Azure Model Router actually served, stamped by ``AzureModelRouterConfig``. + + Reading this beats re-deriving the route from a model string: the stamp is set on the + code path that was actually taken, so it holds no matter what the caller named the model. + """ + if not hidden_params: + return None + selected: Final = hidden_params.get(AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY) + if isinstance(selected, str) and selected: + return selected + return None + + @staticmethod + def is_model_router_call( + model: str | None = None, + hidden_params: Mapping[str, object] | None = None, + ) -> bool: + """Whether a request went down the Azure Model Router route. + + Prefers the response stamp, then the deployment's litellm model path, and only then the + caller-supplied name. The last two go through ``get_azure_ai_route`` so the model-router + name heuristic lives in exactly one place. + """ + if AzureFoundryModelInfo.get_model_router_selected_model(hidden_params) is not None: + return True + deployment_model: Final = (hidden_params or {}).get("litellm_model_name") or (hidden_params or {}).get("model") + return any( + isinstance(candidate, str) and AzureFoundryModelInfo.get_azure_ai_route(candidate) == "model_router" + for candidate in (deployment_model, model) + ) + @staticmethod def get_api_base(api_base: str | None = None) -> str | None: return api_base or litellm.api_base or get_secret_str("AZURE_AI_API_BASE") diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e5714ef66fb..37a947f7f19 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1136,24 +1136,25 @@ async def open_sse_before_first_byte( ) -def _is_azure_model_router_request(model: str) -> bool: +def _is_azure_model_router_request(model: str, hidden_params: Mapping[str, object] | None = None) -> bool: """ - Check if the requested model is an Azure Model Router. + Check if a request went down the Azure Model Router route. - Azure Model Router models follow the pattern: - - azure_ai/model_router/ - - azure_ai/model-router - - model_router/ - - model-router + ``model`` here is what the *client* sent, a model group alias with no ``model_router/`` + prefix, so matching on it alone only works when the operator happened to put "model-router" + in the alias. Where the response is in hand its stamp answers this outright, so callers + should pass ``hidden_params``. Args: model: The requested model name + hidden_params: ``_hidden_params`` from the response, when the caller has it Returns: bool: True if this is an Azure Model Router request """ - model_lower: Final = model.lower() - return "model-router" in model_lower or "model_router" in model_lower + from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo + + return AzureFoundryModelInfo.is_model_router_call(model=model, hidden_params=hidden_params) def _override_openai_response_model( @@ -1221,7 +1222,7 @@ def _override_openai_response_model( return # Check if this is an Azure Model Router request - if so, preserve the actual model used - if _is_azure_model_router_request(requested_model): + if _is_azure_model_router_request(requested_model, hidden_params): verbose_proxy_logger.debug( "%s: Azure Model Router detected - preserving actual model used from response instead of overriding to router model.", log_context, diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 82de634b488..a0f7de6b320 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -6,9 +6,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path import time @@ -277,9 +275,7 @@ def test_response_cost_calculator_uses_router_model_id_from_litellm_metadata(): assert cost is not None, "Cost should not be None" expected_cost = (10 * custom_input_cost) + (5 * custom_output_cost) - assert cost == pytest.approx( - expected_cost - ), f"Expected {expected_cost}, got {cost}" + assert cost == pytest.approx(expected_cost), f"Expected {expected_cost}, got {cost}" finally: litellm.model_cost.pop(custom_model_id, None) @@ -876,13 +872,8 @@ async def test_datadog_logger_not_shadowed_by_llm_obs(monkeypatch): # Regression check: we expect a distinct DataDogLogger, not the LLM Obs logger assert type(datadog_logger) is DataDogLogger - assert any( - isinstance(cb, DataDogLLMObsLogger) - for cb in logging_module._in_memory_loggers - ) - assert any( - type(cb) is DataDogLogger for cb in logging_module._in_memory_loggers - ) + assert any(isinstance(cb, DataDogLLMObsLogger) for cb in logging_module._in_memory_loggers) + assert any(type(cb) is DataDogLogger for cb in logging_module._in_memory_loggers) finally: logging_module._in_memory_loggers.clear() @@ -893,9 +884,7 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch): # Required env vars for Logfire integration monkeypatch.setenv("LOGFIRE_TOKEN", "test-token") - monkeypatch.setenv( - "LOGFIRE_BASE_URL", "https://logfire-api-custom.pydantic.dev" - ) # no trailing slash on purpose + monkeypatch.setenv("LOGFIRE_BASE_URL", "https://logfire-api-custom.pydantic.dev") # no trailing slash on purpose # Import after env vars are set (important if module-level caching exists) from litellm.integrations.opentelemetry import OpenTelemetry # logger class @@ -914,9 +903,7 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch): # Sanity: we got the right logger type and it is cached assert type(logger) is OpenTelemetry - assert any( - type(cb) is OpenTelemetry for cb in logging_module._in_memory_loggers - ) + assert any(type(cb) is OpenTelemetry for cb in logging_module._in_memory_loggers) # Core regression check: base URL env var should influence the exporter endpoint. # @@ -927,9 +914,7 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch): or getattr(logger, "config", None) or getattr(logger, "_otel_config", None) ) - assert ( - cfg is not None - ), "Expected OpenTelemetry logger to keep an otel config on the instance" + assert cfg is not None, "Expected OpenTelemetry logger to keep an otel config on the instance" endpoint = getattr(cfg, "endpoint", None) or getattr(cfg, "otlp_endpoint", None) assert endpoint is not None, "Expected otel config to expose the OTLP endpoint" @@ -1087,9 +1072,7 @@ async def test_logging_non_streaming_request(): # Use the filtered call for assertions call_args = calls_with_expected_input[0] - standard_logging_object = call_args.kwargs["kwargs"][ - "standard_logging_object" - ] + standard_logging_object = call_args.kwargs["kwargs"]["standard_logging_object"] assert standard_logging_object["stream"] is not True finally: # Restore original callbacks to ensure test isolation @@ -1107,18 +1090,14 @@ async def test_logging_non_streaming_request(): "agenerate_content_stream", ], ) -def test_success_handler_skips_sync_callbacks_for_async_requests( - logging_obj, async_flag -): +def test_success_handler_skips_sync_callbacks_for_async_requests(logging_obj, async_flag): """Ensure sync success callbacks are skipped when async call type flags are set.""" from litellm.integrations.custom_logger import CustomLogger class DummyLogger(CustomLogger): pass - logging_obj.stream = ( - False # simulate non-streaming request where sync callbacks would normally run - ) + logging_obj.stream = False # simulate non-streaming request where sync callbacks would normally run logging_obj.model_call_details["litellm_params"] = {async_flag: True} logging_obj.litellm_params = logging_obj.model_call_details["litellm_params"] @@ -1194,21 +1173,11 @@ def test_success_handler_runs_sync_callbacks_for_sync_requests(logging_obj, call def test_is_sync_litellm_request(): assert LitellmLogging._is_sync_litellm_request({}) is True assert LitellmLogging._is_sync_litellm_request({"acompletion": True}) is False - assert ( - LitellmLogging._is_sync_litellm_request({"allm_passthrough_route": True}) - is False - ) - assert ( - LitellmLogging._is_sync_litellm_request({"aanthropic_messages": True}) is False - ) + assert LitellmLogging._is_sync_litellm_request({"allm_passthrough_route": True}) is False + assert LitellmLogging._is_sync_litellm_request({"aanthropic_messages": True}) is False assert LitellmLogging._is_sync_litellm_request({"agenerate_content": True}) is False - assert ( - LitellmLogging._is_sync_litellm_request({"agenerate_content_stream": True}) - is False - ) - assert ( - LitellmLogging._is_sync_litellm_request({"aanthropic_messages": False}) is True - ) + assert LitellmLogging._is_sync_litellm_request({"agenerate_content_stream": True}) is False + assert LitellmLogging._is_sync_litellm_request({"aanthropic_messages": False}) is True def test_get_litellm_params_propagates_allm_passthrough_route(): @@ -1255,9 +1224,7 @@ async def test_dispatch_success_handlers_invokes_callbacks_once_for_final_stream logging_obj.model_call_details["litellm_params"] = {"acompletion": True} with ( - patch.object( - mock_callback, "async_log_success_event", new_callable=AsyncMock - ) as mock_async_log, + patch.object(mock_callback, "async_log_success_event", new_callable=AsyncMock) as mock_async_log, patch.object(mock_callback, "log_success_event") as mock_sync_log, patch.object( logging_obj, @@ -1318,9 +1285,7 @@ async def test_dispatch_success_handlers_sync_path_invokes_callback_once_for_fin with ( patch.object(mock_callback, "log_success_event") as mock_sync_log, - patch.object( - mock_callback, "async_log_success_event", new_callable=AsyncMock - ) as mock_async_log, + patch.object(mock_callback, "async_log_success_event", new_callable=AsyncMock) as mock_async_log, patch.object( logging_obj, "_success_handler_helper_fn", @@ -1362,20 +1327,14 @@ async def test_dispatch_prefer_async_handlers_runs_legacy_callbacks( logging_obj.model_call_details["litellm_params"] = {} with ( - patch.object( - logging_obj, "async_success_handler", new_callable=AsyncMock - ) as mock_async, - patch.object( - logging_obj, "success_handler", new_callable=MagicMock - ) as mock_sync, + patch.object(logging_obj, "async_success_handler", new_callable=AsyncMock) as mock_async, + patch.object(logging_obj, "success_handler", new_callable=MagicMock) as mock_sync, patch.object( logging_obj, "_should_run_sync_callbacks_for_async_calls", return_value=True, ), - patch( - "litellm.litellm_core_utils.litellm_logging.executor.submit" - ) as mock_submit, + patch("litellm.litellm_core_utils.litellm_logging.executor.submit") as mock_submit, ): await logging_obj.dispatch_success_handlers( result=result, @@ -1409,9 +1368,7 @@ async def test_dispatch_success_handlers_invokes_async_callback_for_pass_through try: with ( - patch.object( - mock_callback, "async_log_success_event", new_callable=AsyncMock - ) as mock_async_log, + patch.object(mock_callback, "async_log_success_event", new_callable=AsyncMock) as mock_async_log, patch.object(mock_callback, "log_success_event") as mock_sync_log, ): await logging_obj.dispatch_success_handlers(result={"id": "pt-1"}) @@ -1438,20 +1395,14 @@ async def test_dispatch_failure_handlers_prefer_async_does_not_submit_sync_handl logging_obj.model_call_details["litellm_params"] = {} with ( - patch.object( - logging_obj, "async_failure_handler", new_callable=AsyncMock - ) as mock_async, - patch.object( - logging_obj, "failure_handler", new_callable=MagicMock - ) as mock_sync, + patch.object(logging_obj, "async_failure_handler", new_callable=AsyncMock) as mock_async, + patch.object(logging_obj, "failure_handler", new_callable=MagicMock) as mock_sync, patch.object( logging_obj, "_should_run_sync_failure_callbacks_for_async_calls", return_value=False, ), - patch( - "litellm.litellm_core_utils.litellm_logging.executor.submit" - ) as mock_submit, + patch("litellm.litellm_core_utils.litellm_logging.executor.submit") as mock_submit, ): await logging_obj.dispatch_failure_handlers( exception, @@ -1534,12 +1485,8 @@ async def test_dispatch_failure_handlers_submits_sync_handler_for_failure_only_c patch.object(litellm, "success_callback", []), patch.object(litellm, "failure_callback", [_sync_failure_callback]), patch.object(logging_obj, "async_failure_handler", new_callable=AsyncMock), - patch.object( - logging_obj, "failure_handler", new_callable=MagicMock - ) as mock_sync, - patch( - "litellm.litellm_core_utils.litellm_logging.executor.submit" - ) as mock_submit, + patch.object(logging_obj, "failure_handler", new_callable=MagicMock) as mock_sync, + patch("litellm.litellm_core_utils.litellm_logging.executor.submit") as mock_submit, ): await logging_obj.dispatch_failure_handlers( exception, @@ -1566,15 +1513,9 @@ async def test_dispatch_failure_handlers_sync_sdk_shortcut_runs_sync_handler_inl logging_obj.model_call_details["litellm_params"] = {} with ( - patch.object( - logging_obj, "async_failure_handler", new_callable=AsyncMock - ) as mock_async, - patch.object( - logging_obj, "failure_handler", new_callable=MagicMock - ) as mock_sync, - patch( - "litellm.litellm_core_utils.litellm_logging.executor.submit" - ) as mock_submit, + patch.object(logging_obj, "async_failure_handler", new_callable=AsyncMock) as mock_async, + patch.object(logging_obj, "failure_handler", new_callable=MagicMock) as mock_sync, + patch("litellm.litellm_core_utils.litellm_logging.executor.submit") as mock_submit, ): await logging_obj.dispatch_failure_handlers( exception, @@ -1621,14 +1562,10 @@ def test_success_handler_skips_guardrail_logging_hook_when_disabled(logging_obj) event_hook=GuardrailEventHooks.logging_only, ) guardrail.should_run_guardrail = MagicMock(return_value=False) - guardrail.logging_hook = MagicMock( - return_value=(logging_obj.model_call_details, model_response) - ) + guardrail.logging_hook = MagicMock(return_value=(logging_obj.model_call_details, model_response)) dummy_logger = DummyLogger() - dummy_logger.logging_hook = MagicMock( - return_value=(logging_obj.model_call_details, model_response) - ) + dummy_logger.logging_hook = MagicMock(return_value=(logging_obj.model_call_details, model_response)) with patch.object( logging_obj, @@ -1762,11 +1699,7 @@ def test_get_request_tags_from_metadata_and_litellm_metadata(): # Test case 2: Tags in litellm_metadata only tags = StandardLoggingPayloadSetup._get_request_tags( - litellm_params={ - "litellm_metadata": { - "tags": ["litellm-metadata-tag-1", "litellm-metadata-tag-2"] - } - }, + litellm_params={"litellm_metadata": {"tags": ["litellm-metadata-tag-1", "litellm-metadata-tag-2"]}}, proxy_server_request={}, ) assert "litellm-metadata-tag-1" in tags @@ -1871,15 +1804,9 @@ def test_get_request_tags_does_not_mutate_original_tags(): user_agent_count_2 = len([t for t in tags2 if t.startswith("User-Agent:")]) user_agent_count_3 = len([t for t in tags3 if t.startswith("User-Agent:")]) - assert ( - user_agent_count_1 == 2 - ), f"Expected 2 User-Agent tags, got {user_agent_count_1}" - assert ( - user_agent_count_2 == 2 - ), f"Expected 2 User-Agent tags, got {user_agent_count_2}" - assert ( - user_agent_count_3 == 2 - ), f"Expected 2 User-Agent tags, got {user_agent_count_3}" + assert user_agent_count_1 == 2, f"Expected 2 User-Agent tags, got {user_agent_count_1}" + assert user_agent_count_2 == 2, f"Expected 2 User-Agent tags, got {user_agent_count_2}" + assert user_agent_count_3 == 2, f"Expected 2 User-Agent tags, got {user_agent_count_3}" # Verify all returned lists are independent (different objects) assert tags1 is not tags2 @@ -1912,9 +1839,7 @@ def test_get_extra_header_tags(): # Test case 3: Extra headers configured but request has no headers dict litellm.extra_spend_tag_headers = ["x-custom", "x-tenant"] - result = StandardLoggingPayloadSetup._get_extra_header_tags( - proxy_server_request={"headers": "not-a-dict"} - ) + result = StandardLoggingPayloadSetup._get_extra_header_tags(proxy_server_request={"headers": "not-a-dict"}) assert result is None # Test case 4: Extra headers configured but none match request headers @@ -2215,9 +2140,7 @@ def test_get_masked_values(): "presidio_anonymizer_api_base": None, "vertex_credentials": "{sensitive_api_key}", } - masked_values = _get_masked_values( - sensitive_object, unmasked_length=4, number_of_asterisks=4 - ) + masked_values = _get_masked_values(sensitive_object, unmasked_length=4, number_of_asterisks=4) assert masked_values["presidio_anonymizer_api_base"] is None assert masked_values["vertex_credentials"] == "{s****y}" @@ -2242,9 +2165,7 @@ async def test_e2e_generate_cold_storage_object_key_successful(): patch("litellm.integrations.s3.get_s3_object_key") as mock_get_s3_key, ): # Mock the S3 object key generation to return a predictable result - mock_get_s3_key.return_value = ( - "2025-01-15/time-10-30-45-123456_chatcmpl-test-12345.json" - ) + mock_get_s3_key.return_value = "2025-01-15/time-10-30-45-123456_chatcmpl-test-12345.json" # Call the function result = StandardLoggingPayloadSetup._generate_cold_storage_object_key( @@ -2285,16 +2206,12 @@ async def test_e2e_generate_cold_storage_object_key_with_custom_logger_s3_path() with ( patch("litellm.cold_storage_custom_logger", "s3_v2"), - patch( - "litellm.logging_callback_manager.get_active_custom_logger_for_callback_name" - ) as mock_get_logger, + patch("litellm.logging_callback_manager.get_active_custom_logger_for_callback_name") as mock_get_logger, patch("litellm.integrations.s3.get_s3_object_key") as mock_get_s3_key, ): # Setup mocks mock_get_logger.return_value = mock_custom_logger - mock_get_s3_key.return_value = ( - "storage/2025-01-15/time-10-30-45-123456_chatcmpl-test-12345.json" - ) + mock_get_s3_key.return_value = "storage/2025-01-15/time-10-30-45-123456_chatcmpl-test-12345.json" # Call the function result = StandardLoggingPayloadSetup._generate_cold_storage_object_key( @@ -2313,9 +2230,7 @@ async def test_e2e_generate_cold_storage_object_key_with_custom_logger_s3_path() ) # Verify the result - assert ( - result == "storage/2025-01-15/time-10-30-45-123456_chatcmpl-test-12345.json" - ) + assert result == "storage/2025-01-15/time-10-30-45-123456_chatcmpl-test-12345.json" @pytest.mark.asyncio @@ -2338,16 +2253,12 @@ async def test_e2e_generate_cold_storage_object_key_with_logger_no_s3_path(): with ( patch("litellm.cold_storage_custom_logger", "s3_v2"), - patch( - "litellm.logging_callback_manager.get_active_custom_logger_for_callback_name" - ) as mock_get_logger, + patch("litellm.logging_callback_manager.get_active_custom_logger_for_callback_name") as mock_get_logger, patch("litellm.integrations.s3.get_s3_object_key") as mock_get_s3_key, ): # Setup mocks mock_get_logger.return_value = mock_custom_logger - mock_get_s3_key.return_value = ( - "2025-01-15/time-10-30-45-123456_chatcmpl-test-12345.json" - ) + mock_get_s3_key.return_value = "2025-01-15/time-10-30-45-123456_chatcmpl-test-12345.json" # Call the function result = StandardLoggingPayloadSetup._generate_cold_storage_object_key( @@ -2463,9 +2374,7 @@ def test_get_usage_as_dict(): assert result == {"prompt_tokens": 20, "completion_tokens": 30} # Test case 5: response_obj with no usage key returns empty - result = StandardLoggingPayloadSetup.get_usage_as_dict( - response_obj={"id": "resp-1", "choices": []} - ) + result = StandardLoggingPayloadSetup.get_usage_as_dict(response_obj={"id": "resp-1", "choices": []}) assert result == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0} @@ -2478,26 +2387,20 @@ def test_append_system_prompt_messages(): # Test case 1: system in kwargs with existing messages kwargs = {"system": "You are a helpful assistant"} messages = [{"role": "user", "content": "Hello"}] - result = StandardLoggingPayloadSetup.append_system_prompt_messages( - kwargs=kwargs, messages=messages - ) + result = StandardLoggingPayloadSetup.append_system_prompt_messages(kwargs=kwargs, messages=messages) assert len(result) == 2 assert result[0] == {"role": "system", "content": "You are a helpful assistant"} assert result[1] == {"role": "user", "content": "Hello"} # Test case 2: system in kwargs with None messages kwargs = {"system": "You are a helpful assistant"} - result = StandardLoggingPayloadSetup.append_system_prompt_messages( - kwargs=kwargs, messages=None - ) + result = StandardLoggingPayloadSetup.append_system_prompt_messages(kwargs=kwargs, messages=None) assert len(result) == 1 assert result[0] == {"role": "system", "content": "You are a helpful assistant"} # Test case 3: system in kwargs with empty messages list kwargs = {"system": "You are a helpful assistant"} - result = StandardLoggingPayloadSetup.append_system_prompt_messages( - kwargs=kwargs, messages=[] - ) + result = StandardLoggingPayloadSetup.append_system_prompt_messages(kwargs=kwargs, messages=[]) assert len(result) == 1 assert result[0] == {"role": "system", "content": "You are a helpful assistant"} @@ -2507,24 +2410,18 @@ def test_append_system_prompt_messages(): {"role": "system", "content": "You are a helpful assistant"}, {"role": "user", "content": "Hello"}, ] - result = StandardLoggingPayloadSetup.append_system_prompt_messages( - kwargs=kwargs, messages=messages - ) + result = StandardLoggingPayloadSetup.append_system_prompt_messages(kwargs=kwargs, messages=messages) assert len(result) == 2 assert result[0] == {"role": "system", "content": "You are a helpful assistant"} # Test case 5: no system in kwargs returns messages unchanged kwargs = {} messages = [{"role": "user", "content": "Hello"}] - result = StandardLoggingPayloadSetup.append_system_prompt_messages( - kwargs=kwargs, messages=messages - ) + result = StandardLoggingPayloadSetup.append_system_prompt_messages(kwargs=kwargs, messages=messages) assert result == messages # Test case 6: None kwargs returns messages unchanged - result = StandardLoggingPayloadSetup.append_system_prompt_messages( - kwargs=None, messages=messages - ) + result = StandardLoggingPayloadSetup.append_system_prompt_messages(kwargs=None, messages=messages) assert result == messages @@ -2585,12 +2482,11 @@ async def test_async_success_handler_sets_standard_logging_object_for_pass_throu # Verify that standard_logging_object was set assert "standard_logging_object" in logging_obj.model_call_details, ( - "standard_logging_object should be set for pass-through endpoints " - "even when complete_streaming_response is None" + "standard_logging_object should be set for pass-through endpoints even when complete_streaming_response is None" + ) + assert logging_obj.model_call_details["standard_logging_object"] is not None, ( + "standard_logging_object should not be None for pass-through endpoints" ) - assert ( - logging_obj.model_call_details["standard_logging_object"] is not None - ), "standard_logging_object should not be None for pass-through endpoints" # Verify that async_complete_streaming_response was set to prevent re-processing # This is consistent with the existing code pattern for regular streaming @@ -2598,15 +2494,13 @@ async def test_async_success_handler_sets_standard_logging_object_for_pass_throu "async_complete_streaming_response should be set to prevent re-processing, " "consistent with the existing code pattern" ) - assert ( - logging_obj.model_call_details["async_complete_streaming_response"] is result - ), "async_complete_streaming_response should be set to the result" + assert logging_obj.model_call_details["async_complete_streaming_response"] is result, ( + "async_complete_streaming_response should be set to the result" + ) # Verify that response_cost is set to None (cost calculation not possible for pass-through) # This is consistent with the error handling in the non-pass-through code path - assert ( - "response_cost" in logging_obj.model_call_details - ), "response_cost should be set for pass-through endpoints" + assert "response_cost" in logging_obj.model_call_details, "response_cost should be set for pass-through endpoints" assert logging_obj.model_call_details["response_cost"] is None, ( "response_cost should be None for pass-through endpoints since " "StandardPassThroughResponseObject doesn't have standard usage info" @@ -2665,14 +2559,10 @@ async def test_async_success_handler_prevents_reprocessing_for_pass_through_endp # Verify first call set the values assert "standard_logging_object" in logging_obj.model_call_details assert "async_complete_streaming_response" in logging_obj.model_call_details - first_standard_logging_object = logging_obj.model_call_details[ - "standard_logging_object" - ] + first_standard_logging_object = logging_obj.model_call_details["standard_logging_object"] # Second call - should return early due to async_complete_streaming_response guard - with patch.object( - logging_obj, "get_combined_callback_list", return_value=[] - ) as mock_callbacks: + with patch.object(logging_obj, "get_combined_callback_list", return_value=[]) as mock_callbacks: await logging_obj.async_success_handler( result=result, start_time=start_time, @@ -2683,10 +2573,9 @@ async def test_async_success_handler_prevents_reprocessing_for_pass_through_endp mock_callbacks.assert_not_called() # Verify standard_logging_object wasn't modified by second call - assert ( - logging_obj.model_call_details["standard_logging_object"] - is first_standard_logging_object - ), "standard_logging_object should not be modified on re-processing" + assert logging_obj.model_call_details["standard_logging_object"] is first_standard_logging_object, ( + "standard_logging_object should not be modified on re-processing" + ) @pytest.mark.asyncio @@ -2725,9 +2614,7 @@ async def test_async_success_handler_sets_standard_logging_object_for_streaming_ } # Create a pass-through response object (simulating unparseable streaming response) - result = StandardPassThroughResponseObject( - response='data: {"chunk": 1}\ndata: {"chunk": 2}\ndata: [DONE]' - ) + result = StandardPassThroughResponseObject(response='data: {"chunk": 1}\ndata: {"chunk": 2}\ndata: [DONE]') start_time = datetime.now() end_time = datetime.now() @@ -2747,9 +2634,9 @@ async def test_async_success_handler_sets_standard_logging_object_for_streaming_ "standard_logging_object should be set for streaming pass-through endpoints " "even when the response cannot be parsed into a ModelResponse" ) - assert ( - logging_obj.model_call_details["standard_logging_object"] is not None - ), "standard_logging_object should not be None for streaming pass-through endpoints" + assert logging_obj.model_call_details["standard_logging_object"] is not None, ( + "standard_logging_object should not be None for streaming pass-through endpoints" + ) def test_get_error_information_error_code_priority(): @@ -2791,30 +2678,22 @@ def test_get_error_information_error_code_priority(): self.message = message super().__init__(message) - both_exception = BothAttributesException( - code="400", status_code=500, message="Bad Request" - ) + both_exception = BothAttributesException(code="400", status_code=500, message="Bad Request") result = StandardLoggingPayloadSetup.get_error_information(both_exception) assert result["error_code"] == "400" # Should prefer 'code' over 'status_code' # Test case 4: Exception with 'code' as empty string - should fall back to 'status_code' - empty_code_exception = BothAttributesException( - code="", status_code=404, message="Not Found" - ) + empty_code_exception = BothAttributesException(code="", status_code=404, message="Not Found") result = StandardLoggingPayloadSetup.get_error_information(empty_code_exception) assert result["error_code"] == "404" # Should fall back to status_code # Test case 5: Exception with 'code' as "None" string - should fall back to 'status_code' - none_string_exception = BothAttributesException( - code="None", status_code=503, message="Service Unavailable" - ) + none_string_exception = BothAttributesException(code="None", status_code=503, message="Service Unavailable") result = StandardLoggingPayloadSetup.get_error_information(none_string_exception) assert result["error_code"] == "503" # Should fall back to status_code # Test case 6: Exception with 'code' as None - should fall back to 'status_code' - none_code_exception = BothAttributesException( - code=None, status_code=401, message="Unauthorized" - ) + none_code_exception = BothAttributesException(code=None, status_code=401, message="Unauthorized") result = StandardLoggingPayloadSetup.get_error_information(none_code_exception) assert result["error_code"] == "401" # Should fall back to status_code @@ -2863,9 +2742,7 @@ def test_get_error_information_prefers_message_attribute_over_str(): ) result = StandardLoggingPayloadSetup.get_error_information(exc) - assert ( - result["error_message"] == msg - ), f"expected message from .message attribute, got {result['error_message']!r}" + assert result["error_message"] == msg, f"expected message from .message attribute, got {result['error_message']!r}" assert result["error_code"] == "401" assert result["error_class"] == "ProxyExceptionLike" @@ -2940,8 +2817,7 @@ def test_get_error_information_preserves_explicit_empty_message(): exc = ProxyExceptionLike(message="", code=500) result = StandardLoggingPayloadSetup.get_error_information(exc) assert result["error_message"] == "", ( - "explicit empty .message must survive verbatim; got " - f"{result['error_message']!r}" + f"explicit empty .message must survive verbatim; got {result['error_message']!r}" ) @@ -3204,9 +3080,7 @@ def test_process_hidden_params_recalculates_cost_after_failure_handler_zero(): choices=[{"message": {"role": "assistant", "content": "ok"}}], usage=Usage(prompt_tokens=9698, completion_tokens=30, total_tokens=9728), ) - logging_obj._process_hidden_params_and_response_cost( - result, datetime.now(), datetime.now() - ) + logging_obj._process_hidden_params_and_response_cost(result, datetime.now(), datetime.now()) cost = logging_obj.model_call_details.get("response_cost") assert cost is not None and cost > 0 @@ -3230,9 +3104,7 @@ def test_process_hidden_params_preserves_zero_cost_in_hidden_params(): litellm_call_id="test-hidden-zero-cost", function_id="test-hidden-zero-cost", ) - logging_obj.model_call_details["litellm_params"] = { - "model": "gemini-2.5-flash-lite" - } + logging_obj.model_call_details["litellm_params"] = {"model": "gemini-2.5-flash-lite"} logging_obj.optional_params = {} result = ModelResponse( @@ -3242,9 +3114,7 @@ def test_process_hidden_params_preserves_zero_cost_in_hidden_params(): ) result._hidden_params = {"response_cost": 0.0} - logging_obj._process_hidden_params_and_response_cost( - result, datetime.now(), datetime.now() - ) + logging_obj._process_hidden_params_and_response_cost(result, datetime.now(), datetime.now()) assert logging_obj.model_call_details.get("response_cost") == 0.0 slo = logging_obj.model_call_details.get("standard_logging_object") or {} @@ -3293,9 +3163,7 @@ def test_process_hidden_params_uses_hidden_params_cost_after_failure_handler_zer ) result._hidden_params = {"response_cost": passthrough_cost} - logging_obj._process_hidden_params_and_response_cost( - result, datetime.now(), datetime.now() - ) + logging_obj._process_hidden_params_and_response_cost(result, datetime.now(), datetime.now()) assert logging_obj.model_call_details.get("response_cost") == passthrough_cost slo = logging_obj.model_call_details.get("standard_logging_object") or {} @@ -3352,9 +3220,7 @@ def test_function_setup_litellm_metadata_populates_metadata(): assert litellm_metadata.get("user_api_key_hash") == test_api_key_hash # metadata should be a COPY, not an alias — mutating one must not affect the other - assert ( - metadata is not litellm_metadata - ), "litellm_params['metadata'] should be a copy, not the same object" + assert metadata is not litellm_metadata, "litellm_params['metadata'] should be a copy, not the same object" def test_function_setup_litellm_metadata_guardrail_writes_visible_after_setup(): @@ -3399,9 +3265,9 @@ def test_function_setup_litellm_metadata_guardrail_writes_visible_after_setup(): litellm_params = logging_obj.model_call_details.get("litellm_params", {}) litellm_metadata = litellm_params.get("litellm_metadata") assert litellm_metadata is not None - assert litellm_metadata.get("standard_logging_guardrail_information") == [ - guardrail_entry - ], "guardrail writes after function_setup must be visible to the logging object" + assert litellm_metadata.get("standard_logging_guardrail_information") == [guardrail_entry], ( + "guardrail writes after function_setup must be visible to the logging object" + ) assert litellm_metadata.get("applied_guardrails") == ["pam-ethical-request"] merged = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params) @@ -3570,9 +3436,7 @@ def test_failure_handler_skips_sync_callbacks_for_pass_through_requests(logging_ @pytest.mark.parametrize("call_type", ["completion", "acompletion"]) -def test_failure_handler_runs_sync_callbacks_for_non_pass_through_requests( - logging_obj, call_type -): +def test_failure_handler_runs_sync_callbacks_for_non_pass_through_requests(logging_obj, call_type): """Ensure sync failure callbacks still fire for normal (non-pass-through) requests.""" from litellm.integrations.custom_logger import CustomLogger @@ -3733,9 +3597,7 @@ def test_standard_logging_hidden_params_backfills_response_cost_without_mutating ) response._hidden_params = {"response_cost": None, "model_id": "mid-test"} - payload = logging_obj._build_standard_logging_payload( - response, datetime.now(), datetime.now() - ) + payload = logging_obj._build_standard_logging_payload(response, datetime.now(), datetime.now()) assert payload is not None assert payload["hidden_params"]["response_cost"] == 0.002 @@ -3789,10 +3651,7 @@ def test_merge_hidden_params_from_response_into_metadata_no_op_when_empty(): _hidden_params = {} logging_obj._merge_hidden_params_from_response_into_metadata(_NoHp()) - assert ( - "hidden_params" - not in logging_obj.model_call_details["litellm_params"]["metadata"] - ) + assert "hidden_params" not in logging_obj.model_call_details["litellm_params"]["metadata"] # ── StandardLoggingPayloadSetup.get_additional_headers ─────────────────────── @@ -3870,6 +3729,82 @@ def test_get_standard_logging_object_payload_includes_litellm_call_id(logging_ob assert payload["litellm_call_id"] == call_id +# ── Azure Model Router selected-model attribution ──────────────────────────── + + +def _model_router_response(selected_model: str, stamp: bool): + """A ModelResponse as AzureModelRouterConfig hands it back, with or without the stamp.""" + from litellm.llms.azure_ai.common_utils import ( + AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY, + ) + from litellm.types.utils import ModelResponse + + response = ModelResponse(model=selected_model) + response._hidden_params = {AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY: selected_model} if stamp else {} + return response + + +def test_standard_logging_payload_uses_stamped_model_router_model(logging_obj): + """ + The selected model must win off the stamp, not off "model-router" appearing in the + requested model. An operator whose model group is named anything else was invisible + to the name check, so their logs and spend rows named the router instead. + """ + import datetime + + from litellm.litellm_core_utils.litellm_logging import ( + get_standard_logging_object_payload, + ) + + now = datetime.datetime.now() + payload = get_standard_logging_object_payload( + kwargs={ + "model": "azure_ai/smart-pick", + "custom_llm_provider": "azure_ai", + "messages": [], + "litellm_params": {"metadata": {}}, + }, + init_response_obj=_model_router_response("azure_ai/grok-4-1-fast-reasoning", stamp=True), + start_time=now, + end_time=now, + logging_obj=logging_obj, + status="success", + ) + + assert payload is not None + assert payload["model"] == "azure_ai/grok-4-1-fast-reasoning" + + +def test_standard_logging_payload_keeps_requested_model_without_router_stamp(logging_obj): + """ + Control for the test above: an ordinary azure_ai deployment is unaffected, so the stamp + is what redirects attribution rather than the response model winning unconditionally. + """ + import datetime + + from litellm.litellm_core_utils.litellm_logging import ( + get_standard_logging_object_payload, + ) + + now = datetime.datetime.now() + payload = get_standard_logging_object_payload( + kwargs={ + "model": "azure_ai/smart-pick", + "custom_llm_provider": "azure_ai", + "messages": [], + "litellm_params": {"metadata": {}}, + }, + init_response_obj=_model_router_response("azure_ai/grok-4-1-fast-reasoning", stamp=False), + start_time=now, + end_time=now, + logging_obj=logging_obj, + status="success", + ) + + assert payload is not None + assert payload["model"] == "azure_ai/smart-pick" + + def _make_dict_logging_obj(): """Build a Logging instance configured for a non-streaming dict result.""" obj = LitellmLogging( @@ -3905,9 +3840,7 @@ def test_success_handler_computes_cost_for_dict_response(): "_build_standard_logging_payload", return_value={"response_cost": expected_cost}, ), - patch( - "litellm.litellm_core_utils.litellm_logging.emit_standard_logging_payload" - ), + patch("litellm.litellm_core_utils.litellm_logging.emit_standard_logging_payload"), patch.object( logging_obj, "_is_recognized_call_type_for_logging", @@ -3944,9 +3877,7 @@ def test_success_handler_preserves_precomputed_cost_for_dict_response(): "_build_standard_logging_payload", return_value={"response_cost": precomputed_cost}, ), - patch( - "litellm.litellm_core_utils.litellm_logging.emit_standard_logging_payload" - ), + patch("litellm.litellm_core_utils.litellm_logging.emit_standard_logging_payload"), patch.object( logging_obj, "_is_recognized_call_type_for_logging", @@ -3985,9 +3916,7 @@ def test_success_handler_unified_helper_runs_for_typed_results(): "_build_standard_logging_payload", return_value={"response_cost": expected_cost}, ), - patch( - "litellm.litellm_core_utils.litellm_logging.emit_standard_logging_payload" - ), + patch("litellm.litellm_core_utils.litellm_logging.emit_standard_logging_payload"), patch.object( logging_obj, "_is_recognized_call_type_for_logging", @@ -4042,9 +3971,7 @@ class TestFirstApiCallStartTimeSetOnce: assert first == obj.model_call_details["api_call_start_time"] # Set on the logging object only — user metadata untouched. assert user_meta == {} - assert ( - "first_api_call_start_time" not in obj.model_call_details["litellm_params"] - ) + assert "first_api_call_start_time" not in obj.model_call_details["litellm_params"] time.sleep(0.002) # ensure a distinct retry timestamp obj.pre_call(input="hi", api_key="sk-test") @@ -4061,18 +3988,16 @@ def test_get_error_information_for_logging_payload_ignores_spoofed_disconnect_wi baseline = StandardLoggingPayloadSetup.get_error_information( original_exception=ValueError("provider failure"), ) - error_information, error_str = ( - StandardLoggingPayloadSetup.get_error_information_for_logging_payload( - metadata={ - "error_information": { - "error_code": "499", - "error_message": "Client disconnected the request", - "error_class": "ClientDisconnected", - } - }, - original_exception=ValueError("provider failure"), - error_str="provider failure", - ) + error_information, error_str = StandardLoggingPayloadSetup.get_error_information_for_logging_payload( + metadata={ + "error_information": { + "error_code": "499", + "error_message": "Client disconnected the request", + "error_class": "ClientDisconnected", + } + }, + original_exception=ValueError("provider failure"), + error_str="provider failure", ) assert error_information == baseline assert error_str == "provider failure" @@ -4086,22 +4011,18 @@ def test_get_error_information_for_logging_payload_client_disconnect(): "error_message": "Client disconnected the request", "error_class": "ClientDisconnected", } - error_information, error_str = ( - StandardLoggingPayloadSetup.get_error_information_for_logging_payload( - metadata={"client_disconnected": True, "error_information": custom_error}, - original_exception=None, - error_str=None, - ) + error_information, error_str = StandardLoggingPayloadSetup.get_error_information_for_logging_payload( + metadata={"client_disconnected": True, "error_information": custom_error}, + original_exception=None, + error_str=None, ) assert error_information == custom_error assert error_str == "Client disconnected the request" - error_information, error_str = ( - StandardLoggingPayloadSetup.get_error_information_for_logging_payload( - metadata={"client_disconnected": True}, - original_exception=None, - error_str="existing error", - ) + error_information, error_str = StandardLoggingPayloadSetup.get_error_information_for_logging_payload( + metadata={"client_disconnected": True}, + original_exception=None, + error_str="existing error", ) assert error_information["error_code"] == "499" assert error_str == "existing error" @@ -4109,12 +4030,10 @@ def test_get_error_information_for_logging_payload_client_disconnect(): baseline = StandardLoggingPayloadSetup.get_error_information( original_exception=None, ) - error_information, error_str = ( - StandardLoggingPayloadSetup.get_error_information_for_logging_payload( - metadata={}, - original_exception=None, - error_str=None, - ) + error_information, error_str = StandardLoggingPayloadSetup.get_error_information_for_logging_payload( + metadata={}, + original_exception=None, + error_str=None, ) assert error_information == baseline assert error_str is None @@ -4149,9 +4068,7 @@ def test_get_error_information_prefers_message_attribute_over_empty_str(): def __str__(self): return "" - info = StandardLoggingPayloadSetup.get_error_information( - original_exception=_SilentExc() - ) + info = StandardLoggingPayloadSetup.get_error_information(original_exception=_SilentExc()) assert info["error_message"] == "real failure detail" assert info["error_code"] == "401" @@ -4182,9 +4099,7 @@ def _responses_api_response_with_text(text="hello world"): type="message", role="assistant", status="completed", - content=[ - ResponseOutputText(annotations=[], text=text, type="output_text") - ], + content=[ResponseOutputText(annotations=[], text=text, type="output_text")], ) ], usage=ResponseAPIUsage(input_tokens=11, output_tokens=7, total_tokens=18), @@ -4199,9 +4114,7 @@ def _responses_api_response_with_text(text="hello world"): ("ResponseFailedEvent", "response.failed"), ], ) -def test_handle_anthropic_messages_response_logging_translates_terminal_responses_api_event( - event_cls, event_type -): +def test_handle_anthropic_messages_response_logging_translates_terminal_responses_api_event(event_cls, event_type): """Regression for #28595 / #28943. When anthropic_messages routes to the OpenAI Responses backend and stream=True, success_handler receives a terminal Responses API event. The handler must translate it to a ModelResponse whose choices carry @@ -4240,10 +4153,7 @@ def test_handle_anthropic_messages_response_logging_passes_model_response_throug """Anthropic-native path already yields a ModelResponse; it must be returned unchanged.""" logging_obj = _anthropic_messages_logging_obj() model_response = ModelResponse() - assert ( - logging_obj._handle_anthropic_messages_response_logging(result=model_response) - is model_response - ) + assert logging_obj._handle_anthropic_messages_response_logging(result=model_response) is model_response def test_handle_anthropic_messages_response_logging_degrades_on_unparseable_responses_payload(): @@ -4539,9 +4449,7 @@ def test_non_image_response_has_no_output_image_count(logging_obj): def test_zero_token_video_usage_preserves_duration_seconds(logging_obj): """Video usage bills by duration; the payload must keep duration_seconds even with zero tokens.""" - payload = _build_payload_for_media_response( - logging_obj, {"id": "video-1", "usage": {"duration_seconds": 4.0}} - ) + payload = _build_payload_for_media_response(logging_obj, {"id": "video-1", "usage": {"duration_seconds": 4.0}}) assert payload is not None assert payload["metadata"]["usage_object"]["duration_seconds"] == 4.0 diff --git a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py index 900372f3e54..1f684ecc3b5 100644 --- a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py +++ b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -5,9 +5,7 @@ from unittest.mock import MagicMock, patch import pytest -sys.path.insert( - 0, os.path.abspath("../../../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../../../..")) # Adds the parent directory to the system path from litellm.llms.azure_ai.azure_model_router.transformation import ( AzureModelRouterConfig, ) @@ -120,9 +118,7 @@ def test_azure_ai_grok_stop_parameter_handling(): # Test supported parameters for Grok models for model in ("grok-4-fast", "grok-4.3"): grok_params = config.get_supported_openai_params(model) - assert ( - "stop" not in grok_params - ), "Grok models should not support stop parameter" + assert "stop" not in grok_params, "Grok models should not support stop parameter" # Test supported parameters for non-Grok models gpt_params = config.get_supported_openai_params("gpt-4") @@ -201,11 +197,84 @@ def test_azure_model_router_response_shows_actual_model(): # Verify that the response contains the actual model used, not the router model assert result.model == "azure_ai/gpt-5-nano-2025-08-07", ( - f"Expected model to be 'azure_ai/gpt-5-nano-2025-08-07' (actual model used), " - f"but got '{result.model}'" + f"Expected model to be 'azure_ai/gpt-5-nano-2025-08-07' (actual model used), but got '{result.model}'" ) +def test_azure_model_router_stamps_selected_model_on_hidden_params(): + """ + The selected model must be stamped on _hidden_params, not left for downstream code to + re-derive by looking for "model-router" in the model string. Deployments whose alias + does not contain that text are invisible to the string check. + """ + from httpx import Response + + from litellm.llms.azure_ai.common_utils import ( + AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY, + AzureFoundryModelInfo, + ) + from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj + from litellm.types.utils import ModelResponse + + raw_response_json = { + "id": "chatcmpl-test456", + "object": "chat.completion", + "created": 1234567890, + "model": "grok-4-1-fast-reasoning", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "pong"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + + mock_response = MagicMock(spec=Response) + mock_response.json.return_value = raw_response_json + mock_response.text = json.dumps(raw_response_json) + mock_response.headers = {} + + logging_obj = MagicMock(spec=LiteLLMLoggingObj) + logging_obj.post_call = MagicMock() + logging_obj.model_call_details = {} + + result = AzureModelRouterConfig().transform_response( + model="smart-pick", + raw_response=mock_response, + model_response=ModelResponse(), + logging_obj=logging_obj, + request_data={}, + messages=[{"role": "user", "content": "Reply with just pong"}], + optional_params={}, + litellm_params={"model": "azure_ai/model_router/smart-pick"}, + encoding=None, + api_key="test-key", + json_mode=False, + ) + + assert result._hidden_params[AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY] == result.model + assert result._hidden_params[AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY] == "azure_ai/grok-4-1-fast-reasoning" + assert AzureFoundryModelInfo.get_model_router_selected_model(result._hidden_params) == ( + "azure_ai/grok-4-1-fast-reasoning" + ) + assert AzureFoundryModelInfo.is_model_router_call(model="smart-pick", hidden_params=result._hidden_params) is True + + +def test_azure_model_router_stamp_does_not_leak_across_responses(): + """ + ModelResponse declares _hidden_params as a class-level dict, so the stamp has to be written + as a fresh dict. Mutating in place would bleed the selected model into unrelated responses. + """ + from litellm.llms.azure_ai.common_utils import AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY + from litellm.types.utils import ModelResponse + + untouched = ModelResponse() + + assert AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY not in (untouched._hidden_params or {}) + + def test_drop_tool_level_extra_fields_strips_copilot_mcp_server_name(): """ Regression test: Azure AI returns 400 when tools contain copilot_mcp_server_name. @@ -226,14 +295,10 @@ def test_drop_tool_level_extra_fields_strips_copilot_mcp_server_name(): mock_response.text = error_text mock_response.json.return_value = json.loads(error_text) mock_response.status_code = 400 - e = httpx.HTTPStatusError( - message="400", request=MagicMock(), response=mock_response - ) + e = httpx.HTTPStatusError(message="400", request=MagicMock(), response=mock_response) assert config._error_has_tool_level_extra_fields(error_text) is True - assert ( - config.should_retry_llm_api_inside_llm_translation_on_http_error(e, {}) is True - ) + assert config.should_retry_llm_api_inside_llm_translation_on_http_error(e, {}) is True request_data = { "model": "FW-Kimi-K2.6", @@ -354,9 +419,7 @@ def test_azure_ai_stripping_does_not_mutate_caller_messages(): { "role": "assistant", "content": "I can help.", - "thinking_blocks": [ - {"type": "thinking", "thinking": "Reading the file.", "signature": "sig"} - ], + "thinking_blocks": [{"type": "thinking", "thinking": "Reading the file.", "signature": "sig"}], "provider_specific_fields": {"thought_signature": "sig-top"}, "tool_calls": [ { diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 716fba370df..562046b2f81 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -126,16 +126,12 @@ class TestProxyBaseLLMRequestProcessing: assert json.loads(result.body) == guardrailed_body @pytest.mark.asyncio - async def test_handle_non_streaming_allm_passthrough_route_forwards_upstream_headers( - self, monkeypatch - ): + async def test_handle_non_streaming_allm_passthrough_route_forwards_upstream_headers(self, monkeypatch): """The guardrail JSON path must forward upstream response headers (e.g. x-amzn-requestid) alongside the x-litellm-* headers, matching the non-guardrail passthrough path, while dropping length headers that no longer match the rewritten body.""" - processing_obj = ProxyBaseLLMRequestProcessing( - data={"custom_llm_provider": "bedrock"} - ) + processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"}) monkeypatch.setattr( processing_obj, "_has_post_call_guardrails_for_passthrough", @@ -175,14 +171,10 @@ class TestProxyBaseLLMRequestProcessing: assert result.headers["content-length"] == str(len(result.body)) @pytest.mark.asyncio - async def test_handle_event_stream_allm_passthrough_route_forwards_upstream_headers( - self, monkeypatch - ): + async def test_handle_event_stream_allm_passthrough_route_forwards_upstream_headers(self, monkeypatch): """The guardrail event-stream branch must also forward upstream response headers alongside the x-litellm-* headers.""" - processing_obj = ProxyBaseLLMRequestProcessing( - data={"custom_llm_provider": "bedrock"} - ) + processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"}) monkeypatch.setattr( processing_obj, "_has_post_call_guardrails_for_passthrough", @@ -224,15 +216,11 @@ class TestProxyBaseLLMRequestProcessing: assert result.headers["x-litellm-call-id"] == "test-call-id" @pytest.mark.asyncio - async def test_handle_non_streaming_allm_passthrough_route_applies_response_headers_hook( - self, monkeypatch - ): + async def test_handle_non_streaming_allm_passthrough_route_applies_response_headers_hook(self, monkeypatch): """Guardrailed non-streaming passthrough responses must include headers injected by post_call_response_headers_hook, matching the headers a non-guardrailed passthrough response would carry.""" - processing_obj = ProxyBaseLLMRequestProcessing( - data={"custom_llm_provider": "bedrock"} - ) + processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"}) monkeypatch.setattr( processing_obj, "_has_post_call_guardrails_for_passthrough", @@ -251,9 +239,7 @@ class TestProxyBaseLLMRequestProcessing: return kwargs["response"] proxy_logging_obj.post_call_success_hook = fake_post_call_success_hook - proxy_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value={"x-litellm-custom": "from-hook"} - ) + proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={"x-litellm-custom": "from-hook"}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=upstream, @@ -2221,6 +2207,52 @@ class TestOverrideOpenAIResponseModel: assert response_obj.model == actual_model_used assert response_obj.model != requested_model + def test_override_model_preserves_model_router_model_for_alias_without_router_in_name(self): + """ + The client sends a model group alias, which carries no model_router/ prefix, so the + name check alone only fires when the operator happened to put "model-router" in the + alias. With the stamp on the response the actual model survives whatever it is named. + """ + from litellm.llms.azure_ai.common_utils import ( + AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY, + ) + + requested_model = "smart-pick" + actual_model_used = "azure_ai/grok-4-1-fast-reasoning" + + response_obj = MagicMock() + response_obj.model = actual_model_used + response_obj._hidden_params = { + "additional_headers": {}, + AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY: actual_model_used, + } + + _override_openai_response_model( + response_obj=response_obj, + requested_model=requested_model, + log_context="test_context", + ) + assert response_obj.model == actual_model_used + + def test_override_model_still_restamps_non_router_alias_without_stamp(self): + """ + Control for the test above: absent the stamp, an ordinary deployment keeps being + restamped to the requested model, so the stamp is doing the work rather than the + preserve branch having gone unconditional. + """ + requested_model = "smart-pick" + + response_obj = MagicMock() + response_obj.model = "azure_ai/grok-4-1-fast-reasoning" + response_obj._hidden_params = {"additional_headers": {}} + + _override_openai_response_model( + response_obj=response_obj, + requested_model=requested_model, + log_context="test_context", + ) + assert response_obj.model == requested_model + def test_override_model_uses_winning_model_for_fastest_response(self): """ Test that when fastest_response batch completion is used with a @@ -2793,9 +2825,7 @@ class TestStreamCloseOnDisconnect: finally: closed.set() - response = _UpstreamClosingStreamingResponse( - body(), media_type="text/event-stream" - ) + response = _UpstreamClosingStreamingResponse(body(), media_type="text/event-stream") async def receive(): await asyncio.Event().wait() @@ -2826,9 +2856,7 @@ class TestStreamCloseOnDisconnect: finally: closed.set() - response = _UpstreamClosingStreamingResponse( - body(), media_type="text/event-stream" - ) + response = _UpstreamClosingStreamingResponse(body(), media_type="text/event-stream") async def receive(): await disconnected.wait() @@ -2899,9 +2927,7 @@ class TestStreamCloseOnDisconnect: finally: inner_closed.set() - response = await create_response( - generator=wrapped(), media_type="text/event-stream", headers={} - ) + response = await create_response(generator=wrapped(), media_type="text/event-stream", headers={}) async def receive(): await asyncio.Event().wait() @@ -3097,9 +3123,7 @@ class TestStreamCloseOnDisconnect: with pytest.raises(_ClientDisconnectedBeforeFirstChunk): await asyncio.wait_for( - _buffer_first_chunk_honoring_disconnect( - AcloseRaises(), request=self._request_that_disconnects() - ), + _buffer_first_chunk_honoring_disconnect(AcloseRaises(), request=self._request_that_disconnects()), timeout=5, ) @@ -3115,9 +3139,7 @@ class TestStreamCloseOnDisconnect: with pytest.raises(_ClientDisconnectedBeforeFirstChunk): await asyncio.wait_for( - _buffer_first_chunk_honoring_disconnect( - blocking_gen(), request=self._request_that_disconnects() - ), + _buffer_first_chunk_honoring_disconnect(blocking_gen(), request=self._request_that_disconnects()), timeout=5, ) assert closed.is_set() @@ -3133,9 +3155,7 @@ class TestHandleLLMApiExceptionRetryAfter: user_api_key_dict = UserAPIKeyAuth(api_key="sk-test") proxy_logging_obj = MagicMock() proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) - proxy_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value=callback_headers or {} - ) + proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value=callback_headers or {}) try: await processor._handle_llm_api_exception( @@ -3187,9 +3207,7 @@ class TestHandleLLMApiExceptionRetryAfter: enable_pre_call_checks=False, cooldown_list=[], ) - proxy_exc = await self._invoke( - exc, callback_headers={"retry-after": "", "x-custom": "1"} - ) + proxy_exc = await self._invoke(exc, callback_headers={"retry-after": "", "x-custom": "1"}) assert proxy_exc.headers["retry-after"] == "43" assert proxy_exc.headers["x-custom"] == "1" @@ -3385,9 +3403,7 @@ class TestDisconnectGatherCleanup: return Request(scope={"type": "http", "headers": []}, receive=receive) @pytest.mark.asyncio - async def test_base_process_llm_request_raises_499_on_client_disconnect( - self, monkeypatch - ): + async def test_base_process_llm_request_raises_499_on_client_disconnect(self, monkeypatch): """With cancel_on_disconnect enabled, base_process_llm_request returns 499.""" import asyncio @@ -3416,9 +3432,7 @@ class TestDisconnectGatherCleanup: "common_processing_pre_call_logic", AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)), ) - monkeypatch.setattr( - processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False) - ) + monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)) with pytest.raises(HTTPException) as exc_info: await processing_obj.base_process_llm_request( @@ -3436,9 +3450,7 @@ class TestDisconnectGatherCleanup: assert "disconnected" in exc_info.value.detail.lower() @pytest.mark.asyncio - async def test_base_process_llm_request_reraises_cancelled_error_without_client_disconnect( - self, monkeypatch - ): + async def test_base_process_llm_request_reraises_cancelled_error_without_client_disconnect(self, monkeypatch): import asyncio import litellm.proxy.common_request_processing as cpr @@ -3463,9 +3475,7 @@ class TestDisconnectGatherCleanup: "common_processing_pre_call_logic", AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)), ) - monkeypatch.setattr( - processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False) - ) + monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)) monkeypatch.setattr( cpr, "route_request", @@ -3526,9 +3536,7 @@ class TestDisconnectGatherCleanup: "common_processing_pre_call_logic", AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)), ) - monkeypatch.setattr( - processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False) - ) + monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)) with pytest.raises(HTTPException): await processing_obj.base_process_llm_request( @@ -3579,9 +3587,7 @@ class TestDisconnectGatherCleanup: assert task.done() @pytest.mark.asyncio - async def test_base_process_llm_request_preserves_llm_error_after_gather( - self, monkeypatch - ): + async def test_base_process_llm_request_preserves_llm_error_after_gather(self, monkeypatch): import litellm.proxy.common_request_processing as cpr from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing @@ -3610,9 +3616,7 @@ class TestDisconnectGatherCleanup: "common_processing_pre_call_logic", AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)), ) - monkeypatch.setattr( - processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False) - ) + monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)) mock_request = MagicMock(spec=Request) mock_request.is_disconnected = AsyncMock(return_value=False) @@ -3649,19 +3653,13 @@ class TestStreamingClientDisconnectLogging: "litellm_params": {"metadata": {}}, } - recorded = await _record_streaming_client_disconnect_if_needed( - mock_request, request_data - ) + recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data) assert recorded is True assert request_data["metadata"]["client_disconnected"] is True + assert request_data["metadata"]["error_information"]["error_code"] == "499" assert ( - request_data["metadata"]["error_information"]["error_code"] == "499" - ) - assert ( - mock_logging_obj.model_call_details["litellm_params"]["metadata"][ - "error_information" - ]["error_code"] + mock_logging_obj.model_call_details["litellm_params"]["metadata"]["error_information"]["error_code"] == "499" ) @@ -3675,9 +3673,7 @@ class TestStreamingClientDisconnectLogging: mock_request.is_disconnected = AsyncMock(return_value=False) request_data = {"metadata": {}} - recorded = await _record_streaming_client_disconnect_if_needed( - mock_request, request_data - ) + recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data) assert recorded is False assert "client_disconnected" not in request_data["metadata"] @@ -3702,22 +3698,12 @@ class TestStreamingClientDisconnectLogging: "litellm_params": {"metadata": {}}, } - recorded = await _record_streaming_client_disconnect_if_needed( - mock_request, request_data - ) + recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data) assert recorded is True assert request_data["metadata"]["client_disconnected"] is True - assert ( - mock_logging_obj.model_call_details["litellm_params"]["metadata"][ - "client_disconnected" - ] - is True - ) - assert ( - mock_logging_obj.model_call_details["metadata"]["client_disconnected"] - is True - ) + assert mock_logging_obj.model_call_details["litellm_params"]["metadata"]["client_disconnected"] is True + assert mock_logging_obj.model_call_details["metadata"]["client_disconnected"] is True @pytest.mark.asyncio async def test_record_streaming_client_disconnect_handles_none_request_data_metadata(self): @@ -3733,15 +3719,11 @@ class TestStreamingClientDisconnectLogging: "litellm_params": {"metadata": None}, } - recorded = await _record_streaming_client_disconnect_if_needed( - mock_request, request_data - ) + recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data) assert recorded is True assert request_data["metadata"]["client_disconnected"] is True - assert ( - request_data["litellm_params"]["metadata"]["client_disconnected"] is True - ) + assert request_data["litellm_params"]["metadata"]["client_disconnected"] is True @pytest.mark.asyncio async def test_apply_client_disconnect_metadata_none_returns_early(self): @@ -3752,9 +3734,7 @@ class TestStreamingClientDisconnectLogging: _apply_client_disconnect_metadata(None) @pytest.mark.asyncio - async def test_finalize_streaming_generator_cleanup_fires_deferred_logging( - self, monkeypatch - ): + async def test_finalize_streaming_generator_cleanup_fires_deferred_logging(self, monkeypatch): from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, ) @@ -3786,9 +3766,7 @@ class TestStreamingClientDisconnectLogging: assert request_data["metadata"]["error_information"]["error_code"] == "499" @pytest.mark.asyncio - async def test_finalize_streaming_generator_cleanup_skips_disconnect_after_completion( - self, monkeypatch - ): + async def test_finalize_streaming_generator_cleanup_skips_disconnect_after_completion(self, monkeypatch): from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, ) @@ -3818,9 +3796,7 @@ class TestStreamingClientDisconnectLogging: assert "client_disconnected" not in request_data["metadata"] @pytest.mark.asyncio - async def test_async_streaming_data_generator_records_499_on_early_aclose( - self, monkeypatch - ): + async def test_async_streaming_data_generator_records_499_on_early_aclose(self, monkeypatch): from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, ) @@ -3835,9 +3811,7 @@ class TestStreamingClientDisconnectLogging: yield {"choices": [{"delta": {"content": " there"}}]} mock_proxy_logging = MagicMock(spec=ProxyLogging) - mock_proxy_logging.async_post_call_streaming_iterator_hook = ( - mock_streaming_iterator - ) + mock_proxy_logging.async_post_call_streaming_iterator_hook = mock_streaming_iterator ProxyLogging._callback_capabilities_cache.clear() mock_request = MagicMock(spec=Request) @@ -3848,9 +3822,7 @@ class TestStreamingClientDisconnectLogging: "model": "gemini-2.0-flash", "metadata": {}, "litellm_params": {"metadata": {}}, - "litellm_logging_obj": MagicMock( - model_call_details={"metadata": {}, "litellm_params": {}} - ), + "litellm_logging_obj": MagicMock(model_call_details={"metadata": {}, "litellm_params": {}}), } gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( @@ -3869,6 +3841,8 @@ class TestStreamingClientDisconnectLogging: assert request_data["metadata"]["error_information"]["error_code"] == "499" ProxyLogging._callback_capabilities_cache.clear() + + class TestCancelOnDisconnect: """ Coverage for the opt-in `general_settings.cancel_on_disconnect` flag: @@ -3895,23 +3869,17 @@ class TestCancelOnDisconnect: llm_call = asyncio.get_running_loop().create_future() disconnect_event = asyncio.Event() - await _cancel_llm_call_on_client_disconnect( - request, llm_call, disconnect_event - ) + await _cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event) assert llm_call.cancelled() assert disconnect_event.is_set() async def test_monitor_is_noop_while_client_stays_connected(self): - request = self._request( - [{"type": "http.request", "body": b"", "more_body": False}] - ) + request = self._request([{"type": "http.request", "body": b"", "more_body": False}]) llm_call = asyncio.get_running_loop().create_future() disconnect_event = asyncio.Event() - monitor = asyncio.create_task( - _cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event) - ) + monitor = asyncio.create_task(_cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event)) await asyncio.sleep(0.01) assert not monitor.done() @@ -3930,9 +3898,7 @@ class TestCancelOnDisconnect: llm_call = asyncio.get_running_loop().create_future() disconnect_event = asyncio.Event() - await _cancel_llm_call_on_client_disconnect( - request, llm_call, disconnect_event - ) + await _cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event) assert not llm_call.cancelled() assert not disconnect_event.is_set() @@ -3947,9 +3913,7 @@ class TestCancelOnDisconnect: with pytest.raises(asyncio.CancelledError): await _await_llm_call_cancelling_on_disconnect(request, llm_call) - async def _drive_base_process_llm_request( - self, monkeypatch, general_settings: dict, llm_call, request: Request - ): + async def _drive_base_process_llm_request(self, monkeypatch, general_settings: dict, llm_call, request: Request): from litellm.proxy._types import UserAPIKeyAuth logging_obj = MagicMock() @@ -3958,9 +3922,7 @@ class TestCancelOnDisconnect: logging_obj._on_deferred_stream_complete = None logging_obj.cost_breakdown = None - processor = ProxyBaseLLMRequestProcessing( - data={"model": "fake-model", "litellm_logging_obj": logging_obj} - ) + processor = ProxyBaseLLMRequestProcessing(data={"model": "fake-model", "litellm_logging_obj": logging_obj}) proxy_logging_obj = MagicMock(spec=ProxyLogging) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) @@ -3968,9 +3930,7 @@ class TestCancelOnDisconnect: proxy_logging_obj.post_call_success_hook = AsyncMock( side_effect=lambda data, user_api_key_dict, response: response ) - proxy_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value=None - ) + proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value=None) async def fake_route_request(**kwargs): return llm_call() @@ -4049,9 +4009,7 @@ class TestCancelOnDisconnect: with pytest.raises(ProxyException) as exc_info: await processor._handle_llm_api_exception( - e=HTTPException( - status_code=499, detail="Client disconnected the request" - ), + e=HTTPException(status_code=499, detail="Client disconnected the request"), user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), proxy_logging_obj=proxy_logging_obj, ) @@ -4117,7 +4075,9 @@ class TestAllmPassthroughRoutePostCallGuardrails: proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", capture_hook) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): + with patch.object( + ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -4167,7 +4127,9 @@ class TestAllmPassthroughRoutePostCallGuardrails: proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", non_dict_hook) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): + with patch.object( + ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -4205,7 +4167,9 @@ class TestAllmPassthroughRoutePostCallGuardrails: hook_spy = AsyncMock() monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", hook_spy) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): + with patch.object( + ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -4246,7 +4210,9 @@ class TestAllmPassthroughRoutePostCallGuardrails: hook_spy = AsyncMock() monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", hook_spy) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=False): + with patch.object( + ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=False + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -4358,7 +4324,9 @@ class TestEventStreamAllmPassthroughRoute: "content-length": "99", } - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): + with patch.object( + ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=mock_response, @@ -4389,9 +4357,7 @@ class TestAllmPassthroughStreamingProviderGate: de-anonymized. """ - def _build_processing_obj( - self, custom_llm_provider: str, endpoint: str = "" - ) -> ProxyBaseLLMRequestProcessing: + def _build_processing_obj(self, custom_llm_provider: str, endpoint: str = "") -> ProxyBaseLLMRequestProcessing: logging_obj = MagicMock() logging_obj.litellm_call_id = "call-123" logging_obj.cost_breakdown = None @@ -4442,14 +4408,17 @@ class TestAllmPassthroughStreamingProviderGate: processing_obj = self._build_processing_obj("anthropic") chunks = [b"chunk-1", b"chunk-2"] - with patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails", - return_value=False, - ), patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails_for_passthrough", - return_value=True, + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=True, + ), ): result = await self._run(processing_obj, monkeypatch, chunks) @@ -4458,27 +4427,27 @@ class TestAllmPassthroughStreamingProviderGate: assert streamed == chunks @pytest.mark.asyncio - async def test_bedrock_converse_stream_is_buffered_through_handler( - self, monkeypatch - ): - processing_obj = self._build_processing_obj( - "bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream" - ) + async def test_bedrock_converse_stream_is_buffered_through_handler(self, monkeypatch): + processing_obj = self._build_processing_obj("bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream") chunks = [b"raw-1", b"raw-2"] - with patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails", - return_value=False, - ), patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails_for_passthrough", - return_value=True, - ), patch( - "litellm.llms.bedrock.passthrough.guardrail_translation.handler." - "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", - new=AsyncMock(return_value=b"modified-body"), - ) as mock_handler: + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=True, + ), + patch( + "litellm.llms.bedrock.passthrough.guardrail_translation.handler." + "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", + new=AsyncMock(return_value=b"modified-body"), + ) as mock_handler, + ): result = await self._run(processing_obj, monkeypatch, chunks) assert isinstance(result, Response) @@ -4494,19 +4463,23 @@ class TestAllmPassthroughStreamingProviderGate: ) chunks = [b"raw-1", b"raw-2"] - with patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails", - return_value=False, - ), patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails_for_passthrough", - return_value=True, - ), patch( - "litellm.llms.bedrock.passthrough.guardrail_translation.handler." - "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", - new=AsyncMock(return_value=b"modified-body"), - ) as mock_handler: + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=True, + ), + patch( + "litellm.llms.bedrock.passthrough.guardrail_translation.handler." + "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", + new=AsyncMock(return_value=b"modified-body"), + ) as mock_handler, + ): result = await self._run(processing_obj, monkeypatch, chunks) assert isinstance(result, StreamingResponse) @@ -4528,14 +4501,17 @@ class TestAllmPassthroughStreamingProviderGate: ) chunks = [b"raw-1", b"raw-2"] - with patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails", - return_value=False, - ), patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails_for_passthrough", - return_value=False, + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=False, + ), ): result = await self._run(processing_obj, monkeypatch, chunks) @@ -4554,14 +4530,17 @@ class TestAllmPassthroughStreamingProviderGate: processing_obj = self._build_processing_obj("anthropic") chunks = [b"chunk-1", b"chunk-2"] - with patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails", - return_value=False, - ), patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails_for_passthrough", - return_value=False, + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=False, + ), ): result = await self._run(processing_obj, monkeypatch, chunks) @@ -4902,7 +4881,6 @@ class TestResponseCostHeaderForTypedDictResponses: class TestPreCallWithFallbacksOnLocalRateLimit: - @pytest.mark.asyncio async def test_fallback_triggered_on_local_rate_limit(self): from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError @@ -5054,9 +5032,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: mock_router.fallbacks = [{"gpt-4": ["gpt-3.5-turbo"]}] user_api_key_dict = MagicMock() - user_api_key_dict.router_settings = { - "fallbacks": [{"gpt-4": ["claude-3-haiku"]}] - } + user_api_key_dict.router_settings = {"fallbacks": [{"gpt-4": ["claude-3-haiku"]}]} with patch.object( processor, @@ -5087,9 +5063,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing - processor = ProxyBaseLLMRequestProcessing( - data={"model": "gpt-4", "disable_fallbacks": True} - ) + processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-4", "disable_fallbacks": True}) async def mock_pre_call_logic(**kwargs): raise ProxyRateLimitError( @@ -5215,9 +5189,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: # Real per-key per-model TPM limiter + a key carrying the customer's # `model_tpm_limit` metadata (only the primary is capped). - limiter = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + limiter = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key="sk-lit3890", metadata={"model_tpm_limit": {primary_model: 100}}, @@ -5225,10 +5197,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: # Pre-seed the primary's per-model token counter at the cap so the very # next request trips it. The counter key uses the *hashed* api_key. - counter_key = ( - f"{user_api_key_dict.api_key}::{primary_model}" - f"::{precise_minute}::request_count" - ) + counter_key = f"{user_api_key_dict.api_key}::{primary_model}::{precise_minute}::request_count" await limiter.internal_usage_cache.async_set_cache( key=counter_key, value={"current_requests": 0, "current_tpm": 100, "current_rpm": 0}, @@ -5259,9 +5228,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: mock_router = MagicMock() mock_router.fallbacks = [{primary_model: [fallback_model]}] - with patch( - "litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock - ): + with patch("litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock): with patch.object( processor, "common_processing_pre_call_logic", @@ -5291,9 +5258,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: # Sanity-check the premise: the limiter genuinely raises a # ProxyRateLimitError for the capped primary under the frozen clock. - with patch( - "litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock - ): + with patch("litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock): with pytest.raises(ProxyRateLimitError): await limiter.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -5654,16 +5619,12 @@ class TestStreamingClientDisconnectBilling: prompt_tokens=1000, completion_tokens=10, total_tokens=1010, - prompt_tokens_details=PromptTokensDetailsWrapper( - cached_tokens=500 - ), + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=500), ), ) ) - event = await self._bill_and_collect_success_event( - append_openai_style_cached_usage_chunk - ) + event = await self._bill_and_collect_success_event(append_openai_style_cached_usage_chunk) usage = event["response_obj"].usage assert getattr(usage, "cache_read_input_tokens", None) == 500 @@ -6433,9 +6394,7 @@ class TestInjectCostIntoUsageDict: logging_obj.model_call_details["custom_llm_provider"] = "anthropic" assert logging_obj.cost_breakdown is None - model_response = ModelResponse( - usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224) - ) + model_response = ModelResponse(usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224)) cost = ProxyBaseLLMRequestProcessing._logging_obj_cost_or_none(model_response, logging_obj) assert cost is not None and cost > 0 @@ -6464,9 +6423,7 @@ class TestInjectCostIntoUsageDict: ) existing = logging_obj.cost_breakdown - model_response = ModelResponse( - usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224) - ) + model_response = ModelResponse(usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224)) ProxyBaseLLMRequestProcessing._logging_obj_cost_or_none(model_response, logging_obj) assert logging_obj.cost_breakdown is existing @@ -6761,9 +6718,7 @@ def test_ttft_keepalive_interval_only_arms_for_a_streaming_request(request_data, @pytest.mark.asyncio @pytest.mark.parametrize("stream_requested, expect_ping", [(True, True), (False, False)]) -async def test_base_process_llm_request_pings_while_the_upstream_call_is_still_running( - stream_requested, expect_ping -): +async def test_base_process_llm_request_pings_while_the_upstream_call_is_still_running(stream_requested, expect_ping): """The wiring, not the helper: every route funnels through this method, and the whole time-to-first-token is spent inside the call it wraps.""" @@ -6909,9 +6864,7 @@ async def test_a_late_failure_is_reported_to_the_failure_hook(): async def record(exc): audited.append(exc) - response = await open_sse_before_first_byte( - slow_failure(), ping_interval_seconds=0.05, on_late_failure=record - ) + response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05, on_late_failure=record) collected = await _drain(response) assert [type(exc).__name__ for exc in audited] == ["HTTPException"] @@ -6928,9 +6881,7 @@ async def test_a_failing_audit_hook_never_costs_the_client_its_error_frame(): async def broken_hook(exc): raise RuntimeError("the audit backend is down") - response = await open_sse_before_first_byte( - slow_failure(), ping_interval_seconds=0.05, on_late_failure=broken_hook - ) + response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05, on_late_failure=broken_hook) collected = await _drain(response) error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip()) @@ -6984,9 +6935,7 @@ async def test_base_process_llm_request_audits_a_failure_that_lands_after_its_ke [(0, False), (None, True)], ids=["operator-hard-disabled-this-deployment", "deployment-says-nothing"], ) -async def test_base_process_llm_request_honours_a_deployment_hard_disable( - deployment_keepalive, expect_ping -): +async def test_base_process_llm_request_honours_a_deployment_hard_disable(deployment_keepalive, expect_ping): """`keepalive_seconds: 0` is documented as a disable a request cannot lift. The funnel has to hand its router to the gate for that to hold before the upstream has answered, since no deployment has served the request yet.""" @@ -7032,9 +6981,7 @@ async def test_a_hook_returning_a_replacement_decides_what_the_client_sees(): async def sanitize(exc): return HTTPException(status_code=502, detail="upstream unavailable") - response = await open_sse_before_first_byte( - slow_failure(), ping_interval_seconds=0.05, on_late_failure=sanitize - ) + response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05, on_late_failure=sanitize) collected = await _drain(response) error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip()) @@ -7073,9 +7020,7 @@ async def test_a_hook_that_returns_nothing_leaves_the_real_error_intact(): async def audit_only(exc): return None - response = await open_sse_before_first_byte( - slow_failure(), ping_interval_seconds=0.05, on_late_failure=audit_only - ) + response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05, on_late_failure=audit_only) collected = await _drain(response) error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip()) @@ -7092,9 +7037,7 @@ async def test_a_broken_hook_does_not_replace_the_real_error_with_its_own_bug(): async def broken_hook(exc): raise RuntimeError("the audit backend is down") - response = await open_sse_before_first_byte( - slow_failure(), ping_interval_seconds=0.05, on_late_failure=broken_hook - ) + response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05, on_late_failure=broken_hook) collected = await _drain(response) error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip()) From c3bcb6f64f787e256d106d98d3ef17dea525c78b Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Tue, 25 Aug 2026 10:50:35 -0700 Subject: [PATCH 85/93] test(mcp): drain the logging worker after each test so queued callbacks cannot leak into the next test (#38228) LoggingWorker now carries still-queued coroutines onto the next event loop (12a34a10d8). Under xdist, a success-logging coroutine queued by test_acompletion_mcp_respects_manual_approval ran nine seconds later inside test_mcp_tool_call_hook on the same worker, resolved litellm.callbacks at run time and overwrote that test's captured payload with a gpt-4o-mini completion (assert 1.35e-05 == 1.42). Run clear_queue() in the suite's autouse teardown so every coroutine a test enqueues finishes before the next test registers its callbacks, and add a subprocess regression test that runs the real conftest against a stopped worker with work still queued. --- tests/mcp_tests/conftest.py | 3 ++ tests/mcp_tests/test_mcp_logging.py | 45 +++++++++++++++++++++++++++++ 2 files changed, 48 insertions(+) diff --git a/tests/mcp_tests/conftest.py b/tests/mcp_tests/conftest.py index d1dc3ec7216..5823893afc0 100644 --- a/tests/mcp_tests/conftest.py +++ b/tests/mcp_tests/conftest.py @@ -7,6 +7,7 @@ import pytest import litellm import asyncio +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER @pytest.fixture(scope="session") @@ -38,6 +39,8 @@ def setup_and_teardown(): yield # Teardown code (executes after the yield point) + # LoggingWorker carries still-queued coroutines onto the next test's loop, where they'd log into that test's callbacks + asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue()) loop.close() # Close the loop created earlier asyncio.set_event_loop(None) # Remove the reference to the loop diff --git a/tests/mcp_tests/test_mcp_logging.py b/tests/mcp_tests/test_mcp_logging.py index 1903f29001f..fc9f675f837 100644 --- a/tests/mcp_tests/test_mcp_logging.py +++ b/tests/mcp_tests/test_mcp_logging.py @@ -1,6 +1,9 @@ import os import pytest import asyncio +import subprocess +import sys +from pathlib import Path from typing import Optional from unittest.mock import AsyncMock, patch @@ -458,3 +461,45 @@ async def test_mcp_tool_call_hook(): logged_standard_logging_payload is not None ), "Standard logging payload should not be None" assert logged_standard_logging_payload["response_cost"] == 1.42 + + +_QUEUED_LOGGING_OUTLIVES_TEST = ''' +import time + +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + +ran_at = [] + + +async def _record_run(): + ran_at.append(time.monotonic()) + + +async def test_1_leaves_logging_queued_behind_a_stopped_worker(): + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(_record_run()) + await GLOBAL_LOGGING_WORKER.stop() + assert ran_at == [] + + +async def test_2_starts_after_the_previous_tests_logging_ran(): + started_at = time.monotonic() + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(_record_run()) + await GLOBAL_LOGGING_WORKER.flush() + assert [t < started_at for t in ran_at] == [True, False] +''' + + +def test_logging_queued_by_one_test_is_drained_before_the_next(tmp_path: Path): + """Regression: a logging coroutine queued by one test must not run inside a later test (it would log into that + test's callbacks, which is how test_mcp_tool_call_hook captured a gpt-4o-mini payload under xdist).""" + (tmp_path / "conftest.py").write_text((Path(__file__).parent / "conftest.py").read_text()) + (tmp_path / "pyproject.toml").write_text('[tool.pytest.ini_options]\nasyncio_mode = "auto"\n') + (tmp_path / "test_queued_logging.py").write_text(_QUEUED_LOGGING_OUTLIVES_TEST) + result = subprocess.run( + [sys.executable, "-m", "pytest", "-q", "-p", "no:cacheprovider", "test_queued_logging.py"], + cwd=tmp_path, + capture_output=True, + text=True, + timeout=120, + ) + assert result.returncode == 0, result.stdout + result.stderr From e8bdbcd1cf176914a2b110f95e8fc9d87c574436 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:56:55 -0700 Subject: [PATCH 86/93] fix(bedrock_mantle): parse converse passthrough bodies with the converse shape config for logging --- .../passthrough/transformation.py | 29 +++++++++++- ...drock_mantle_passthrough_transformation.py | 45 +++++++++++++++++++ 2 files changed, 73 insertions(+), 1 deletion(-) diff --git a/litellm/llms/bedrock_mantle/passthrough/transformation.py b/litellm/llms/bedrock_mantle/passthrough/transformation.py index 1393ac7c6e7..e6b831efa57 100644 --- a/litellm/llms/bedrock_mantle/passthrough/transformation.py +++ b/litellm/llms/bedrock_mantle/passthrough/transformation.py @@ -1,12 +1,19 @@ from collections.abc import Mapping -from typing import Final, Literal +from typing import TYPE_CHECKING, Final, Literal, Optional +from httpx import Response + +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.llms.bedrock.passthrough.transformation import BedrockPassthroughConfig from litellm.llms.bedrock_mantle.common_utils import ( MANTLE_HOST_RE, resolve_mantle_bearer_token, resolve_mantle_region, ) +from litellm.types.utils import LlmProviders + +if TYPE_CHECKING: + from litellm.types.utils import CostResponseTypes class BedrockMantlePassthroughConfig(BedrockPassthroughConfig): @@ -42,3 +49,23 @@ class BedrockMantlePassthroughConfig(BedrockPassthroughConfig): def get_bedrock_bearer_token(self, litellm_params: Mapping[str, object]) -> str | None: api_key: Final = litellm_params.get("api_key") return resolve_mantle_bearer_token(api_key if isinstance(api_key, str) else None) + + def logging_non_streaming_response( + self, + model: str, + custom_llm_provider: str, + httpx_response: Response, + request_data: dict, # mutable-ok: mirrors the inherited BedrockPassthroughConfig signature + logging_obj: Logging, + endpoint: str, + ) -> Optional["CostResponseTypes"]: + is_converse: Final = "invoke" not in endpoint and "converse" in endpoint + shape_provider: Final = LlmProviders.BEDROCK.value if is_converse else custom_llm_provider + return super().logging_non_streaming_response( + model=model, + custom_llm_provider=shape_provider, + httpx_response=httpx_response, + request_data=request_data, + logging_obj=logging_obj, + endpoint=endpoint, + ) diff --git a/tests/test_litellm/llms/bedrock_mantle/passthrough/test_bedrock_mantle_passthrough_transformation.py b/tests/test_litellm/llms/bedrock_mantle/passthrough/test_bedrock_mantle_passthrough_transformation.py index b7f9e492e14..8c6eda605ca 100644 --- a/tests/test_litellm/llms/bedrock_mantle/passthrough/test_bedrock_mantle_passthrough_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/passthrough/test_bedrock_mantle_passthrough_transformation.py @@ -14,6 +14,7 @@ from litellm.utils import ProviderConfigManager MANTLE_API_BASE = "https://bedrock-mantle.us-east-2.api.aws" INVOKE_ENDPOINT = "model/us.openai.gpt-5.6-sol/invoke" +CONVERSE_ENDPOINT = "model/us.openai.gpt-5.6-sol/converse" REQUEST_BODY = {"messages": [{"role": "user", "content": "say pong"}], "max_completion_tokens": 64} @@ -150,3 +151,47 @@ def test_invoke_passthrough_route_reaches_bedrock_runtime_for_a_mantle_deploymen assert str(sent["url"]) == f"https://bedrock-runtime.us-east-2.amazonaws.com/{INVOKE_ENDPOINT}" assert sent["headers"]["Authorization"] == f"Bearer {expected_bearer}" assert json.loads(sent["content"]) == REQUEST_BODY + + +def _logged_model_response(endpoint, body): + request = httpx.Request("POST", f"https://bedrock-runtime.us-east-1.amazonaws.com/{endpoint}") + return BedrockMantlePassthroughConfig().logging_non_streaming_response( + model="us.openai.gpt-5.6-sol", + custom_llm_provider="bedrock_mantle", + httpx_response=httpx.Response(200, json=body, request=request), + request_data={"messages": [{"role": "user", "content": [{"text": "say pong"}]}]}, + logging_obj=MagicMock(), + endpoint=endpoint, + ) + + +def test_converse_logging_parses_the_converse_response_shape(): + result = _logged_model_response( + CONVERSE_ENDPOINT, + { + "metrics": {"latencyMs": 800.0}, + "output": {"message": {"content": [{"text": "pong"}], "role": "assistant"}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 8, "outputTokens": 5, "totalTokens": 13}, + }, + ) + assert result.choices[0].message.content == "pong" + assert result.usage.prompt_tokens == 8 + assert result.usage.completion_tokens == 5 + + +def test_invoke_logging_parses_the_openai_chat_response_shape(): + result = _logged_model_response( + INVOKE_ENDPOINT, + { + "choices": [{"finish_reason": "stop", "index": 0, "message": {"content": "pong", "role": "assistant"}}], + "created": 1787677792, + "id": "chatcmpl-regression", + "model": "us.openai.gpt-5.6-sol", + "object": "chat.completion", + "usage": {"completion_tokens": 5, "prompt_tokens": 8, "total_tokens": 13}, + }, + ) + assert result.choices[0].message.content == "pong" + assert result.usage.prompt_tokens == 8 + assert result.usage.completion_tokens == 5 From 5470c1bccbaa31aa1fccc5a5801c402588b43e80 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Tue, 25 Aug 2026 10:59:26 -0700 Subject: [PATCH 87/93] fix(ui): forward OAuth issuer/authorization/token/registration URLs from the MCP server edit form (#38154) The edit form's Authorize & Fetch Token button built its temporary OAuth session payload without issuer, authorization_url, token_url, or registration_url, unlike the create form's equivalent payload builder. The backend's temporary-session endpoint builds its ephemeral server purely from that payload, so any admin-configured OAuth endpoints on an existing server were silently dropped, endpoint discovery fell back to (and failed against) the plain server url, and Authorize & Fetch Token 400'd with "authorization url is not configured" even though the saved server had those fields filled in. Add the four missing fields to the edit form's temporary payload builder, mirroring the create form. --- .../_components/mcp_server_edit.test.tsx | 38 +++++++++++++++++++ .../_components/mcp_server_edit.tsx | 4 ++ 2 files changed, 42 insertions(+) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx index 438caa2f5e6..5aec78ba926 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx @@ -381,6 +381,44 @@ describe("MCPServerEdit (true passthrough warning)", () => { }); }); +describe("MCPServerEdit (OAuth authorize temp payload)", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("forwards issuer/authorization_url/token_url/registration_url to the temp OAuth session payload", async () => { + // Without these fields the ephemeral server the temp OAuth session endpoint builds has no + // admin-configured OAuth endpoints on it, discovery falls back to (and fails against) the + // plain server url, and Authorize & Fetch Token 400s with "authorization url is not + // configured" even though the saved server (and the visible form) has all four fields filled in. + render( + , + ); + + await waitFor(() => { + expect(mockOauth.getTemporaryPayload).toBeTruthy(); + }); + const payload = mockOauth.getTemporaryPayload!(); + expect(payload).toBeTruthy(); + expect(payload?.issuer).toBe("https://github.com/login/oauth"); + expect(payload?.authorization_url).toBe("https://github.com/login/oauth/authorize"); + expect(payload?.token_url).toBe("https://github.com/login/oauth/access_token"); + expect(payload?.registration_url).toBe("https://github.com/login/oauth/register"); + }); +}); + describe("MCPServerEdit (auth type switch)", () => { beforeEach(() => { vi.clearAllMocks(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx index 5e79b20825c..8793c45371a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx @@ -282,6 +282,10 @@ const MCPServerEdit: React.FC = ({ credentials: isClientForwardedTokenMode(values.auth_type) ? preservedAdminCredentials(values.credentials) : values.credentials, + issuer: values.issuer, + authorization_url: values.authorization_url, + token_url: values.token_url, + registration_url: values.registration_url, mcp_access_groups: values.mcp_access_groups || mcpServer.mcp_access_groups, static_headers: staticHeaders, command: values.command, From 559588a4732e94971781983eaf433cb11da9fbb9 Mon Sep 17 00:00:00 2001 From: milan Date: Tue, 25 Aug 2026 18:23:09 +0000 Subject: [PATCH 88/93] fix(azure_ai): remove new LIT002 violations to satisfy type-discipline budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/azure_ai/azure_model_router/transformation.py | 2 +- litellm/llms/azure_ai/common_utils.py | 6 +++++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/litellm/llms/azure_ai/azure_model_router/transformation.py b/litellm/llms/azure_ai/azure_model_router/transformation.py index d33564c0f9a..61cbc213b11 100644 --- a/litellm/llms/azure_ai/azure_model_router/transformation.py +++ b/litellm/llms/azure_ai/azure_model_router/transformation.py @@ -99,7 +99,7 @@ class AzureModelRouterConfig(AzureAIStudioConfig): if selected_model: # Rebuilt rather than mutated in place: ModelResponseBase declares _hidden_params as a # class-level dict, so an in-place write can bleed into unrelated responses. - transformed_response._hidden_params = { # pyright: ignore[reportPrivateUsage] # ModelResponse exposes no public hidden-params setter + transformed_response._hidden_params = { # pyright: ignore[reportPrivateUsage] # ModelResponse exposes no public hidden-params setter # mutable-ok: ModelResponse requires _hidden_params to be a plain dict **get_hidden_params_dict(transformed_response), AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY: selected_model, } diff --git a/litellm/llms/azure_ai/common_utils.py b/litellm/llms/azure_ai/common_utils.py index e8431f29f07..55a51176666 100644 --- a/litellm/llms/azure_ai/common_utils.py +++ b/litellm/llms/azure_ai/common_utils.py @@ -112,7 +112,11 @@ class AzureFoundryModelInfo(BaseLLMModelInfo): """ if AzureFoundryModelInfo.get_model_router_selected_model(hidden_params) is not None: return True - deployment_model: Final = (hidden_params or {}).get("litellm_model_name") or (hidden_params or {}).get("model") + deployment_model: Final = ( + hidden_params.get("litellm_model_name") or hidden_params.get("model") + if hidden_params is not None + else None + ) return any( isinstance(candidate, str) and AzureFoundryModelInfo.get_azure_ai_route(candidate) == "model_router" for candidate in (deployment_model, model) From 77f22be2075858dae62de1f44ecec85e6131ab86 Mon Sep 17 00:00:00 2001 From: milan Date: Tue, 25 Aug 2026 18:27:19 +0000 Subject: [PATCH 89/93] style: apply ruff format to common_utils Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/azure_ai/common_utils.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/litellm/llms/azure_ai/common_utils.py b/litellm/llms/azure_ai/common_utils.py index 55a51176666..26a90157455 100644 --- a/litellm/llms/azure_ai/common_utils.py +++ b/litellm/llms/azure_ai/common_utils.py @@ -113,9 +113,7 @@ class AzureFoundryModelInfo(BaseLLMModelInfo): if AzureFoundryModelInfo.get_model_router_selected_model(hidden_params) is not None: return True deployment_model: Final = ( - hidden_params.get("litellm_model_name") or hidden_params.get("model") - if hidden_params is not None - else None + hidden_params.get("litellm_model_name") or hidden_params.get("model") if hidden_params is not None else None ) return any( isinstance(candidate, str) and AzureFoundryModelInfo.get_azure_ai_route(candidate) == "model_router" From 6cd1fcdcf05734cff34d0dfed7e098da73a0a913 Mon Sep 17 00:00:00 2001 From: milan Date: Tue, 25 Aug 2026 18:37:26 +0000 Subject: [PATCH 90/93] test: drop unneeded proxy_server patches and ratchet lint budgets Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ruff-strict-budget.json | 2 +- .../proxy/spend_tracking/test_spend_tracking_utils.py | 11 ++--------- type-discipline-budget.json | 2 +- 3 files changed, 4 insertions(+), 11 deletions(-) diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 03318718fb5..9815ccfa23e 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -168,7 +168,7 @@ "limit": 3 }, "RET504": { - "limit": 176 + "limit": 175 }, "RUF012": { "limit": 240 diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index dab9415de77..843e1d1296f 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -3,13 +3,10 @@ import datetime import json from datetime import timezone from typing import Any, Final, cast - -from typing_extensions import ReadOnly, TypedDict +from unittest.mock import AsyncMock, MagicMock, patch import pytest - - -from unittest.mock import AsyncMock, MagicMock, patch +from typing_extensions import ReadOnly, TypedDict import litellm from litellm.constants import ( @@ -3526,8 +3523,6 @@ def _model_router_spend_log_kwargs(slp_model: str | None) -> _ModelRouterSpendLo } -@patch("litellm.proxy.proxy_server.master_key", None) -@patch("litellm.proxy.proxy_server.general_settings", {}) def test_get_logging_payload_uses_standard_logging_payload_model(): payload = get_logging_payload( kwargs=_model_router_spend_log_kwargs(slp_model="azure_ai/gpt-5-mini"), @@ -3538,8 +3533,6 @@ def test_get_logging_payload_uses_standard_logging_payload_model(): assert payload["model"] == "azure_ai/gpt-5-mini" -@patch("litellm.proxy.proxy_server.master_key", None) -@patch("litellm.proxy.proxy_server.general_settings", {}) def test_get_logging_payload_falls_back_to_kwargs_model_when_slp_model_missing(): payload = get_logging_payload( kwargs=_model_router_spend_log_kwargs(slp_model=None), diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 05098546325..542762f1cea 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -30,7 +30,7 @@ "limit": 16673 }, "LIT011": { - "limit": 5588 + "limit": 5587 }, "LIT012": { "limit": 4510 From e25ce85dbc0dfc9059f11e02387ba0cf7a916a3b Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 25 Aug 2026 12:44:16 -0700 Subject: [PATCH 91/93] fix(proxy): stop expected 4xx responses from saturating worker CPU on failure logging (#38102) * fix(proxy): stop expected 4xx responses from saturating worker CPU on failure logging Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep threaded sync failure handler so CustomLogger sync callbacks still run Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: remove redundant comments per review Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/__init__.py | 1 + litellm/litellm_core_utils/core_helpers.py | 20 ++ litellm/litellm_core_utils/litellm_logging.py | 14 +- litellm/proxy/auth/auth_exception_handler.py | 8 +- litellm/proxy/common_request_processing.py | 9 +- litellm/proxy/utils.py | 42 ++-- .../litellm_core_utils/test_core_helpers.py | 23 ++ .../test_litellm_logging.py | 59 +++++ .../proxy/auth/test_auth_exception_handler.py | 55 +++++ .../proxy/test_common_request_processing.py | 28 +++ tests/test_litellm/proxy/test_proxy_utils.py | 218 +++++++++++------- 11 files changed, 373 insertions(+), 104 deletions(-) diff --git a/litellm/__init__.py b/litellm/__init__.py index ee2c551481c..f42b8280761 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -199,6 +199,7 @@ standard_logging_payload_excluded_fields: Optional[List[str]] = ( None # Fields to exclude from StandardLoggingPayload before callbacks receive it ) log_raw_request_response: bool = False +log_client_error_tracebacks: bool = False request_correlation_in_logs: bool = False redact_messages_in_exceptions: Optional[bool] = False redact_user_api_key_info: Optional[bool] = False diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index de1092bc02f..b71ef4c8e06 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -58,6 +58,26 @@ def safe_divide( return numerator / denominator +def is_expected_client_error(exception: BaseException | None) -> bool: + """ + True when the exception maps to an HTTP 4xx status. + + ProxyException stores the status on .code (as a str), HTTPException and + litellm exceptions on .status_code. + """ + if exception is None: + return False + code: Final[object] = getattr(exception, "code", None) + status_code: Final[object] = code if code is not None else getattr(exception, "status_code", None) + if status_code is None or isinstance(status_code, bool): + return False + try: + status: Final = int(str(status_code)) + except ValueError: + return False + return 400 <= status < 500 + + def coerce_token_limit(value: object) -> int | None: """ Coerce a max_input_tokens / max_output_tokens value to an int, treating a diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 960ed266e75..626af7530a4 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -62,7 +62,7 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.deepeval.deepeval import DeepEvalLogger from litellm.integrations.mlflow import MlflowLogger from litellm.integrations.sqs import SQSLogger -from litellm.litellm_core_utils.core_helpers import reconstruct_model_name +from litellm.litellm_core_utils.core_helpers import is_expected_client_error, reconstruct_model_name from litellm.litellm_core_utils.get_litellm_params import get_litellm_params from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import ( cost_breakdown_with_guardrail, @@ -3124,6 +3124,13 @@ class Logging(LiteLLMLoggingBaseClass): if not hasattr(self, "model_call_details"): self.model_call_details = {} + if ( + self.model_call_details.get("log_event_type") == "failed_api_call" + and self.model_call_details.get("exception") is exception + and self.model_call_details.get("standard_logging_object") is not None + ): + return start_time, self.model_call_details["end_time"] + self.model_call_details["log_event_type"] = "failed_api_call" self.model_call_details["exception"] = exception self.model_call_details["traceback_exception"] = ( @@ -5455,9 +5462,10 @@ class StandardLoggingPayloadSetup: error_class: Final[str] = str(original_exception.__class__.__name__) if original_exception else "" _llm_provider_in_exception: Final = getattr(original_exception, "llm_provider", "") - # Get traceback information (first 100 lines) traceback_info = traceback_str or "" - if original_exception: + if original_exception and ( + litellm.log_client_error_tracebacks or not is_expected_client_error(original_exception) + ): tb: Final[TracebackType | None] = getattr(original_exception, "__traceback__", None) if tb: tb_lines: Final = traceback.format_tb(tb) diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index 233679126f8..a42187b3a44 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -11,6 +11,7 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.constants import EMPTY_MAPPING from litellm.integrations.otel.runtime import seed_request_identity +from litellm.litellm_core_utils.core_helpers import is_expected_client_error from litellm.proxy._types import ( LitellmUserRoles, ProxyErrorTypes, @@ -109,7 +110,12 @@ class UserAPIKeyAuthExceptionHandler: request=request, use_x_forwarded_for=general_settings.get("use_x_forwarded_for") is True, ) - verbose_proxy_logger.exception( + log_fn: Final = ( + verbose_proxy_logger.error + if is_expected_client_error(e) and not litellm.log_client_error_tracebacks + else verbose_proxy_logger.exception + ) + log_fn( "litellm.proxy.proxy_server.user_api_key_auth(): Exception occured - %s\nRequester IP Address:%s", e, requester_ip, diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 413fead9bd4..6fd743deb44 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -33,7 +33,7 @@ from litellm.constants import ( UNSAFE_PROXY_RESPONSE_HEADERS, ) from litellm.integrations.custom_guardrail import CustomGuardrail -from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket +from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket, is_expected_client_error from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer from litellm.litellm_core_utils.get_supported_openai_params import ( get_supported_openai_params, @@ -1380,7 +1380,12 @@ def _log_llm_api_exception(e: Exception) -> None: "litellm.proxy.proxy_server._handle_llm_api_exception(): client disconnected, upstream LLM request cancelled" ) return - verbose_proxy_logger.exception("litellm.proxy.proxy_server._handle_llm_api_exception(): Exception occured - %s", e) + log_fn: Final = ( + verbose_proxy_logger.error + if is_expected_client_error(e) and not litellm.log_client_error_tracebacks + else verbose_proxy_logger.exception + ) + log_fn("litellm.proxy.proxy_server._handle_llm_api_exception(): Exception occured - %s", e) async def _cancel_llm_call_on_client_disconnect( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 86d954c0913..2bbdad9c487 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -91,7 +91,7 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.prometheus import PrometheusLogger from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert -from litellm.litellm_core_utils.core_helpers import coerce_token_limit +from litellm.litellm_core_utils.core_helpers import coerce_token_limit, is_expected_client_error from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.safe_json_loads import safe_json_loads @@ -2575,20 +2575,36 @@ class ProxyLogging: api_key="", ) - # log the custom exception - await litellm_logging_obj.async_failure_handler( - exception=original_exception, - traceback_exception=traceback.format_exc(), + await self._dispatch_proxy_only_failure_handlers( + litellm_logging_obj=litellm_logging_obj, + original_exception=original_exception, ) - threading.Thread( - target=litellm_logging_obj.failure_handler, - args=( - original_exception, - traceback.format_exc(), - ), - daemon=True, - ).start() + @staticmethod + async def _dispatch_proxy_only_failure_handlers( + litellm_logging_obj: Logging, + original_exception: Exception | None, + ) -> None: + """Runs the async failure handler plus the threaded sync handler. Expected + client (4xx) errors skip traceback formatting unless + litellm.log_client_error_tracebacks is set.""" + include_traceback: Final = litellm.log_client_error_tracebacks or not is_expected_client_error( + original_exception + ) + traceback_str: Final = traceback.format_exc() if include_traceback else "" + await litellm_logging_obj.async_failure_handler( + exception=original_exception, + traceback_exception=traceback_str, + ) + + threading.Thread( + target=litellm_logging_obj.failure_handler, + args=( + original_exception, + traceback_str, + ), + daemon=True, + ).start() async def post_call_success_hook( self, diff --git a/tests/test_litellm/litellm_core_utils/test_core_helpers.py b/tests/test_litellm/litellm_core_utils/test_core_helpers.py index aa7f2990c13..33a33711698 100644 --- a/tests/test_litellm/litellm_core_utils/test_core_helpers.py +++ b/tests/test_litellm/litellm_core_utils/test_core_helpers.py @@ -255,3 +255,26 @@ class TestRedactNestedMatchAndRegexKeys: def test_passes_through_none_and_str(self): assert redact_nested_match_and_regex_keys(None) is None assert redact_nested_match_and_regex_keys("plain") == "plain" + + +class TestIsExpectedClientError: + def test_status_ranges(self): + from litellm.litellm_core_utils.core_helpers import is_expected_client_error + + class WithStatusCode(Exception): + def __init__(self, status_code): + self.status_code = status_code + + class WithCode(Exception): + def __init__(self, code): + self.code = code + + assert is_expected_client_error(WithStatusCode(400)) is True + assert is_expected_client_error(WithStatusCode(429)) is True + assert is_expected_client_error(WithStatusCode(499)) is True + assert is_expected_client_error(WithStatusCode(500)) is False + assert is_expected_client_error(WithStatusCode(399)) is False + assert is_expected_client_error(WithCode("403")) is True + assert is_expected_client_error(WithCode("invalid_request_error")) is False + assert is_expected_client_error(Exception("no status")) is False + assert is_expected_client_error(None) is False diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index cc738d68016..f93daa61570 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -5678,3 +5678,62 @@ def test_get_custom_logger_compatible_class_finds_v2_newrelic(monkeypatch): logging_module._in_memory_loggers.clear() monkeypatch.delenv("LITELLM_OTEL_V2", raising=False) is_otel_v2_enabled.cache_clear() + + +class _ClientError(Exception): + def __init__(self, status_code, message): + self.status_code = status_code + self.message = message + super().__init__(message) + + +def _raise_and_catch(exc): + try: + raise exc + except Exception as caught: + return caught + + +def test_get_error_information_skips_traceback_for_expected_4xx(monkeypatch): + """Regression for LIT-6043: expected client (4xx) errors must not pay for + traceback.format_tb on every rejected request unless + litellm.log_client_error_tracebacks is enabled.""" + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + client_exc = _raise_and_catch(_ClientError(status_code=403, message="team does not allow model")) + assert client_exc.__traceback__ is not None + result = StandardLoggingPayloadSetup.get_error_information(client_exc) + assert result["traceback"] == "" + + server_exc = _raise_and_catch(_ClientError(status_code=500, message="boom")) + result = StandardLoggingPayloadSetup.get_error_information(server_exc) + assert "test_litellm_logging" in result["traceback"] + + monkeypatch.setattr(litellm, "log_client_error_tracebacks", True) + result = StandardLoggingPayloadSetup.get_error_information(client_exc) + assert "test_litellm_logging" in result["traceback"] + + +def test_failure_handler_helper_fn_builds_payload_once_per_exception(): + """Regression for LIT-6043: async and sync failure handlers both call + _failure_handler_helper_fn for the same failed request; the standardized + payload must be built once, not once per handler.""" + obj = LitellmLogging( + model="gpt-4o", + messages=[{"role": "user", "content": "Hey"}], + stream=False, + call_type="acompletion", + start_time=time.time(), + litellm_call_id="lit-6043-1", + function_id="f", + ) + exc = _raise_and_catch(_ClientError(status_code=400, message="invalid model")) + obj._failure_handler_helper_fn(exception=exc, traceback_exception="") + first_payload = obj.model_call_details["standard_logging_object"] + assert first_payload is not None + obj._failure_handler_helper_fn(exception=exc, traceback_exception="") + assert obj.model_call_details["standard_logging_object"] is first_payload + + other_exc = _raise_and_catch(_ClientError(status_code=429, message="rate limited")) + obj._failure_handler_helper_fn(exception=other_exc, traceback_exception="") + assert obj.model_call_details["standard_logging_object"] is not first_payload diff --git a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py index b0094b81112..90b3b29d919 100644 --- a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py +++ b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py @@ -701,3 +701,58 @@ async def test_auth_failure_ip_stamp_does_not_mutate_callers_request_data(): ) assert request_data == {"model": "gpt-4o"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "auth_error,expect_traceback", + [ + pytest.param( + ProxyException( + message="Authentication Error", type=ProxyErrorTypes.auth_error, param=None, code=401 + ), + False, + id="expected_401_no_traceback", + ), + pytest.param(ValueError("unexpected internal error"), True, id="unexpected_error_keeps_traceback"), + ], +) +async def test_handle_authentication_error_traceback_only_for_unexpected_errors(auth_error, expect_traceback, caplog): + """Regression for LIT-6043: expected 4xx auth rejections must not format a + traceback via logger.exception; unexpected errors must keep it.""" + handler = UserAPIKeyAuthExceptionHandler() + + with ( + patch( # test-quality-ok: handler reads proxy_server globals at call time + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ), + patch( # test-quality-ok: handler reads proxy_server globals at call time + "litellm.proxy.auth.auth_exception_handler.seed_request_identity", + ), + patch( # test-quality-ok: handler reads proxy_server globals at call time + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + verbose_proxy_logger.propagate = True + try: + try: + raise auth_error + except (ProxyException, ValueError) as caught: + with caplog.at_level("ERROR", logger="LiteLLM Proxy"), pytest.raises(ProxyException): + await handler._handle_authentication_error( + caught, + MagicMock(), + {}, + "/v1/chat/completions", + None, + "sk-bad-key", + ) + finally: + verbose_proxy_logger.propagate = False + + records = [r for r in caplog.records if "user_api_key_auth(): Exception occured" in r.getMessage()] + assert len(records) == 1 + assert (records[0].exc_info is not None) is expect_traceback diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index c9b3e5faf8b..b595c44d2ce 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -7343,3 +7343,31 @@ class TestRouterModelNameOnNonStreamingResponse: ) assert response[ROUTER_MODEL_NAME_RESPONSE_FIELD] == "deep-model" + + +@pytest.mark.parametrize( + "exc,expect_traceback", + [ + pytest.param(HTTPException(status_code=400, detail="Invalid model name passed in"), False, id="expected_400"), + pytest.param(ValueError("unexpected internal error"), True, id="unexpected_error"), + ], +) +def test_log_llm_api_exception_traceback_only_for_unexpected_errors(exc, expect_traceback, caplog): + """Regression for LIT-6043: expected 4xx errors log without formatting a + traceback; unexpected errors keep logger.exception behavior.""" + from litellm._logging import verbose_proxy_logger + from litellm.proxy.common_request_processing import _log_llm_api_exception + + verbose_proxy_logger.propagate = True + try: + with caplog.at_level("ERROR", logger="LiteLLM Proxy"): + try: + raise exc + except Exception as raised: + _log_llm_api_exception(raised) + finally: + verbose_proxy_logger.propagate = False + + records = [r for r in caplog.records if "_handle_llm_api_exception(): Exception occured" in r.getMessage()] + assert len(records) == 1 + assert (records[0].exc_info is not None) is expect_traceback diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index fb01216982f..6920cc0dae3 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -11,7 +11,6 @@ from litellm.proxy.utils import ProxyLogging from litellm.types.guardrails import GuardrailEventHooks - from unittest.mock import MagicMock, patch from litellm.proxy.utils import get_custom_url, join_paths @@ -82,9 +81,7 @@ async def test_proxy_only_error_log_marks_no_upstream_llm_call(): captured = {} def fake_pre_call(self, *args, **kwargs): - captured["flag"] = self.model_call_details.get( - LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL - ) + captured["flag"] = self.model_call_details.get(LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL) from litellm.litellm_core_utils.litellm_logging import Logging @@ -102,9 +99,7 @@ async def test_proxy_only_error_log_marks_no_upstream_llm_call(): "model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}], }, - user_api_key_dict=UserAPIKeyAuth( - api_key="sk-bad", request_route="/v1/chat/completions" - ), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-bad", request_route="/v1/chat/completions"), route="/v1/chat/completions", original_exception=Exception("bad key"), ) @@ -148,13 +143,9 @@ async def test_proxy_only_error_log_keeps_litellm_metadata_in_litellm_params(): request_data={ "model": "gpt-4o", "input": "blocked prompt", - "litellm_metadata": { - "standard_logging_guardrail_information": guardrail_info - }, + "litellm_metadata": {"standard_logging_guardrail_information": guardrail_info}, }, - user_api_key_dict=UserAPIKeyAuth( - api_key="sk-1234", request_route="/v1/responses" - ), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", request_route="/v1/responses"), route="/v1/responses", original_exception=HTTPException(status_code=400, detail="blocked"), ) @@ -163,12 +154,7 @@ async def test_proxy_only_error_log_keeps_litellm_metadata_in_litellm_params(): Logging.pre_call = orig_pre_call Logging.async_failure_handler = orig_async_failure - assert ( - captured["litellm_params"]["litellm_metadata"][ - "standard_logging_guardrail_information" - ] - == guardrail_info - ) + assert captured["litellm_params"]["litellm_metadata"]["standard_logging_guardrail_information"] == guardrail_info assert "litellm_metadata" not in captured["optional_params"] @@ -206,9 +192,7 @@ def test_get_model_group_info_order(): def test_join_paths_no_duplication(): """Test that join_paths doesn't duplicate route when base_path already ends with it""" - result = join_paths( - base_path="http://0.0.0.0:4000/my-custom-path/", route="/my-custom-path" - ) + result = join_paths(base_path="http://0.0.0.0:4000/my-custom-path/", route="/my-custom-path") assert result == "http://0.0.0.0:4000/my-custom-path" @@ -814,9 +798,9 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: estimated = request_data["combined_usage_object"] assert isinstance(estimated, Usage) - expected = litellm_module.token_counter(model="gpt-3.5-turbo", messages=messages) + litellm_module.token_counter( - model="gpt-3.5-turbo", text=system_prompt - ) + expected = litellm_module.token_counter( + model="gpt-3.5-turbo", messages=messages + ) + litellm_module.token_counter(model="gpt-3.5-turbo", text=system_prompt) assert estimated.prompt_tokens == expected @pytest.mark.asyncio @@ -834,7 +818,9 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: estimated = request_data["combined_usage_object"] assert isinstance(estimated, Usage) - expected = litellm_module.token_counter(model="gpt-3.5-turbo", messages=messages) + litellm_module.token_counter( + expected = litellm_module.token_counter( + model="gpt-3.5-turbo", messages=messages + ) + litellm_module.token_counter( model="gpt-3.5-turbo", text="part one of the system prompt. part two of the system prompt." ) assert estimated.prompt_tokens == expected @@ -870,9 +856,9 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: estimated = request_data["combined_usage_object"] assert isinstance(estimated, Usage) - expected = litellm_module.token_counter(model="gpt-3.5-turbo", messages=messages) + litellm_module.token_counter( - model="gpt-3.5-turbo", text=system_prompt - ) + expected = litellm_module.token_counter( + model="gpt-3.5-turbo", messages=messages + ) + litellm_module.token_counter(model="gpt-3.5-turbo", text=system_prompt) assert estimated.prompt_tokens == expected @pytest.mark.asyncio @@ -890,9 +876,9 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: estimated = request_data["combined_usage_object"] assert isinstance(estimated, Usage) - expected = litellm_module.token_counter(model="gpt-3.5-turbo", messages=messages) + litellm_module.token_counter( - model="gpt-3.5-turbo", text=dispatched_system - ) + expected = litellm_module.token_counter( + model="gpt-3.5-turbo", messages=messages + ) + litellm_module.token_counter(model="gpt-3.5-turbo", text=dispatched_system) assert estimated.prompt_tokens == expected @@ -916,9 +902,7 @@ def test_create_model_info_response_includes_max_tokens_from_lookup(): model_id="some-model", provider="openai", llm_router=None, - get_model_info=lambda _model: _fake_model_info( - max_input_tokens=128000, max_output_tokens=16384 - ), + get_model_info=lambda _model: _fake_model_info(max_input_tokens=128000, max_output_tokens=16384), ) assert response["id"] == "some-model" @@ -935,9 +919,7 @@ def test_create_model_info_response_does_not_call_router_group_info(): model_id="some-model", provider="openai", llm_router=router, - get_model_info=lambda _model: _fake_model_info( - max_input_tokens=128000, max_output_tokens=16384 - ), + get_model_info=lambda _model: _fake_model_info(max_input_tokens=128000, max_output_tokens=16384), ) router.get_model_group_info.assert_not_called() @@ -968,9 +950,7 @@ def test_create_model_info_response_deployment_limits_override_cost_map(): model_id="gpt-4o", provider="openai", llm_router=router, - get_model_info=lambda _model: _fake_model_info( - max_input_tokens=128000, max_output_tokens=16384 - ), + get_model_info=lambda _model: _fake_model_info(max_input_tokens=128000, max_output_tokens=16384), ) assert response["max_input_tokens"] == 200000 @@ -1008,9 +988,7 @@ def test_create_model_info_response_survives_malformed_cost_map_limits(bad_value model_id="some-model", provider="openai", llm_router=None, - get_model_info=lambda _model: _fake_model_info( - max_input_tokens=bad_value, max_output_tokens=bad_value - ), + get_model_info=lambda _model: _fake_model_info(max_input_tokens=bad_value, max_output_tokens=bad_value), ) assert response["id"] == "some-model" @@ -1023,9 +1001,7 @@ def test_create_model_info_response_keeps_valid_cost_map_limit_beside_malformed_ model_id="some-model", provider="openai", llm_router=None, - get_model_info=lambda _model: _fake_model_info( - max_input_tokens="128,000", max_output_tokens=16384 - ), + get_model_info=lambda _model: _fake_model_info(max_input_tokens="128,000", max_output_tokens=16384), ) assert "max_input_tokens" not in response @@ -1068,9 +1044,7 @@ def test_create_model_info_response_emits_integer_token_counts(): model_id="some-model", provider="openai", llm_router=None, - get_model_info=lambda _model: _fake_model_info( - max_input_tokens=128000, max_output_tokens=16384 - ), + get_model_info=lambda _model: _fake_model_info(max_input_tokens=128000, max_output_tokens=16384), ) assert isinstance(response["max_input_tokens"], int) @@ -1119,9 +1093,7 @@ def test_create_model_info_response_no_router_keeps_base_fields(): def test_create_model_info_response_reads_real_cost_map(): - response = create_model_info_response( - model_id="gpt-4o", provider="openai", llm_router=None - ) + response = create_model_info_response(model_id="gpt-4o", provider="openai", llm_router=None) assert isinstance(response["max_input_tokens"], int) assert response["max_input_tokens"] > 0 @@ -1205,10 +1177,7 @@ class TestPostCallFailureHookLLMExceptionAlerting: @pytest.mark.asyncio async def test_http_exception_does_not_alert(self): - assert ( - await self._alerted(HTTPException(status_code=400, detail="blocked")) - is False - ) + assert await self._alerted(HTTPException(status_code=400, detail="blocked")) is False @pytest.mark.asyncio async def test_genuine_llm_api_error_still_alerts(self): @@ -1241,9 +1210,7 @@ class TestPostCallFailureHookProxyExceptionLogging: await proxy_logging_obj.post_call_failure_hook( request_data={}, original_exception=exc, - user_api_key_dict=UserAPIKeyAuth( - api_key="sk-test", request_route=request_route - ), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", request_route=request_route), ) return handle_mock.await_count > 0 @@ -1260,20 +1227,12 @@ class TestPostCallFailureHookProxyExceptionLogging: @pytest.mark.asyncio async def test_proxy_exception_on_llm_route_is_logged(self): - assert ( - await self._logged(self._block(), request_route="/v1/chat/completions") - is True - ) + assert await self._logged(self._block(), request_route="/v1/chat/completions") is True @pytest.mark.asyncio async def test_generic_exception_on_llm_route_is_not_logged(self): # A raw provider/unknown exception is logged by the LLM call path, not here. - assert ( - await self._logged( - Exception("upstream 503"), request_route="/v1/chat/completions" - ) - is False - ) + assert await self._logged(Exception("upstream 503"), request_route="/v1/chat/completions") is False class TestShouldUseSmtpSsl: @@ -1307,9 +1266,7 @@ class TestCreateSmtpConnection: patch("smtplib.SMTP_SSL") as mock_smtp_ssl, patch("smtplib.SMTP") as mock_smtp, ): - result = _create_smtp_connection( - smtp_host="mail.example.com", smtp_port=465 - ) + result = _create_smtp_connection(smtp_host="mail.example.com", smtp_port=465) mock_smtp.assert_not_called() assert result is mock_smtp_ssl.return_value @@ -1329,9 +1286,7 @@ class TestCreateSmtpConnection: patch("smtplib.SMTP_SSL") as mock_smtp_ssl, patch("smtplib.SMTP") as mock_smtp, ): - result = _create_smtp_connection( - smtp_host="mail.example.com", smtp_port=587 - ) + result = _create_smtp_connection(smtp_host="mail.example.com", smtp_port=587) mock_smtp_ssl.assert_not_called() assert result is mock_smtp.return_value @@ -1352,9 +1307,7 @@ class TestSendEmailStartTls: monkeypatch.delenv("SMTP_USE_SSL", raising=False) mock_server = MagicMock(spec=smtplib.SMTP) - with patch( - "litellm.proxy.utils._create_smtp_connection" - ) as mock_create_connection: + with patch("litellm.proxy.utils._create_smtp_connection") as mock_create_connection: mock_create_connection.return_value.__enter__.return_value = mock_server await send_email( receiver_email="receiver@example.com", @@ -1690,9 +1643,7 @@ def test_a_failed_dispatch_is_estimated_as_input_only(): usage = _estimate_dispatched_failure_usage(FAILURE_USAGE_MODEL, ONE_USER_MESSAGE, None) assert usage is not None - assert usage.prompt_tokens == _count_request_input_tokens( - FAILURE_USAGE_MODEL, ONE_USER_MESSAGE, None - ) + assert usage.prompt_tokens == _count_request_input_tokens(FAILURE_USAGE_MODEL, ONE_USER_MESSAGE, None) assert usage.completion_tokens == 0 assert usage.total_tokens == usage.prompt_tokens @@ -1759,9 +1710,7 @@ def test_a_request_that_reached_a_provider_bills_its_input_at_no_cost(): def test_a_failure_that_cost_the_provider_nothing_lifts_nothing(model_call_details, dispatched): from litellm.proxy.utils import _failure_usage_to_lift - assert _failure_usage_to_lift( - model_call_details=model_call_details, request_body={}, dispatched=dispatched - ) is None + assert _failure_usage_to_lift(model_call_details=model_call_details, request_body={}, dispatched=dispatched) is None def test_the_no_upstream_call_key_the_module_uses_is_the_one_asserted_above(): @@ -1829,3 +1778,102 @@ def test_a_dispatched_failure_lifts_the_four_fields_the_spend_log_needs(): assert lifted["response_cost"] == 0.0 assert lifted["combined_usage_object"].prompt_tokens > 0 assert lifted["standard_logging_object"] == {"id": "log-1"} + + +@pytest.mark.asyncio +async def test_proxy_only_error_expected_4xx_skips_traceback_for_both_handlers(monkeypatch): + """Regression for LIT-6043: an expected 4xx must not format a traceback for + either the async or the threaded sync failure handler.""" + import asyncio + + import litellm + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.proxy._types import UserAPIKeyAuth + + monkeypatch.setattr(litellm, "failure_callback", []) + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + captured = {} + sync_ran = asyncio.Event() + loop = asyncio.get_running_loop() + + async def fake_async_failure(self, exception, traceback_exception, *args, **kwargs): + captured["async_traceback"] = traceback_exception + + def fake_sync_failure(self, exception, traceback_exception, *args, **kwargs): + captured["sync_traceback"] = traceback_exception + loop.call_soon_threadsafe(sync_ran.set) + + orig_async_failure = Logging.async_failure_handler + orig_sync_failure = Logging.failure_handler + Logging.async_failure_handler = fake_async_failure + Logging.failure_handler = fake_sync_failure + try: + try: + raise HTTPException(status_code=400, detail="Invalid model name passed in") + except HTTPException as exc: + await proxy_logging_obj._handle_logging_proxy_only_error( + request_data={ + "model": "does-not-exist", + "messages": [{"role": "user", "content": "hi"}], + }, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", request_route="/v1/chat/completions"), + route="/v1/chat/completions", + original_exception=exc, + ) + await asyncio.wait_for(sync_ran.wait(), timeout=5) + finally: + Logging.async_failure_handler = orig_async_failure + Logging.failure_handler = orig_sync_failure + + assert captured["async_traceback"] == "" + assert captured["sync_traceback"] == "" + + +@pytest.mark.asyncio +async def test_proxy_only_error_5xx_keeps_traceback_and_runs_sync_callbacks(monkeypatch): + """Unexpected (5xx) errors keep the full traceback, and a configured + sync-only failure callback still gets its threaded handler.""" + import asyncio + + import litellm + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.proxy._types import UserAPIKeyAuth + + def _custom_sync_callback(kwargs, completion_response, start_time, end_time): + pass + + monkeypatch.setattr(litellm, "failure_callback", [_custom_sync_callback]) + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + captured = {} + sync_ran = asyncio.Event() + loop = asyncio.get_running_loop() + + async def fake_async_failure(self, exception, traceback_exception, *args, **kwargs): + captured["async_traceback"] = traceback_exception + + def fake_sync_failure(self, *args, **kwargs): + loop.call_soon_threadsafe(sync_ran.set) + + orig_async_failure = Logging.async_failure_handler + orig_sync_failure = Logging.failure_handler + Logging.async_failure_handler = fake_async_failure + Logging.failure_handler = fake_sync_failure + try: + try: + raise HTTPException(status_code=500, detail="internal error") + except HTTPException as exc: + await proxy_logging_obj._handle_logging_proxy_only_error( + request_data={ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + }, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", request_route="/v1/chat/completions"), + route="/v1/chat/completions", + original_exception=exc, + ) + await asyncio.wait_for(sync_ran.wait(), timeout=5) + finally: + Logging.async_failure_handler = orig_async_failure + Logging.failure_handler = orig_sync_failure + + assert "test_proxy_utils" in captured["async_traceback"] From 896e2598dadcab650e4109320c69e527ede62b18 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 25 Aug 2026 12:45:23 -0700 Subject: [PATCH 92/93] fix(caching): keep upstream RedisCluster on redis-py with per-connection recovery (#38171) Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../caching/redis_cluster_node_isolation.py | 34 ++++++++++++++-- .../test_redis_cluster_node_isolation.py | 39 ++++++++++++++++++- 2 files changed, 69 insertions(+), 4 deletions(-) diff --git a/litellm/caching/redis_cluster_node_isolation.py b/litellm/caching/redis_cluster_node_isolation.py index 8b0c120e80c..ae8c78709d9 100644 --- a/litellm/caching/redis_cluster_node_isolation.py +++ b/litellm/caching/redis_cluster_node_isolation.py @@ -18,6 +18,14 @@ already does when one of its pooled connections errors), leaving every other nod connections untouched. Every other branch (MOVED, ASK, CLUSTERDOWN, slot-not-covered, retry-exhaustion) is unchanged from upstream, since those already carry real evidence the topology changed. + +redis-py 8.x fixed this upstream with gentler machinery than this override's +``node.disconnect()`` (which also kills connections other coroutines are mid-operation +on, so one timeout cascades into a reconnect storm and, with TLS, a fresh handshake per +killed connection): it marks in-use connections for reconnect only after their current +operation completes, disconnects only the idle pooled ones, and defers reinitialization +to the outer retry loop. When the installed ``ClusterNode`` has that per-connection +recovery API, the factory returns the base ``RedisCluster`` unmodified. """ import asyncio @@ -72,8 +80,16 @@ class _ClusterAttrs(Protocol): _VERIFIED_REDIS_VERSIONS: Final = frozenset({"5.3.1"}) -def get_litellm_async_redis_cluster_class() -> type["_AsyncRedisClusterType"]: - """Builds the ``RedisCluster`` subclass with the per-node isolation fix. +def get_litellm_async_redis_cluster_class( + cluster_node_class: type | None = None, +) -> type["_AsyncRedisClusterType"]: + """Returns the base ``RedisCluster`` when the installed redis-py already recovers a + node-level connection error per-connection (8.x+), else builds the ``RedisCluster`` + subclass with the per-node isolation fix for older versions whose upstream branch + tears down the whole cluster client. + + ``cluster_node_class`` exists for dependency injection in tests; production callers + leave it unset and the installed ``ClusterNode`` is used. Imported lazily because this module is reachable from a base ``import litellm`` while redis is not a base dependency. Cheap to call repeatedly: the underlying redis @@ -81,7 +97,10 @@ def get_litellm_async_redis_cluster_class() -> type["_AsyncRedisClusterType"]: """ import redis from redis.asyncio.cluster import ( - RedisCluster as _BaseAsyncRedisCluster, # pyright: ignore[reportUnknownVariableType] # redis-py ships no resolvable stub for this class under the repo's current (stale) types-redis pin + ClusterNode as _AsyncClusterNode, # pyright: ignore[reportUnknownVariableType] # redis-py ships no resolvable stub for this class under the repo's current (stale) types-redis pin + ) + from redis.asyncio.cluster import ( + RedisCluster as _BaseAsyncRedisCluster, # pyright: ignore[reportUnknownVariableType] # same stale-stub gap as the import above ) from redis.cluster import get_node_name from redis.commands import READ_COMMANDS @@ -98,6 +117,15 @@ def get_litellm_async_redis_cluster_class() -> type["_AsyncRedisClusterType"]: from redis.exceptions import ConnectionError as _RedisConnectionError from redis.exceptions import TimeoutError as _RedisTimeoutError + node_class: Final = cluster_node_class if cluster_node_class is not None else _AsyncClusterNode + if hasattr(node_class, "update_active_connections_for_reconnect"): + verbose_logger.debug( + "redis-py %s recovers a node-level connection error per-connection upstream; " + "using the base RedisCluster without litellm's node-isolation override.", + redis.__version__, + ) + return _BaseAsyncRedisCluster + if redis.__version__ not in _VERIFIED_REDIS_VERSIONS: verbose_logger.warning( "redis-py %s is not in the set this cluster-teardown-storm fix was verified " diff --git a/tests/test_litellm/caching/test_redis_cluster_node_isolation.py b/tests/test_litellm/caching/test_redis_cluster_node_isolation.py index f4cd3ab20ef..c16ceec8c31 100644 --- a/tests/test_litellm/caching/test_redis_cluster_node_isolation.py +++ b/tests/test_litellm/caching/test_redis_cluster_node_isolation.py @@ -28,6 +28,14 @@ if TYPE_CHECKING: from redis.asyncio.cluster import RedisCluster as _AsyncRedisClusterType +class _NodeClassWithPerConnectionRecovery: + def update_active_connections_for_reconnect(self) -> None: ... + + +class _NodeClassWithoutPerConnectionRecovery: + pass + + class _FakeClusterNode: def __init__(self, name: str, raises: Exception | None = None, response: object = None) -> None: self.name = name @@ -47,7 +55,9 @@ class _FakeNodesManager: def _build_cluster_instance() -> "_AsyncRedisClusterType": - cluster_cls = get_litellm_async_redis_cluster_class() + cluster_cls = get_litellm_async_redis_cluster_class( + cluster_node_class=_NodeClassWithoutPerConnectionRecovery + ) instance = cluster_cls.__new__(cluster_cls) instance.RedisClusterRequestTTL = 1 instance.reinitialize_counter = 0 @@ -58,6 +68,33 @@ def _build_cluster_instance() -> "_AsyncRedisClusterType": return instance +def test_per_connection_recovery_redis_py_gets_the_unmodified_upstream_class() -> None: + """Regression (redis-py 8.x): when upstream ClusterNode already recovers a node-level + connection error per-connection, the factory must NOT install the copied override, + whose node.disconnect() also kills connections other coroutines are mid-operation on.""" + from redis.asyncio.cluster import RedisCluster + + cluster_cls = get_litellm_async_redis_cluster_class( + cluster_node_class=_NodeClassWithPerConnectionRecovery + ) + + assert cluster_cls is RedisCluster + + +def test_pre_recovery_redis_py_still_gets_the_node_isolation_override() -> None: + """Old redis-py (5.x) responds to a node-level error with a full-cluster aclose(), + so those versions must keep litellm's per-node isolation override.""" + from redis.asyncio.cluster import RedisCluster + + cluster_cls = get_litellm_async_redis_cluster_class( + cluster_node_class=_NodeClassWithoutPerConnectionRecovery + ) + + assert cluster_cls is not RedisCluster + assert issubclass(cluster_cls, RedisCluster) + assert "_execute_command" in cluster_cls.__dict__ + + @pytest.mark.asyncio @pytest.mark.parametrize("error_cls", [RedisConnectionError, RedisTimeoutError]) async def test_node_level_error_resets_only_that_node_not_the_whole_client(error_cls: type[Exception]) -> None: From 0458accaa0b968ca37b98322012512572fadd0c7 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 25 Aug 2026 12:47:37 -0700 Subject: [PATCH 93/93] perf(auth): drop guaranteed-miss internal-cache Redis read from team object lookup (#38073) Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 14 ---- .../access_group_endpoints.py | 2 - .../proxy/auth/test_auth_checks.py | 82 +++++++++++++++++-- 3 files changed, 76 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 4af942a357e..5d52af3b364 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -71,7 +71,6 @@ from litellm.proxy.auth.budget_throttle import ( ) from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import publish_auth_cache_invalidation -from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, _safe_get_request_query_params, @@ -2741,20 +2740,9 @@ async def _get_team_object_from_user_api_key_cache( async def _get_team_object_from_cache( key: str, - proxy_logging_obj: ProxyLogging | None, user_api_key_cache: UserApiKeyCache, parent_otel_span: Span | None, ) -> LiteLLM_TeamTableCachedObj | None: - ## INTERNAL USAGE CACHE (plain DualCache) — checked before UserApiKeyCache stores ## - if proxy_logging_obj is not None and proxy_logging_obj.internal_usage_cache.dual_cache: - cached_raw: Final = await proxy_logging_obj.internal_usage_cache.dual_cache.async_get_cache( - key=key, parent_otel_span=parent_otel_span - ) - if cached_raw is not None: - from_internal: Final = CacheCodec.deserialize(cached_raw, LiteLLM_TeamTableCachedObj) - if from_internal is not None: - return from_internal - decoded: Final = await user_api_key_cache.async_get_cache( key=key, parent_otel_span=parent_otel_span, @@ -2790,7 +2778,6 @@ async def get_team_object( if not check_db_only: cached_team_obj: Final = await _get_team_object_from_cache( key=key, - proxy_logging_obj=proxy_logging_obj, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, ) @@ -2953,7 +2940,6 @@ async def get_team_object_by_alias( cached_team_obj: Final = await _get_team_object_from_cache( key=cache_key, - proxy_logging_obj=proxy_logging_obj, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, ) diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index 2271501d480..c803d87e436 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -210,7 +210,6 @@ async def _patch_team_caches_add_access_group( for team_id in team_ids: cached_team = await _get_team_object_from_cache( key=f"team_id:{team_id}", - proxy_logging_obj=proxy_logging_obj, user_api_key_cache=user_api_key_cache, parent_otel_span=None, ) @@ -240,7 +239,6 @@ async def _patch_team_caches_remove_access_group( for team_id in team_ids: cached_team = await _get_team_object_from_cache( key=f"team_id:{team_id}", - proxy_logging_obj=proxy_logging_obj, user_api_key_cache=user_api_key_cache, parent_otel_span=None, ) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index abe73f7d05c..06fd3730f7d 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -4869,10 +4869,9 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): 3. When team_alias is None, NO alias-key operation happens (no delete of an empty-keyed entry, no spurious write). 4. DELETES the team_id-keyed entry from the internal usage cache - BEFORE the fresh write (LIT-4391). `_get_team_object_from_cache` - consults the internal usage cache first, so a leftover copy there - (backfilled from a Redis shared with `user_api_key_cache`) would - keep serving the pre-update team allowlist. + BEFORE the fresh write (LIT-4391). `_get_team_object_from_cache` no + longer reads the internal usage cache (LIT-5944), but the delete + protects mixed-version rolling deploys where older workers still do. """ from unittest.mock import AsyncMock, MagicMock @@ -4981,8 +4980,10 @@ async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): Regression test for LIT-4391: keys with models=["all-team-models"] kept getting 403 team_model_access_denied for models added via /team/update. - `_get_team_object_from_cache` consults the internal usage cache BEFORE - `user_api_key_cache`. When both share one Redis (enable_redis_auth_cache), + `_get_team_object_from_cache` used to consult the internal usage cache + BEFORE `user_api_key_cache` (removed in LIT-5944; this test now also + guards against reintroducing that read). + When both share one Redis (enable_redis_auth_cache), any team read backfills the internal cache's in-memory tier with the team object. `_cache_team_object` (the /team/update refresh) only wrote `user_api_key_cache`, so that backfilled copy kept shadowing the update @@ -5054,6 +5055,75 @@ async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): ) +class _CountingFakeRedis(_SharedFakeRedis): + """Counts per-key Redis round-trips so tests can pin the number of + network operations a code path issues.""" + + def __init__(self): + super().__init__() + self.get_calls: int = 0 + + async def async_get_cache(self, key, **kwargs): + self.get_calls += 1 + return await super().async_get_cache(key, **kwargs) + + +@pytest.mark.asyncio +async def test_warm_team_object_reads_issue_no_redis_ops_lit_5944(): + """ + Regression test for LIT-5944: project/team-scoped virtual-key requests + paid ~4 awaited Redis GETs per request just to re-read the team object. + + `_get_team_object_from_cache` used to consult + `proxy_logging_obj.internal_usage_cache.dual_cache` (in-memory TTL 1s, + Redis-backed) BEFORE `user_api_key_cache`. Nothing writes team objects + into that internal cache — `_cache_team_object` only DELETES the key + there — so when `user_api_key_cache` has no Redis tier the shared Redis + key stays absent forever and every team lookup in the auth hot path + (4 call sites per chat-completion request) became a guaranteed-miss + Redis round-trip, saturating the event loop at high TPS. + + Pins: once `_cache_team_object` has cached a team, repeated + `get_team_object` reads are served from `user_api_key_cache`'s in-memory + tier and issue ZERO Redis operations. + """ + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.auth.auth_checks import _cache_team_object, get_team_object + + team_id = "team-lit-5944" + counting_redis = _CountingFakeRedis() + user_api_key_cache = UserApiKeyCache() + proxy_logging_obj = MagicMock() + proxy_logging_obj.internal_usage_cache.dual_cache = DualCache( + redis_cache=counting_redis, + default_in_memory_ttl=1, + ) + prisma_client = MagicMock() + + await _cache_team_object( + team_id=team_id, + team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a"]), + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + for _ in range(4): + team_obj = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert team_obj is not None and team_obj.models == ["model-a"] + + assert counting_redis.get_calls == 0, ( + "Warm team-object reads must be served from user_api_key_cache's " + "in-memory tier without any Redis round-trips. " + f"Got {counting_redis.get_calls} Redis GETs for 4 get_team_object calls." + ) + + @pytest.mark.asyncio async def test_cache_team_object_tolerates_cache_invalidation_failures(): """