From 5d777c16d9e59c690886d978f7546b130bd2432d Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 26 Sep 2026 16:53:02 -0700 Subject: [PATCH 01/65] fix(mcp): align hub publication status and controls (#43241) * fix(mcp): align hub publication status and controls * refactor(mcp): keep hub visibility guard outside table rendering --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> --- cookbook/litellm_proxy_server/mcp/README.md | 37 ++++ .../mcp_server/mcp_server_manager.py | 24 +-- .../mcp_management_endpoints.py | 35 ++-- .../public_endpoints/public_endpoints.py | 14 +- .../mcp_server/test_mcp_server_manager.py | 54 ++++++ .../test_mcp_management_endpoints.py | 156 ++++++++++++++++ .../public_endpoints/test_public_endpoints.py | 66 +++++-- .../_components/MCPPermissionManagement.tsx | 4 +- .../_components/MCPServerCard.test.tsx | 12 ++ .../mcp-servers/_components/MCPServerCard.tsx | 19 +- .../_components/mcp_server_view.test.tsx | 11 +- .../_components/mcp_server_view.tsx | 21 +-- .../mcp-servers/_components/utils.test.tsx | 29 +++ .../mcp-servers/_components/utils.tsx | 35 +++- .../AIHub/MCPHubTableColumns.test.tsx | 19 +- .../components/AIHub/MCPHubTableColumns.tsx | 6 +- .../components/AIHub/ModelHubTable.test.tsx | 30 ++- .../src/components/AIHub/ModelHubTable.tsx | 15 +- .../AIHub/forms/MakeMCPPublicForm.test.tsx | 174 +++++++++++++----- .../AIHub/forms/MakeMCPPublicForm.tsx | 121 ++++++++---- .../src/components/mcp_tools/types.tsx | 2 + 21 files changed, 720 insertions(+), 164 deletions(-) create mode 100644 cookbook/litellm_proxy_server/mcp/README.md diff --git a/cookbook/litellm_proxy_server/mcp/README.md b/cookbook/litellm_proxy_server/mcp/README.md new file mode 100644 index 00000000000..aeee0719019 --- /dev/null +++ b/cookbook/litellm_proxy_server/mcp/README.md @@ -0,0 +1,37 @@ +# Publish MCP servers in the AI Hub + +Set `litellm_settings.public_mcp_servers` to the concrete IDs of the servers you want listed in the public AI Hub. Pin `server_id` in each configuration entry so the publication list stays stable across deployments + +```yaml +mcp_servers: + documentation: + server_id: documentation-mcp + url: https://mcp.example.com/mcp + transport: http + available_on_public_internet: true + +litellm_settings: + public_mcp_hub_strict_whitelist: true + public_mcp_servers: + - documentation-mcp +``` + +Use `documentation-mcp`, the `server_id`, in the publication list. The configuration key `documentation`, display names, and aliases are not publication IDs. Database-created servers use the ID returned by `/v1/mcp/server` + +The dashboard's **AI Hub > MCP Hub > Manage MCP Hub Visibility** dialog edits this same list. Its YAML example includes the selected server IDs. With database-backed configuration (`store_model_in_db: true`), a value declared in YAML is owned by that file: edit the file and reload, or remove that key from YAML to let the dashboard manage it in the database. File-backed deployments can save the list directly to their configuration file + +To remove all explicit entries, save an empty selection in the dialog or configure: + +```yaml +litellm_settings: + public_mcp_hub_strict_whitelist: true + public_mcp_servers: [] +``` + +## Hub listing and network access + +The **Hub listing** column in AI Hub identifies servers that appear in `/public/mcp_hub`. The dashboard derives this status from the current registry and publication settings. Setting `mcp_info.is_public` on a server does not publish it; that response field is derived metadata. `mcp_info.is_public_explicit` identifies registered servers included in the explicit publication list + +Gateway cards and server details show **All Networks** when `available_on_public_internet` is enabled or the server is explicitly published in `public_mcp_servers`. They show **Internal Only** when both are false. The per-server flag defaults to `true`; explicit publication overrides a disabled flag for compatibility. Older proxies that omit the metadata needed to determine access show **Unknown**. These labels describe allowed client IPs; authentication and tool permissions still apply + +The default `public_mcp_hub_strict_whitelist: true` lists only registered servers in `public_mcp_servers`. Legacy mode (`false`) additionally lists registered servers with `available_on_public_internet: true`. In legacy mode, clearing the explicit publication list leaves these automatically listed servers visible. Enable strict mode when the publication list should fully determine hub visibility diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d0d9100971d..31896d9ddc5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -6799,6 +6799,16 @@ class MCPServerManager: return server return None + @staticmethod + def _is_public_mcp_server(server: MCPServer, public_ids: Container[str]) -> bool: + return server.server_id in public_ids or ( + not litellm.public_mcp_hub_strict_whitelist and server.available_on_public_internet + ) + + def is_mcp_server_public(self, server_id: str) -> bool: + server: Final = self.registry.get(server_id) or self.config_mcp_servers.get(server_id) + return server is not None and self._is_public_mcp_server(server, litellm.public_mcp_servers or ()) + def get_public_mcp_servers(self) -> list[MCPServer]: """ Return the MCP servers published to the AI Hub via /v1/mcp/make_public. @@ -6816,18 +6826,8 @@ class MCPServerManager: deployments that relied on the OR-with-default semantics; will be removed in a future release. """ - if litellm.public_mcp_hub_strict_whitelist: - if litellm.public_mcp_servers is None: - return [] - public_ids = set(litellm.public_mcp_servers) - return [server for server in self.get_registry().values() if server.server_id in public_ids] - - public_ids = set(litellm.public_mcp_servers or []) - return [ - server - for server in self.get_registry().values() - if server.available_on_public_internet or server.server_id in public_ids - ] + public_ids: Final = frozenset(litellm.public_mcp_servers or ()) + return [server for server in self.get_registry().values() if self._is_public_mcp_server(server, public_ids)] def expand_permission_list(self, identifiers: list[str]) -> list[str]: """ diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 346b1a75ac6..deb0e00ff9b 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -650,7 +650,16 @@ if MCP_AVAILABLE: if hasattr(redacted_server, "credentials"): setattr(redacted_server, "credentials", _preserved_admin_config_credentials(redacted_server.credentials)) - return redacted_server + is_public: Final = global_mcp_server_manager.is_mcp_server_public(redacted_server.server_id) + return redacted_server.model_copy( + update={ + "mcp_info": { + **(redacted_server.mcp_info or {}), + "is_public": is_public, + "is_public_explicit": is_public and redacted_server.server_id in (litellm.public_mcp_servers or ()), + } + } + ) def _preserved_admin_config_credentials( credentials: "MCPCredentials | str | None", @@ -832,10 +841,10 @@ if MCP_AVAILABLE: sanitized.updated_at = None # `mcp_info` is arbitrary metadata; keep only an explicit safe subset. - is_public = False - if isinstance(sanitized.mcp_info, dict): - is_public = bool(sanitized.mcp_info.get("is_public")) - sanitized.mcp_info = {"is_public": True} if is_public else None + sanitized.mcp_info = { + "is_public": (sanitized.mcp_info or {}).get("is_public") is True, + "is_public_explicit": (sanitized.mcp_info or {}).get("is_public_explicit") is True, + } return sanitized @@ -1260,14 +1269,6 @@ if MCP_AVAILABLE: for server in redacted_mcp_servers: server.connected_app_reachable = server.server_id in reachable_ids - # augment the mcp servers with public status - if litellm.public_mcp_servers is not None: - for server in redacted_mcp_servers: - if server.server_id in litellm.public_mcp_servers: - if server.mcp_info is None: - server.mcp_info = {} - server.mcp_info["is_public"] = True - # Annotate has_user_credential for BYOK servers (single batched query) from litellm.proxy.proxy_server import prisma_client as _byok_prisma_client @@ -3041,9 +3042,6 @@ if MCP_AVAILABLE: }, ) - if litellm.public_mcp_servers is None: - litellm.public_mcp_servers = [] - for server_id in request.mcp_server_ids: server = global_mcp_server_manager.get_mcp_server_by_id(server_id=server_id) if server is None: @@ -3052,16 +3050,15 @@ if MCP_AVAILABLE: detail=f"MCP Server with ID {server_id} not found", ) - litellm.public_mcp_servers = request.mcp_server_ids - # Update config with new settings if "litellm_settings" not in config or config["litellm_settings"] is None: config["litellm_settings"] = {} - config["litellm_settings"]["public_mcp_servers"] = litellm.public_mcp_servers + config["litellm_settings"]["public_mcp_servers"] = request.mcp_server_ids # Save the updated config await proxy_config.save_config(new_config=config) + litellm.public_mcp_servers = request.mcp_server_ids verbose_proxy_logger.debug( "Updated public mcp servers to: %s by user: %s", litellm.public_mcp_servers, user_api_key_dict.user_id diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index 26a5c44fce1..bba5ef681d0 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -300,7 +300,19 @@ async def get_mcp_servers(): ) public_mcp_servers: Final = global_mcp_server_manager.get_public_mcp_servers() - return [MCPPublicServer.model_validate(server.model_dump()) for server in public_mcp_servers] + return [ + MCPPublicServer.model_validate( + { + **server.model_dump(), + "mcp_info": { + **(server.mcp_info or {}), + "is_public": True, + "is_public_explicit": server.server_id in (litellm.public_mcp_servers or ()), + }, + } + ) + for server in public_mcp_servers + ] @router.get( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index db476e86043..70ef4312f4c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -9566,6 +9566,60 @@ class TestGetPublicMCPServers: manager.config_mcp_servers[s.server_id] = s return manager + @pytest.mark.parametrize("registered_in", ("config", "database", "both", "neither")) + @pytest.mark.parametrize("public_ids", (None, [], ["server-id"], ["server-alias"], ["Server Name"])) + @pytest.mark.parametrize( + "strict,network_access,implicitly_public", + ((True, True, False), (True, False, False), (False, True, True), (False, False, False)), + ) + def test_public_status_agrees_with_hub_membership( + self, + registered_in: Literal["config", "database", "both", "neither"], + public_ids: list[str] | None, + strict: bool, + network_access: bool, + implicitly_public: bool, + ) -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="server-id", + name="server-alias", + alias="server-alias", + server_name="Server Name", + transport=MCPTransport.http, + available_on_public_internet=network_access, + mcp_info={"is_public": True, "description": "Preserve custom metadata"}, + ) + config_server: Final = ( + server.model_copy(update={"available_on_public_internet": not network_access}) + if registered_in == "both" + else server + ) + manager.config_mcp_servers = ( + {server.server_id: config_server} if registered_in in ("config", "both") else {} + ) + manager.registry = {server.server_id: server} if registered_in in ("database", "both") else {} + original_server: Final = server.model_dump() + original_config_server: Final = config_server.model_dump() + expected_public: Final = registered_in != "neither" and ( + public_ids == [server.server_id] or implicitly_public + ) + + with ( + patch("litellm.public_mcp_servers", public_ids), + patch("litellm.public_mcp_hub_strict_whitelist", strict), + ): + public_servers: Final = manager.get_public_mcp_servers() + assert manager.is_mcp_server_public(server.server_id) is expected_public + assert [item.server_id for item in public_servers] == ( + [server.server_id] if expected_public else [] + ) + assert manager.is_mcp_server_public("server-alias") is False + assert manager.is_mcp_server_public("missing-server") is False + + assert server.model_dump() == original_server + assert config_server.model_dump() == original_config_server + @patch("litellm.public_mcp_servers", None) def test_returns_empty_when_whitelist_is_none(self): """No /make_public call yet → hub returns nothing, regardless of diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 8aa16b817fc..11b3dcf54bc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -33,6 +33,7 @@ from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, LiteLLM_MCPServerTable, LitellmUserRoles, + MakeMCPServersPublicRequest, MCPTransport, MCPUserCredentialResponse, NewMCPServerRequest, @@ -154,6 +155,161 @@ def patch_proxy_general_settings(settings: dict): ) +@pytest.mark.asyncio +@pytest.mark.parametrize("from_db", (False, True)) +@pytest.mark.parametrize( + "strict,explicit,expected_public", + ((True, True, True), (True, False, False), (False, False, True)), +) +async def test_mcp_publication_list_and_detail_derive_current_status( + from_db: bool, strict: bool, explicit: bool, expected_public: bool +) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="publication-server", + name="publication-server", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + available_on_public_internet=True, + mcp_info={ + "is_public": not expected_public, + "is_public_explicit": not explicit, + "description": "Keep this description", + }, + ) + manager.registry = {server.server_id: server} if from_db else {} + manager.config_mcp_servers = {} if from_db else {server.server_id: server} + record: Final = manager._build_mcp_server_table(server) + original_metadata: Final = dict(server.mcp_info or {}) + admin: Final = generate_mock_user_api_key_auth() + + with ( + patch("litellm.public_mcp_servers", [server.server_id] if explicit else []), + patch("litellm.public_mcp_hub_strict_whitelist", strict), + patch.object(mgmt_endpoints, "global_mcp_server_manager", manager), + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()), + patch.object(mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=record if from_db else None)), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": "view_all"}), + ): + listing: Final = await mgmt_endpoints.fetch_all_mcp_servers( + user_api_key_dict=admin, team_id=None, connected_app_view=False + ) + detail: Final = await mgmt_endpoints.fetch_mcp_server( + request=_make_mock_request(), server_id=server.server_id, user_api_key_dict=admin + ) + assert len(listing) == 1 + for projected in (listing[0], detail): + assert projected.mcp_info == { + "is_public": expected_public, + "is_public_explicit": explicit, + "description": "Keep this description", + } + assert bool(manager.get_public_mcp_servers()) is expected_public + + assert server.mcp_info == original_metadata + assert record.mcp_info == original_metadata + + +@pytest.mark.parametrize("approval_status", ("pending_review", "rejected", "draft", "active")) +@pytest.mark.parametrize("strict", (False, True)) +def test_mcp_publication_projection_excludes_unregistered_lifecycle_records( + approval_status: str, strict: bool +) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + record: Final = LiteLLM_MCPServerTable( + server_id="unregistered-server", + transport=MCPTransport.http, + approval_status=approval_status, + credentials={"auth_value": "test-secret"}, + available_on_public_internet=True, + mcp_info={"is_public": True, "is_public_explicit": True}, + ) + original: Final = record.model_dump() + with ( + patch("litellm.public_mcp_servers", [record.server_id]), + patch("litellm.public_mcp_hub_strict_whitelist", strict), + patch.object(mgmt_endpoints, "global_mcp_server_manager", MCPServerManager()), + ): + for project in ( + mgmt_endpoints._redact_mcp_credentials, + mgmt_endpoints._sanitize_mcp_server_for_non_admin, + mgmt_endpoints._sanitize_mcp_server_for_virtual_key, + ): + projected: Final = project(record) + assert projected.mcp_info == {"is_public": False, "is_public_explicit": False} + assert projected.credentials is None + assert record.model_dump() == original + + +@pytest.mark.asyncio +@pytest.mark.parametrize("previous_ids", (None, ["old-server"])) +@pytest.mark.parametrize( + "selected_ids,save_error,role,error_status", + ( + (["new-server"], None, LitellmUserRoles.PROXY_ADMIN, None), + ([], None, LitellmUserRoles.PROXY_ADMIN, None), + (["new-server"], HTTPException(400, "Owned by config file"), LitellmUserRoles.PROXY_ADMIN, 400), + (["new-server"], RuntimeError("Database write failed"), LitellmUserRoles.PROXY_ADMIN, 500), + (["missing-server"], None, LitellmUserRoles.PROXY_ADMIN, 404), + (["new-server"], None, LitellmUserRoles.INTERNAL_USER, 403), + ), +) +async def test_mcp_publication_updates_runtime_only_after_successful_save( + previous_ids: list[str] | None, + selected_ids: list[str], + save_error: HTTPException | RuntimeError | None, + role: LitellmUserRoles, + error_status: int | None, +) -> None: + import litellm + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + manager: Final = MCPServerManager() + server: Final = generate_mock_mcp_server_config_record(server_id="new-server") + manager.config_mcp_servers = {server.server_id: server} + expected_config: Final = {"litellm_settings": {"drop_params": True, "public_mcp_servers": selected_ids}} + + async def save_config(new_config: Mapping[str, object]) -> None: + assert litellm.public_mcp_servers is previous_ids + assert new_config == expected_config + if save_error is not None: + raise save_error + + save: Final = AsyncMock(side_effect=save_config) + proxy_config: Final = SimpleNamespace( + get_config=AsyncMock(return_value={"litellm_settings": {"drop_params": True}}), + save_config=save, + ) + request: Final = MakeMCPServersPublicRequest(mcp_server_ids=selected_ids) + caller: Final = generate_mock_user_api_key_auth(user_role=role) + with ( + patch("litellm.public_mcp_servers", previous_ids), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + manager, + ), + ): + if error_status is None: + response: Final = await mgmt_endpoints.make_mcp_servers_public(request, caller) + assert response["public_mcp_servers"] == selected_ids + assert litellm.public_mcp_servers == selected_ids + else: + with pytest.raises(HTTPException) as error: + await mgmt_endpoints.make_mcp_servers_public(request, caller) + assert error.value.status_code == error_status + assert litellm.public_mcp_servers is previous_ids + + if error_status in (403, 404): + save.assert_not_awaited() + else: + save.assert_awaited_once_with(new_config=expected_config) + + class TestMCPCredentialsTokenExchangeProfile: """token_exchange_profile must be a declared MCPCredentials field so the management API can persist the entra_obo profile. An undeclared key is silently stripped by pydantic when the diff --git a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py index 0dec44af402..18839a65d62 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -1086,43 +1086,73 @@ def test_clean_display_name_passthrough_when_no_suffix(): assert _clean_display_name("") == "" -def test_public_mcp_hub_returns_only_whitelisted_servers(): - """Regression: /public/mcp_hub must gate strictly on - litellm.public_mcp_servers, mirroring /public/model_hub and - /public/agent_hub. Servers with available_on_public_internet=True that - are not on the whitelist must not leak.""" +@pytest.mark.parametrize( + "strict,explicit,expected_listed", + ((True, True, True), (True, False, False), (False, True, True), (False, False, True)), +) +@pytest.mark.parametrize("stored_public", (None, False, True)) +def test_public_mcp_hub_derives_publication_metadata_without_mutating_registry( + strict: bool, + explicit: bool, + expected_listed: bool, + stored_public: bool | None, +) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.proxy._types import MCPTransport - app = FastAPI() + app: Final = FastAPI() app.include_router(router) - app.dependency_overrides[user_api_key_auth] = lambda: MagicMock() - client = TestClient(app) + client: Final = TestClient(app) - listed = MCPServer( + server: Final = MCPServer( server_id="listed", name="listed", server_name="listed", transport=MCPTransport.http, available_on_public_internet=True, + mcp_info=( + { + "is_public": stored_public, + "is_public_explicit": not explicit, + "description": "Preserve custom metadata", + } + if stored_public is not None + else None + ), ) - - mock_manager = MagicMock() - mock_manager.get_public_mcp_servers.return_value = [listed] + unlisted: Final = MCPServer( + server_id="unlisted", + name="unlisted", + transport=MCPTransport.http, + available_on_public_internet=False, + mcp_info={"is_public": True, "is_public_explicit": True}, + ) + manager: Final = MCPServerManager() + manager.config_mcp_servers = {server.server_id: server} + manager.registry = {unlisted.server_id: unlisted} + original_registry: Final = {key: value.model_dump() for key, value in manager.get_registry().items()} with ( - patch("litellm.public_mcp_servers", ["listed"]), + patch("litellm.public_mcp_servers", [server.server_id] if explicit else []), + patch("litellm.public_mcp_hub_strict_whitelist", strict), patch( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", - mock_manager, + manager, ), ): - response = client.get("/public/mcp_hub") + response: Final = client.get("/public/mcp_hub") assert response.status_code == 200 - data = response.json() - assert [item["server_id"] for item in data] == ["listed"] - app.dependency_overrides.clear() + data: Final = response.json() + assert [item["server_id"] for item in data] == ([server.server_id] if expected_listed else []) + if expected_listed: + assert data[0]["mcp_info"] == { + **({"description": "Preserve custom metadata"} if stored_public is not None else {}), + "is_public": True, + "is_public_explicit": explicit, + } + assert {key: value.model_dump() for key, value in manager.get_registry().items()} == original_registry def test_public_mcp_hub_returns_empty_when_whitelist_unset(): diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx index cb423b435ae..a48c991bc7a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx @@ -217,12 +217,12 @@ const MCPPermissionManagement: React.FC = ({
Internal network only - +

- Turn on to restrict access to callers within your internal network only. + Turn on to restrict public IPs. Explicitly published server IDs remain accessible from public IPs.

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx index 100b0ea93d3..d298d9d8145 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx @@ -126,3 +126,15 @@ describe("MCPServerCard per-user credentials", () => { expect(screen.queryByRole("button", { name: "Set" })).not.toBeInTheDocument(); }); }); + +describe("MCPServerCard network access", () => { + it("shows effective network access without a hub listing badge", () => { + renderCard({ + available_on_public_internet: false, + mcp_info: { server_name: "demo_server", is_public: true, is_public_explicit: true }, + }); + + expect(screen.getByText("All Networks")).toBeInTheDocument(); + expect(screen.queryByText(/^Hub:/)).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx index 775809e3670..bb153f94665 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx @@ -13,7 +13,7 @@ import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/comp import { cn } from "@/lib/cva.config"; import { AUTH_TYPE, MCP_REACHABLE_DESCRIPTION, type MCPServer } from "@/components/mcp_tools/types"; import { Logo } from "@/components/molecules/logo/Logo"; -import { getMaskedAndFullUrl } from "./utils"; +import { getMaskedAndFullUrl, getMCPNetworkAccess } from "./utils"; interface MCPServerCardProps { server: MCPServer; @@ -70,7 +70,7 @@ const MCPServerCard: FC = ({ server.auth_type === AUTH_TYPE.OAUTH2 && !server.oauth2_flow && !server.delegate_auth_to_upstream; const status = server.status || "unknown"; const healthTone = HEALTH_TONE[status] ?? HEALTH_TONE.unknown; - const isPublic = server.available_on_public_internet; + const networkAccess = getMCPNetworkAccess(server); const accessGroups = (server.mcp_access_groups ?? []).filter((g): g is string => typeof g === "string"); const missing = missingUserFields ?? []; @@ -236,10 +236,17 @@ const MCPServerCard: FC = ({ )} - - - {isPublic ? "Public" : "Internal"} - + + + + {networkAccess.label} + + } + /> + {networkAccess.description} + {accessGroups.slice(0, 2).map((g) => ( { }); it("shows the read-only settings summary before editing", async () => { - renderView({ allow_all_keys: true, available_on_public_internet: false }); + renderView({ + allow_all_keys: true, + available_on_public_internet: false, + mcp_info: { server_name: "demo server", is_public: true, is_public_explicit: true }, + }); await userEvent.click(screen.getByRole("tab", { name: "Settings" })); expect(await screen.findByText("MCP Server Settings")).toBeInTheDocument(); expect(screen.getByText("Allow All Keys")).toBeInTheDocument(); expect(screen.getByText("Enabled")).toBeInTheDocument(); - expect(screen.getByText("Internal only")).toBeInTheDocument(); + expect(screen.getByText("Network access")).toBeInTheDocument(); + expect(screen.getByText("All Networks")).toBeInTheDocument(); + expect(screen.queryByText("MCP Hub")).not.toBeInTheDocument(); + expect(screen.queryByText("Listed")).not.toBeInTheDocument(); expect(screen.queryByText("edit form")).not.toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx index a7ff34301a0..c97596ce0f6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx @@ -13,7 +13,7 @@ import { MCPServerUserCredentialsPanel } from "./MCPServerUserCredentialsPanel"; import { getSecureItem } from "@/utils/secureStorage"; import { isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles"; import MCPServerCostDisplay from "./mcp_server_cost_display"; -import { getMaskedAndFullUrl } from "./utils"; +import { getMaskedAndFullUrl, getMCPNetworkAccess } from "./utils"; import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils"; import { CheckIcon, CopyIcon } from "lucide-react"; @@ -68,6 +68,7 @@ export const MCPServerView: React.FC = ({ const returningFromEditOAuth = isReturningFromEditOAuth(canEdit, mcpServer.server_id); const [editing, setEditing] = useState(isEditing || returningFromEditOAuth); const [showFullUrl, setShowFullUrl] = useState(false); + const networkAccess = getMCPNetworkAccess(mcpServer); const [copiedStates, setCopiedStates] = useState>({}); const [selectedTabIndex, setSelectedTabIndex] = useState(returningFromEditOAuth ? 2 : initialTabIndex); const canViewUserCredentials = userRole !== null && isProxyAdminTierRole(userRole); @@ -318,19 +319,13 @@ export const MCPServerView: React.FC = ({
-

Network Access

+

Network access

- {mcpServer.available_on_public_internet ? ( - - - Public - - ) : ( - - - Internal only - - )} + + + {networkAccess.label} + +

{networkAccess.description}

{handleAuth(mcpServer.auth_type) === "oauth2" && ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx index 3b4fda400c2..bf30821d73a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx @@ -3,11 +3,40 @@ import { extractMCPToken, maskUrl, getMaskedAndFullUrl, + getMCPNetworkAccess, validateMCPServerUrl, validateMCPServerName, normalizeToolOverrideMap, } from "./utils"; +describe("getMCPNetworkAccess", () => { + it.each([ + { publicIp: true, explicit: false, label: "All Networks" }, + { publicIp: false, explicit: true, label: "All Networks" }, + { publicIp: true, explicit: true, label: "All Networks" }, + { publicIp: false, explicit: false, label: "Internal Only" }, + { publicIp: true, explicit: undefined, label: "All Networks" }, + { publicIp: false, explicit: undefined, label: "Unknown" }, + { publicIp: undefined, explicit: false, label: "Unknown" }, + ])("reports $label for network=$publicIp and publication=$explicit", ({ publicIp, explicit, label }) => { + expect( + getMCPNetworkAccess({ + available_on_public_internet: publicIp, + mcp_info: { server_name: "demo", is_public: true, is_public_explicit: explicit }, + }).label, + ).toBe(label); + }); + + it("explains when hub publication permits public IPs", () => { + expect( + getMCPNetworkAccess({ + available_on_public_internet: false, + mcp_info: { server_name: "demo", is_public_explicit: true }, + }).description, + ).toContain("because this server is published in MCP Hub"); + }); +}); + describe("extractMCPToken", () => { it("should extract token after /mcp/", () => { const result = extractMCPToken("https://example.com/mcp/abc123"); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx index 4738e1e8fba..bb72831d92e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx @@ -1,4 +1,37 @@ -import { MCPEnvVar, MCPEnvVarScope } from "@/components/mcp_tools/types"; +import { MCPEnvVar, MCPEnvVarScope, type MCPServer } from "@/components/mcp_tools/types"; + +export const getMCPNetworkAccess = ( + server: Pick, +): { + readonly label: "All Networks" | "Internal Only" | "Unknown"; + readonly dotClassName: string; + readonly description: string; +} => { + const explicitlyPublished = server.mcp_info?.is_public_explicit; + if (server.available_on_public_internet === true || explicitlyPublished === true) { + return { + label: "All Networks", + dotClassName: "bg-success", + description: + server.available_on_public_internet === true + ? "Allows requests from public and internal IPs. Authentication and access permissions still apply" + : "Allows requests from public and internal IPs because this server is published in MCP Hub. Authentication and access permissions still apply", + }; + } + if (server.available_on_public_internet === false && explicitlyPublished === false) { + return { + label: "Internal Only", + dotClassName: "bg-warning", + description: + "Allows requests only from internal/private IP ranges. Authentication and access permissions still apply", + }; + } + return { + label: "Unknown", + dotClassName: "bg-border", + description: "The proxy did not report enough network and publication settings to determine allowed client IPs", + }; +}; export const extractMCPToken = (url: string): { token: string | null; baseUrl: string } => { try { diff --git a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx index 1032b03a3ce..fee58e10fda 100644 --- a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx @@ -1,4 +1,4 @@ -import { render, screen } from "@testing-library/react"; +import { render, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { describe, expect, it, vi } from "vitest"; import { DataTable } from "@/components/shared/DataTable"; @@ -63,6 +63,23 @@ describe("getMCPHubTableColumns", () => { expect(screen.getByText("Auth Type")).toBeInTheDocument(); }); + it("shows hub membership separately from the network setting", () => { + renderTable(vi.fn(), [ + { ...mockServer, available_on_public_internet: false, mcp_info: { is_public: true } }, + { + ...mockServer, + server_id: "network-only", + server_name: "Network-only server", + available_on_public_internet: true, + mcp_info: { is_public: false }, + }, + ]); + + expect(screen.getByText("Hub listing")).toBeInTheDocument(); + expect(within(screen.getByRole("row", { name: /exa_test/ })).getByText("Listed")).toBeInTheDocument(); + expect(within(screen.getByRole("row", { name: /Network-only server/ })).getByText("Unlisted")).toBeInTheDocument(); + }); + it("does not expose a URL column", () => { renderTable(); expect(screen.queryByText("URL")).not.toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx index 20a14bcb476..db53e97569b 100644 --- a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx @@ -203,8 +203,8 @@ export const getMCPHubTableColumns = ({ onServerClick }: MCPHubTableColumnsDeps) { id: "is_public", accessorFn: (row) => row.mcp_info?.is_public === true, - meta: { title: "Public", skeleton: "badge", className: "hidden md:table-cell" }, - header: ({ column }) => , + meta: { title: "Hub listing", skeleton: "badge", className: "hidden md:table-cell" }, + header: ({ column }) => , size: 100, enableSorting: true, sortingFn: (rowA, rowB) => { @@ -214,7 +214,7 @@ export const getMCPHubTableColumns = ({ onServerClick }: MCPHubTableColumnsDeps) }, cell: ({ row }) => { const isPublic = row.original.mcp_info?.is_public === true; - return ; + return ; }, }, { diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx index 27fe2330acd..1f052b8932a 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx @@ -1,5 +1,7 @@ import * as networking from "@/components/networking"; import userEvent from "@testing-library/user-event"; +import { act } from "@testing-library/react"; +import type { MCPServerData } from "./MCPHubTableColumns"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils"; import ModelHubTable from "./ModelHubTable"; @@ -18,6 +20,7 @@ vi.mock("@/components/networking", () => ({ getProxyBaseUrl: vi.fn(() => "http://localhost:4000"), getAgentsList: vi.fn(), fetchMCPServers: vi.fn(), + makeMCPPublicCall: vi.fn(), getUiSettings: vi.fn(), getClaudeCodePluginsList: vi.fn(() => Promise.resolve({ plugins: [] })), })); @@ -202,13 +205,13 @@ describe("ModelHubTable", () => { }); describe("hub tabs", () => { - const renderHub = async (agents: object[] = []) => { + const renderHub = async (agents: object[] = [], mcpServers: Promise = Promise.resolve([])) => { vi.mocked(networking.modelHubCall).mockResolvedValue({ data: [{ model_group: "claude-opus-4-8", providers: ["anthropic"], mode: "chat" }], }); vi.mocked(networking.getConfigFieldSetting).mockResolvedValue({ field_value: false }); vi.mocked(networking.getAgentsList).mockResolvedValue({ agents }); - vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); + vi.mocked(networking.fetchMCPServers).mockReturnValue(mcpServers); vi.mocked(networking.getUiSettings).mockResolvedValue({ values: {} }); mockUseUISettings.mockReturnValue({ data: { values: {} }, isLoading: false }); @@ -219,6 +222,29 @@ describe("ModelHubTable", () => { return { user, search: await screen.findByPlaceholderText("Search model names...") }; }; + it("requires a fresh MCP publication list before and after saving", async () => { + const servers = Promise.withResolvers(); + const { user } = await renderHub([], servers.promise); + await user.click(screen.getByRole("tab", { name: "MCP Hub" })); + + const manageVisibility = screen.getByRole("button", { name: "Manage MCP Hub Visibility" }); + expect(manageVisibility).toBeDisabled(); + await act(async () => servers.resolve([])); + expect(manageVisibility).toBeEnabled(); + + const refresh = Promise.withResolvers(); + vi.mocked(networking.makeMCPPublicCall).mockResolvedValueOnce({}); + vi.mocked(networking.fetchMCPServers).mockReturnValueOnce(refresh.promise); + await user.click(manageVisibility); + await user.click(screen.getByRole("button", { name: "Next" })); + await user.click(screen.getByRole("button", { name: "Save Publication List" })); + + expect(networking.makeMCPPublicCall).toHaveBeenCalledWith("test-token", []); + expect(manageVisibility).toBeDisabled(); + await act(async () => refresh.reject(new Error("Unable to reload the publication list"))); + expect(manageVisibility).toBeDisabled(); + }); + it("keeps the model filter typed on the Model Hub tab after visiting another hub", async () => { const { user, search } = await renderHub(); diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx index 063850b3e72..c4d776ea2fb 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx @@ -49,6 +49,10 @@ interface ModelHubTableProps { userRole: string | null; } +function isMCPHubVisibilityDisabled(isLoading: boolean, servers: readonly MCPServerData[] | null): boolean { + return isLoading || servers === null; +} + function HubEmptyState({ title, body }: { title: string; body: string }) { return (
@@ -359,10 +363,14 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, if (accessToken) { const fetchMcpData = async () => { try { + setMcpLoading(true); const response = await fetchMCPServers(accessToken); setMcpHubData(response); } catch (error) { + setMcpHubData(null); console.error("Error refreshing MCP server data:", error); + } finally { + setMcpLoading(false); } }; fetchMcpData(); @@ -567,7 +575,12 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, {/* Header with Make Public Button */} {publicPage == false && canModify && (
- +
)} diff --git a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx index 5b96e9ad194..881711c1668 100644 --- a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx @@ -1,6 +1,8 @@ import { render, screen, fireEvent, act, waitFor } from "@testing-library/react"; import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; import MakeMCPPublicForm from "./MakeMCPPublicForm"; +import userEvent from "@testing-library/user-event"; +import { toast } from "@/lib/toast"; import { MCPServerData } from "@/components/AIHub/MCPHubTableColumns"; // Mock the networking function @@ -8,6 +10,10 @@ vi.mock("../../networking", () => ({ makeMCPPublicCall: vi.fn(), })); +vi.mock("@/lib/toast", () => ({ + toast: { success: vi.fn(), fromError: vi.fn() }, +})); + // Import the mocked function import { makeMCPPublicCall } from "../../networking"; const mockMakeMCPPublicCall = vi.mocked(makeMCPPublicCall); @@ -28,7 +34,7 @@ describe("MakeMCPPublicForm", () => { url: "http://example.com/server1", transport: "http", status: "active", - mcp_info: { is_public: false }, + mcp_info: { is_public: false, is_public_explicit: false }, allowed_tools: ["tool-1", "tool-2"], auth_type: "bearer", credentials: {}, @@ -50,7 +56,7 @@ describe("MakeMCPPublicForm", () => { url: "http://example.com/server2", transport: "websocket", status: "inactive", - mcp_info: { is_public: true }, + mcp_info: { is_public: true, is_public_explicit: true }, allowed_tools: [], auth_type: "none", credentials: {}, @@ -80,16 +86,16 @@ describe("MakeMCPPublicForm", () => { it("should render the component", () => { render(); - expect(screen.getByText("Make MCP Servers Public")).toBeInTheDocument(); - expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument(); + expect(screen.getByText("Manage MCP Hub Visibility")).toBeInTheDocument(); + expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument(); }); it("should initialize with correct state", () => { render(); // Check that the component renders with the correct title and content - expect(screen.getByText("Make MCP Servers Public")).toBeInTheDocument(); - expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument(); + expect(screen.getByText("Manage MCP Hub Visibility")).toBeInTheDocument(); + expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument(); // Check that all server checkboxes are present const checkboxes = screen.getAllByRole("checkbox"); @@ -104,7 +110,7 @@ describe("MakeMCPPublicForm", () => { render(); // Initially on step 1 - expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument(); + expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument(); // Select all servers using the select all checkbox const selectAllCheckbox = screen.getByRole("checkbox", { name: "Select All (2)" }); @@ -123,7 +129,7 @@ describe("MakeMCPPublicForm", () => { // Should move to step 2 await waitFor(() => { - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); }); }); @@ -145,10 +151,10 @@ describe("MakeMCPPublicForm", () => { // Wait for navigation to complete await waitFor(() => { - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); }); - const submitButton = screen.getByRole("button", { name: "Make Public" }); + const submitButton = screen.getByRole("button", { name: "Save Publication List" }); await act(async () => { fireEvent.click(submitButton); }); @@ -187,29 +193,105 @@ describe("MakeMCPPublicForm", () => { expect(checkboxes[2]).not.toBeChecked(); }); - it("should show error when no servers selected", async () => { + it("submits an empty publication list after the last server is deselected", async () => { + mockMakeMCPPublicCall.mockResolvedValueOnce({}); render(); - // Deselect all servers first - const checkboxes = screen.getAllByRole("checkbox"); - await act(async () => { - fireEvent.click(checkboxes[0]); // Click select all to select all - }); - await act(async () => { - fireEvent.click(checkboxes[0]); // Click select all again to deselect all - }); + fireEvent.click(screen.getAllByRole("checkbox")[2]); + expect(screen.getByRole("button", { name: "Next" })).toBeEnabled(); + fireEvent.click(screen.getByRole("button", { name: "Next" })); + fireEvent.click(screen.getByRole("button", { name: "Save Publication List" })); - // Try to go to next step - const nextButton = screen.getByRole("button", { name: "Next" }); - await act(async () => { - fireEvent.click(nextButton); - }); - - // Should stay on same step - expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument(); + await waitFor(() => expect(mockMakeMCPPublicCall).toHaveBeenCalledWith("test-token", [])); + expect(mockProps.onSuccess).toHaveBeenCalled(); }); - it("should display empty state when no servers are available", () => { + it("keeps legacy listings separate from explicitly published selections", () => { + render( + , + ); + + expect(screen.getAllByRole("checkbox")[1]).not.toBeChecked(); + expect(screen.getAllByRole("checkbox")[2]).toBeChecked(); + expect(screen.getByText("Listed by legacy mode")).toBeInTheDocument(); + }); + + it.each([ + { mode: "all missing, stale true", info: { is_public: true }, mixed: false }, + { mode: "all missing, stale false", info: { is_public: false }, mixed: false }, + { mode: "mixed, stale true", info: { is_public: true }, mixed: true }, + { mode: "mixed, stale false", info: { is_public: false }, mixed: true }, + { mode: "null explicit status", info: { is_public: true, is_public_explicit: null }, mixed: true }, + { mode: "nonboolean explicit status", info: { is_public: true, is_public_explicit: "true" }, mixed: true }, + ])("blocks unknown explicit publication metadata: $mode", ({ info, mixed }) => { + const unknownServer = { ...mockProps.mcpHubData[0], mcp_info: info }; + const catalog = mixed ? [unknownServer, mockProps.mcpHubData[1]] : [unknownServer]; + render(); + + expect(screen.getByRole("alert")).toHaveTextContent("explicit publication status"); + expect(screen.queryByRole("checkbox")).not.toBeInTheDocument(); + expect(screen.queryByText("Configure in YAML")).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Copy code" })).not.toBeInTheDocument(); + const nextButton = screen.getByRole("button", { name: "Next" }); + expect(nextButton).toBeDisabled(); + fireEvent.click(nextButton); + expect(screen.queryByText("Confirm MCP Hub Publication")).not.toBeInTheDocument(); + expect(mockMakeMCPPublicCall).not.toHaveBeenCalled(); + }); + + it.each([true, false])("blocks confirmation when explicit metadata disappears with stale listing %s", (listed) => { + const { rerender } = render(); + fireEvent.click(screen.getByRole("button", { name: "Next" })); + expect(screen.getByRole("button", { name: "Save Publication List" })).toBeEnabled(); + + const catalog = [{ ...mockProps.mcpHubData[0], mcp_info: { is_public: listed } }, mockProps.mcpHubData[1]]; + rerender(); + + expect(screen.getByRole("alert")).toHaveTextContent("explicit publication status"); + expect(screen.queryByText("Confirm MCP Hub Publication")).not.toBeInTheDocument(); + const saveButton = screen.getByRole("button", { name: "Save Publication List" }); + expect(saveButton).toBeDisabled(); + fireEvent.click(saveButton); + expect(mockMakeMCPPublicCall).not.toHaveBeenCalled(); + expect(screen.queryByRole("button", { name: "Copy code" })).not.toBeInTheDocument(); + + const refreshedCatalog = [ + { ...mockProps.mcpHubData[0], mcp_info: { is_public: true, is_public_explicit: true } }, + { ...mockProps.mcpHubData[1], mcp_info: { is_public: false, is_public_explicit: false } }, + ]; + rerender(); + expect(screen.queryByRole("alert")).not.toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Next" })).toBeEnabled(); + expect(screen.getByRole("checkbox", { name: "Publish Test Server 1" })).toBeChecked(); + expect(screen.getByRole("checkbox", { name: "Publish Test Server 2" })).not.toBeChecked(); + }); + + it("copies publication YAML using the selected server IDs", async () => { + const user = userEvent.setup(); + render(); + + await user.click(screen.getByText("Configure in YAML")); + await user.click(screen.getByRole("button", { name: "Copy code" })); + + expect(await navigator.clipboard.readText()).toBe( + 'litellm_settings:\n public_mcp_hub_strict_whitelist: true\n public_mcp_servers:\n - "server-2"', + ); + + await user.click(screen.getByRole("checkbox", { name: "Publish Test Server 2" })); + await user.click(screen.getByRole("button", { name: "Copy code" })); + expect(await navigator.clipboard.readText()).toBe( + "litellm_settings:\n public_mcp_hub_strict_whitelist: true\n public_mcp_servers: []", + ); + }); + + it("allows clearing publication IDs when the loaded server catalog is empty", async () => { + mockMakeMCPPublicCall.mockResolvedValueOnce({}); const emptyProps = { ...mockProps, mcpHubData: [] as MCPServerData[], @@ -223,9 +305,13 @@ describe("MakeMCPPublicForm", () => { const selectAllCheckbox = screen.getByRole("checkbox", { name: "Select All" }); expectDisabledControl(selectAllCheckbox); - // Next button should be disabled const nextButton = screen.getByRole("button", { name: "Next" }); - expect(nextButton).toBeDisabled(); + expect(nextButton).toBeEnabled(); + fireEvent.click(nextButton); + fireEvent.click(screen.getByRole("button", { name: "Save Publication List" })); + + await waitFor(() => expect(mockMakeMCPPublicCall).toHaveBeenCalledWith("test-token", [])); + expect(mockProps.onSuccess).toHaveBeenCalled(); }); it("should handle Cancel button functionality", async () => { @@ -252,7 +338,7 @@ describe("MakeMCPPublicForm", () => { // Verify we're on step 1 await waitFor(() => { - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); }); // Click Previous button @@ -262,7 +348,7 @@ describe("MakeMCPPublicForm", () => { }); // Should go back to step 0 - expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument(); + expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument(); }); it("should handle individual server selection", async () => { @@ -322,8 +408,8 @@ describe("MakeMCPPublicForm", () => { }); it("should handle submit error properly", async () => { - const errorMessage = "Network error"; - mockMakeMCPPublicCall.mockRejectedValueOnce(new Error(errorMessage)); + const error = new Error("Update litellm_settings.public_mcp_servers in your YAML configuration"); + mockMakeMCPPublicCall.mockRejectedValueOnce(error); render(); @@ -333,10 +419,10 @@ describe("MakeMCPPublicForm", () => { }); await waitFor(() => { - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); }); - const submitButton = screen.getByRole("button", { name: "Make Public" }); + const submitButton = screen.getByRole("button", { name: "Save Publication List" }); await act(async () => { fireEvent.click(submitButton); }); @@ -346,6 +432,8 @@ describe("MakeMCPPublicForm", () => { expect(mockMakeMCPPublicCall).toHaveBeenCalledWith("test-token", ["server-2"]); }); + expect(toast.fromError).toHaveBeenCalledWith(error); + // Should not call onSuccess or onClose on error expect(mockProps.onSuccess).not.toHaveBeenCalled(); expect(mockProps.onClose).not.toHaveBeenCalled(); @@ -366,10 +454,10 @@ describe("MakeMCPPublicForm", () => { }); await waitFor(() => { - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); }); - const submitButton = screen.getByRole("button", { name: "Make Public" }); + const submitButton = screen.getByRole("button", { name: "Save Publication List" }); await act(async () => { fireEvent.click(submitButton); }); @@ -381,7 +469,7 @@ describe("MakeMCPPublicForm", () => { expect(mockMakeMCPPublicCall).toHaveBeenCalledTimes(1); expect(mockProps.onSuccess).not.toHaveBeenCalled(); expect(mockProps.onClose).not.toHaveBeenCalled(); - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); resolvePromise({}); await waitFor(() => { @@ -400,7 +488,7 @@ describe("MakeMCPPublicForm", () => { // Modal should not be rendered expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); - expect(screen.queryByText("Make MCP Servers Public")).not.toBeInTheDocument(); + expect(screen.queryByText("Manage MCP Hub Visibility")).not.toBeInTheDocument(); }); it("should preselect already public servers when modal opens", () => { @@ -415,7 +503,7 @@ describe("MakeMCPPublicForm", () => { url: "http://example.com/server1", transport: "http", status: "active", - mcp_info: { is_public: false }, // Not public + mcp_info: { is_public: false, is_public_explicit: false }, // Not public allowed_tools: [], auth_type: "bearer", credentials: {}, @@ -437,7 +525,7 @@ describe("MakeMCPPublicForm", () => { url: "http://example.com/server2", transport: "websocket", status: "inactive", - mcp_info: { is_public: true }, // Already public + mcp_info: { is_public: true, is_public_explicit: true }, // Already public allowed_tools: [], auth_type: "none", credentials: {}, @@ -459,7 +547,7 @@ describe("MakeMCPPublicForm", () => { url: "http://example.com/server3", transport: "sse", status: "healthy", - mcp_info: { is_public: true }, // Already public + mcp_info: { is_public: true, is_public_explicit: true }, // Already public allowed_tools: [], auth_type: "oauth", credentials: {}, diff --git a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx index 8287cf47f1a..2448732a236 100644 --- a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx @@ -1,5 +1,6 @@ import React, { useState, useEffect } from "react"; import { Loader2 } from "lucide-react"; +import CodeBlock from "@/components/CodeBlock"; import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Checkbox } from "@/components/ui/checkbox"; @@ -29,6 +30,11 @@ interface MakeMCPPublicFormProps { onSuccess: () => void; } +interface PublicationSelection { + readonly catalog: MCPServerData[]; + readonly serverIds: Set; +} + const MakeMCPPublicForm: React.FC = ({ visible, onClose, @@ -37,21 +43,28 @@ const MakeMCPPublicForm: React.FC = ({ onSuccess, }) => { const [currentStep, setCurrentStep] = useState(0); - const [selectedServers, setSelectedServers] = useState>(new Set()); + const [selection, setSelection] = useState(null); const [loading, setLoading] = useState(false); + const selectedServers = selection?.serverIds ?? new Set(); + const hasPublicationMetadata = mcpHubData.every((server) => typeof server.mcp_info?.is_public_explicit === "boolean"); + const canManagePublication = hasPublicationMetadata && selection?.catalog === mcpHubData; + const publicationYaml = [ + "litellm_settings:", + " public_mcp_hub_strict_whitelist: true", + selectedServers.size === 0 + ? " public_mcp_servers: []" + : ` public_mcp_servers:\n${Array.from(selectedServers, (id) => ` - ${JSON.stringify(id)}`).join("\n")}`, + ].join("\n"); const handleClose = () => { setCurrentStep(0); - setSelectedServers(new Set()); + setSelection(null); onClose(); }; const handleNext = () => { + if (!canManagePublication) return; if (currentStep === 0) { - if (selectedServers.size === 0) { - toast.fromError("Please select at least one MCP server to make public"); - return; - } setCurrentStep(1); } }; @@ -69,37 +82,32 @@ const MakeMCPPublicForm: React.FC = ({ } else { newSelection.delete(serverId); } - setSelectedServers(newSelection); + setSelection({ catalog: mcpHubData, serverIds: newSelection }); }; const handleSelectAll = (checked: boolean) => { if (checked) { const allServerIds = mcpHubData.map((server) => server.server_id); - setSelectedServers(new Set(allServerIds)); + setSelection({ catalog: mcpHubData, serverIds: new Set(allServerIds) }); } else { - setSelectedServers(new Set()); + setSelection({ catalog: mcpHubData, serverIds: new Set() }); } }; - // Initialize and preselect already public servers when modal opens useEffect(() => { - if (visible && mcpHubData.length > 0) { - // Extract server IDs from servers that are already public - const publicServerIds = mcpHubData - .filter((server) => server.mcp_info?.is_public === true) - .map((server) => server.server_id); - - // Preselect servers that are already public - setSelectedServers(new Set(publicServerIds)); - } - }, [visible]); // Only re-run when modal visibility changes, not when mcpHubData updates - - const handleSubmit = async () => { - if (selectedServers.size === 0) { - toast.fromError("Please select at least one MCP server to make public"); + if (!visible || !hasPublicationMetadata) { + setSelection(null); return; } + const publicServerIds = mcpHubData + .filter((server) => server.mcp_info.is_public_explicit === true) + .map((server) => server.server_id); + setSelection({ catalog: mcpHubData, serverIds: new Set(publicServerIds) }); + setCurrentStep(0); + }, [visible, mcpHubData, hasPublicationMetadata]); + const handleSubmit = async () => { + if (!canManagePublication) return; setLoading(true); try { const serverIdsToMakePublic = Array.from(selectedServers); @@ -107,12 +115,12 @@ const MakeMCPPublicForm: React.FC = ({ // Make batch API call for all servers await makeMCPPublicCall(accessToken, serverIdsToMakePublic); - toast.success(`Successfully made ${serverIdsToMakePublic.length} MCP server(s) public!`); + toast.success("MCP Hub publication list updated"); handleClose(); onSuccess(); } catch (error) { console.error("Error making MCP servers public:", error); - toast.fromError("Failed to make MCP servers public. Please try again."); + toast.fromError(error); } finally { setLoading(false); } @@ -126,7 +134,7 @@ const MakeMCPPublicForm: React.FC = ({ return (
-

Select MCP Servers to Make Public

+

Select MCP Servers for the Hub

- Select the MCP servers you want to be visible on the public model hub. Users will still require a valid - Virtual Key to use these servers. + Select the complete list of MCP servers to publish on the public hub. Uncheck a server to remove it from this + list, or uncheck all to clear it. Authentication and access permissions still apply +

+ +

+ Legacy mode also lists servers with public IP access enabled. Set public_mcp_hub_strict_whitelist to true in + your configuration to use only the publication list

@@ -160,16 +173,22 @@ const MakeMCPPublicForm: React.FC = ({ className="flex items-center space-x-3 p-3 border rounded-lg hover:bg-accent" > handleServerSelection(server.server_id, checked === true)} />

{server.server_name}

- {isPublic && Public} + {isPublic && ( + + {server.mcp_info?.is_public_explicit === false ? "Listed by legacy mode" : "Listed"} + + )} {server.transport} {server.status || "unknown"}
+

{server.server_id}

{server.description || server.url}

@@ -193,6 +212,18 @@ const MakeMCPPublicForm: React.FC = ({
+
+ Configure in YAML +
+

+ Merge these settings into your proxy configuration and reload it. Entries use the server IDs shown above, + not names or aliases. For servers defined in YAML, pin server_id in each existing mcp_servers entry so the + publication list stays stable +

+ +
+
+ {selectedServers.size > 0 && (

@@ -207,19 +238,20 @@ const MakeMCPPublicForm: React.FC = ({ const renderStep2Content = () => { return (

-

Confirm Making MCP Servers Public

+

Confirm MCP Hub Publication

- Warning: Once you make these MCP servers public, anyone who can go to the{" "} - /ui/model_hub_table will be able to know they exist on the proxy. + Anyone who can open /ui/model_hub_table can discover published servers. Explicitly published + server IDs also allow requests from public IPs. Authentication and access permissions still apply

-

MCP Servers to be made public:

+

MCP servers in the publication list:

+ {selectedServers.size === 0 &&

No explicitly published servers

} {Array.from(selectedServers).map((serverId) => { const server = mcpHubData.find((s) => s.server_id === serverId); return ( @@ -248,8 +280,8 @@ const MakeMCPPublicForm: React.FC = ({

- Total: {selectedServers.size} MCP server{selectedServers.size !== 1 ? "s" : ""} will be - made public + Saving replaces the publication list with {selectedServers.size} MCP server + {selectedServers.size !== 1 ? "s" : ""}. Legacy mode may still list servers with public IP access enabled

@@ -257,6 +289,15 @@ const MakeMCPPublicForm: React.FC = ({ }; const renderStepContent = () => { + if (!hasPublicationMetadata) { + return ( +
+ This proxy does not provide explicit publication status for every MCP server. Update the proxy to manage + visibility here, or edit litellm_settings.public_mcp_servers in its existing configuration +
+ ); + } + if (!canManagePublication) return

Loading publication settings

; switch (currentStep) { case 0: return renderStep1Content(); @@ -276,15 +317,15 @@ const MakeMCPPublicForm: React.FC = ({
{currentStep === 0 && ( - )} {currentStep === 1 && ( - )}
@@ -296,7 +337,7 @@ const MakeMCPPublicForm: React.FC = ({ !open && handleClose()} disablePointerDismissal> - Make MCP Servers Public + Manage MCP Hub Visibility
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index be7d39616ca..afeff869b7a 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -321,6 +321,8 @@ export interface MCPServerCostInfo { // Define MCP provider info export interface MCPInfo { server_name: string; + is_public?: boolean; + is_public_explicit?: boolean; description?: string; logo_url?: string; mcp_server_cost_info?: MCPServerCostInfo | null; From 4c2458a0b1738472ad66dec20207a9673245d8b4 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 17:18:23 -0700 Subject: [PATCH 02/65] chore(cost-map): sync openrouter prices for deepseek, minimax, qwen and glm rows (#43384) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 45 ++++++++++--------- model_prices_and_context_window.json | 45 ++++++++++--------- 2 files changed, 46 insertions(+), 44 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 45f5967d372..d807dca329a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -41956,15 +41956,15 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "cache_read_input_token_cost": 8.8e-09, - "input_cost_per_token": 2.64e-07, + "cache_read_input_token_cost": 2.475e-07, + "input_cost_per_token": 2.476e-07, "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 384000, - "max_tokens": 384000, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 7.92e-07, + "output_cost_per_token": 3.5e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42318,13 +42318,13 @@ "max_output_tokens": 8000 }, "openrouter/minimax/minimax-m2": { - "input_cost_per_token": 3e-07, + "input_cost_per_token": 2.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.2e-06, + "output_cost_per_token": 1.02e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -66867,13 +66867,13 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "input_cost_per_token": 3.794e-07, - "output_cost_per_token": 1.1924e-06, - "cache_read_input_token_cost": 7.046e-08, + "input_cost_per_token": 2.38e-07, + "output_cost_per_token": 7.48e-07, + "cache_read_input_token_cost": 3.91e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -67005,7 +67005,7 @@ }, "openrouter/deepseek/deepseek-v4-flash-0731": { "cache_read_input_token_cost": 1.6e-08, - "input_cost_per_token": 2.2e-08, + "input_cost_per_token": 2.1e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, @@ -67719,13 +67719,13 @@ "supports_web_search": false }, "openrouter/minimax/minimax-m2.7": { - "input_cost_per_token": 3e-07, - "output_cost_per_token": 1.2e-06, - "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 2.1e-07, + "output_cost_per_token": 8.4e-07, + "cache_read_input_token_cost": 4.2e-08, "litellm_provider": "openrouter", "max_input_tokens": 204800, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 176947, + "max_tokens": 176947, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -68841,12 +68841,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-30b-a3b": { - "input_cost_per_token": 1.2e-07, - "output_cost_per_token": 5e-07, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 1.3e-07, + "output_cost_per_token": 5.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, - "max_output_tokens": 16384, - "max_tokens": 16384, + "max_output_tokens": 8192, + "max_tokens": 8192, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 45f5967d372..d807dca329a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -41956,15 +41956,15 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "cache_read_input_token_cost": 8.8e-09, - "input_cost_per_token": 2.64e-07, + "cache_read_input_token_cost": 2.475e-07, + "input_cost_per_token": 2.476e-07, "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 384000, - "max_tokens": 384000, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 7.92e-07, + "output_cost_per_token": 3.5e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42318,13 +42318,13 @@ "max_output_tokens": 8000 }, "openrouter/minimax/minimax-m2": { - "input_cost_per_token": 3e-07, + "input_cost_per_token": 2.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.2e-06, + "output_cost_per_token": 1.02e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -66867,13 +66867,13 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "input_cost_per_token": 3.794e-07, - "output_cost_per_token": 1.1924e-06, - "cache_read_input_token_cost": 7.046e-08, + "input_cost_per_token": 2.38e-07, + "output_cost_per_token": 7.48e-07, + "cache_read_input_token_cost": 3.91e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -67005,7 +67005,7 @@ }, "openrouter/deepseek/deepseek-v4-flash-0731": { "cache_read_input_token_cost": 1.6e-08, - "input_cost_per_token": 2.2e-08, + "input_cost_per_token": 2.1e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, @@ -67719,13 +67719,13 @@ "supports_web_search": false }, "openrouter/minimax/minimax-m2.7": { - "input_cost_per_token": 3e-07, - "output_cost_per_token": 1.2e-06, - "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 2.1e-07, + "output_cost_per_token": 8.4e-07, + "cache_read_input_token_cost": 4.2e-08, "litellm_provider": "openrouter", "max_input_tokens": 204800, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 176947, + "max_tokens": 176947, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -68841,12 +68841,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-30b-a3b": { - "input_cost_per_token": 1.2e-07, - "output_cost_per_token": 5e-07, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 1.3e-07, + "output_cost_per_token": 5.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, - "max_output_tokens": 16384, - "max_tokens": 16384, + "max_output_tokens": 8192, + "max_tokens": 8192, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, From e73abe6c72785ad91d4927da26de3a5d1b54300b Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 17:36:34 -0700 Subject: [PATCH 03/65] chore(cost-map): drop stale cache hit field from openrouter deepseek-v4-pro-0813 (#43389) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 1 - model_prices_and_context_window.json | 1 - 2 files changed, 2 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index d807dca329a..09fc442e5a7 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -41958,7 +41958,6 @@ "openrouter/deepseek/deepseek-v4-pro-0813": { "cache_read_input_token_cost": 2.475e-07, "input_cost_per_token": 2.476e-07, - "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index d807dca329a..09fc442e5a7 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -41958,7 +41958,6 @@ "openrouter/deepseek/deepseek-v4-pro-0813": { "cache_read_input_token_cost": 2.475e-07, "input_cost_per_token": 2.476e-07, - "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, From eea1d0f2696d6ab6b67e8b208c85bae9fa624e1e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 18:05:21 -0700 Subject: [PATCH 04/65] fix(responses): stream guardrail pre-call block as SSE with a typed output item (#42507) * fix(responses): stream guardrail pre-call block as SSE with a typed output item Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): import blocked usage helper from the guardrail utils module Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(e2e): drop narrating docstrings and poll without rebinding Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): cover pre-call guardrail block on /v1/responses stream and json Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): audit cells for responses guardrail block contract Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): observe upstream on the recorded chat route for responses denial cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): tidy responses denial audit cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): wait for worker count to recover after SIGKILL Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): require a replacement worker after SIGKILL Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): type the blocked response test helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: yucheng --- .../guardrail_translation/handler.py | 6 +- .../proxy/response_api_endpoints/endpoints.py | 32 +- tests/e2e/coverage_registry/guardrail.yaml | 1 + tests/e2e/guardrails/guardrails_client.py | 29 + ...est_responses_pre_call_block_stream_e2e.py | 154 ++++ .../observability/test_guardrail_effects.py | 695 +++++++++++++++++- .../response_api_endpoints/test_endpoints.py | 148 +++- .../proxy/test_blocked_response_usage.py | 20 +- 8 files changed, 1022 insertions(+), 63 deletions(-) create mode 100644 tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 66cebe0175d..d6d68e0607a 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -1540,7 +1540,7 @@ class OpenAIResponsesHandler(BaseTranslation): from litellm.responses.streaming_iterator import build_synthetic_response_events return build_synthetic_response_events( - transformed=_blocked_response(exc, response_id=f"resp_{uuid.uuid4()}", model=exc.model), + transformed=build_blocked_response(exc), logging_obj=None, chunk_size=max(len(exc.message), 1), ) @@ -1648,6 +1648,10 @@ def _blocked_output_item(exc: "ModifyResponseException") -> GenericResponseOutpu return GenericResponseOutputItem.model_validate(payload) +def build_blocked_response(exc: "ModifyResponseException") -> ResponsesAPIResponse: + return _blocked_response(exc, response_id=f"resp_{uuid.uuid4()}", model=exc.model) + + def _blocked_response( exc: "ModifyResponseException", response_id: str, diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 75eefb2e73b..c5d702ad65a 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -1,13 +1,11 @@ import asyncio import contextlib import json -import time -from collections.abc import AsyncIterator, Awaitable, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Mapping, Sequence from enum import Enum from functools import partial from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, TypeAlias, cast, get_args -from uuid import uuid4 import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, Response @@ -21,8 +19,9 @@ from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger from litellm.constants import EMPTY_MAPPING from litellm.integrations.custom_guardrail import ModifyResponseException -from litellm.llms.base_llm.guardrail_translation.utils import ( - blocked_responses_api_usage as _blocked_responses_api_usage, +from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + build_blocked_response, ) from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import ( @@ -30,7 +29,7 @@ from litellm.proxy.auth.user_api_key_auth import ( user_api_key_auth, user_api_key_auth_websocket, ) -from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing, create_response from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_set_request_parsed_body, @@ -440,17 +439,16 @@ async def responses_api( request_data=_data, ) - violation_text: Final = e.message - response_obj: Final = ResponsesAPIResponse( - id=f"resp_{uuid4()}", - object="response", - created_at=int(time.time()), - model=e.model or data.get("model"), - output=cast(Any, [{"content": [{"type": "text", "text": violation_text}]}]), - status="completed", - usage=_blocked_responses_api_usage(e.original_response), - ) - return response_obj + if data.get("stream") is True: + block_chunks: Final = OpenAIResponsesHandler().build_block_sse_chunks(e) + + async def _blocked_stream() -> AsyncGenerator[str, None]: + for chunk in block_chunks: + yield chunk.decode() + yield "data: [DONE]\n\n" + + return await create_response(generator=_blocked_stream(), media_type="text/event-stream", headers={}) + return build_blocked_response(e) except Exception as e: raise await processor._handle_llm_api_exception( e=e, diff --git a/tests/e2e/coverage_registry/guardrail.yaml b/tests/e2e/coverage_registry/guardrail.yaml index f49568c883b..920a288aea6 100644 --- a/tests/e2e/coverage_registry/guardrail.yaml +++ b/tests/e2e/coverage_registry/guardrail.yaml @@ -37,4 +37,5 @@ - {id: guardrail.mcp_security.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [mcp_operations], source: "guardrail_hooks/mcp_security", rationale: "MCP protocol security"} - {id: guardrail.llm_as_a_judge.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/llm_as_a_judge", rationale: "LLM-based judgment guardrail"} - {id: guardrail.litellm_content_filter.pre_mcp_call.blocks, module: guardrail, tier: P1, hook_point: pre_mcp_call, assertions: [blocks], exercised_on: [mcp_operations], source: "guardrail_hooks/litellm_content_filter/content_filter.py:_scan_mcp_tool_call_arguments", rationale: "A general content-filter guardrail configured mode=pre_mcp_call blocks a banned keyword in an MCP tool call's arguments before it reaches the upstream MCP server; a clean argument passes"} +- {id: guardrail.custom_code.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [responses], source: "response_api_endpoints/endpoints.py ModifyResponseException handler", rationale: "A pre_call custom_code block on /v1/responses must answer in the requested shape: SSE response.completed with a completed assistant output_text message item when stream=true, schema-valid JSON when not, both with zero usage"} - {id: guardrail.dispatch.pre_call.rejects_unknown_name, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "proxy guardrail dispatch (per-request `guardrails` selector)", rationale: "A request naming a guardrail this proxy does not serve must fail closed with a 4xx; today it is silently served unguarded, so a typo'd name drops the protection the caller asked for"} diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py index 17223dc36fa..1f4fc43355b 100644 --- a/tests/e2e/guardrails/guardrails_client.py +++ b/tests/e2e/guardrails/guardrails_client.py @@ -111,6 +111,15 @@ class ToolPermissionParamsBody(GuardrailParamsBase): on_disallowed_action: Literal["block", "rewrite"] = "block" +class CustomCodeParamsBody(GuardrailParamsBase): + """Custom-code guardrail params: `custom_code` is the sandboxed source the + proxy compiles, which must define `apply_guardrail(inputs, request_data, + input_type)` returning `allow()` or `block(reason)`.""" + + guardrail: Literal["custom_code"] = "custom_code" + custom_code: str + + GuardrailParamsBody = ( ContentFilterParamsBody | BedrockGuardrailParamsBody @@ -118,6 +127,7 @@ GuardrailParamsBody = ( | BlockCodeExecutionParamsBody | PresidioParamsBody | ToolPermissionParamsBody + | CustomCodeParamsBody ) @@ -174,6 +184,7 @@ class _ResponsesGuardrailBody(BaseModel): model: str input: str guardrails: list[str] | None = None + stream: bool | None = None @dataclass(frozen=True, slots=True) @@ -509,6 +520,24 @@ class GuardrailsClient: json=_ResponsesGuardrailBody(model=model, input=text, guardrails=guardrails), ) + def responses_stream_raw( + self, + key: str, + model: str, + text: str, + *, + guardrails: list[str] | None = None, + ) -> StreamingResponse: + """Drive /v1/responses with stream=true, returning the raw HTTP outcome: + a streamed block is judged on status, content-type, and the SSE event + sequence, not a typed JSON body.""" + return self.proxy.transport.send( + "/v1/responses", + headers=self.proxy.transport.bearer(key), + json=_ResponsesGuardrailBody(model=model, input=text, guardrails=guardrails, stream=True), + stream=True, + ) + def apply_guardrail(self, key: str, *, name: str, text: str) -> Result[ApplyGuardrailResponse]: return self.proxy.transport.post( "/guardrails/apply_guardrail", diff --git a/tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py b/tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py new file mode 100644 index 00000000000..93512e2a64c --- /dev/null +++ b/tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py @@ -0,0 +1,154 @@ +from __future__ import annotations + +import time +from collections.abc import Callable +from typing import Final + +import pytest +from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker +from e2e_http import StreamingResponse +from guardrails_client import CustomCodeParamsBody, GuardrailsClient +from lifecycle import ResourceManager +from pydantic import BaseModel, TypeAdapter + +pytestmark = pytest.mark.e2e + +DENIAL: Final = "This model is not currently available. Please contact support if you think this is a mistake." + +CUSTOM_CODE: Final = f''' +def apply_guardrail(inputs, request_data, input_type): + return block("{DENIAL}") +''' + + +class _ContentPart(BaseModel): + type: str + text: str | None = None + + +class _OutputItem(BaseModel): + type: str | None = None + id: str | None = None + role: str | None = None + status: str | None = None + content: list[_ContentPart] = [] + + +class _Usage(BaseModel): + total_tokens: int = 0 + + +class _ResponseBody(BaseModel): + output: list[_OutputItem] = [] + usage: _Usage | None = None + + +class _EventHead(BaseModel): + type: str + + +class _CompletedEvent(BaseModel): + type: str + response: _ResponseBody + + +_EVENT_HEAD: Final = TypeAdapter(_EventHead) + + +def _denial_delivered(result: StreamingResponse) -> bool: + if not result.ok: + return False + if DENIAL in result.body: + return True + return any(DENIAL in event for event in result.stream_events) + + +def _poll_terminal(result: StreamingResponse) -> bool: + if _denial_delivered(result): + return True + if result.ok: + return False + return "Guardrail not found" not in result.body and result.status_code not in (-1, 401, 429) + + +def _poll_attempt(call: Callable[[], StreamingResponse], deadline: float) -> StreamingResponse: + result: Final = call() + if _poll_terminal(result) or time.monotonic() >= deadline: + return result + time.sleep(POLL_INTERVAL) + return _poll_attempt(call, deadline) + + +def _poll_for_block(call: Callable[[], StreamingResponse]) -> StreamingResponse: + return _poll_attempt(call, time.monotonic() + POLL_TIMEOUT) + + +def _assert_blocked_response(response: _ResponseBody) -> None: + item = next(iter(response.output), None) + assert item is not None, f"blocked response carried no output item: {response.output!r}" + assert item.type == "message", f"output[0] must be a message item, got {item.type!r}: {item!r}" + assert item.role == "assistant", f"output[0] role must be assistant, got {item.role!r}" + assert item.status == "completed", f"output[0] status must be completed, got {item.status!r}" + part = next(iter(item.content), None) + assert part is not None, f"output[0] carried no content part: {item!r}" + assert part.type == "output_text", f"content[0] must be output_text, got {part.type!r}" + assert part.text == DENIAL, f"content[0] text must be the denial, got {part.text!r}" + assert response.usage is not None and response.usage.total_tokens == 0, ( + f"a blocked response never reached a provider, usage must be zero: {response.usage!r}" + ) + + +class TestResponsesPreCallBlock: + def _register_block(self, client: GuardrailsClient, resources: ResourceManager) -> str: + name: Final = f"e2e-custom-code-responses-block-{unique_marker()}" + guardrail_id: Final = client.register( + name, + CustomCodeParamsBody(mode="pre_call", default_on=False, custom_code=CUSTOM_CODE), + ) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + return name + + @pytest.mark.covers("guardrail.custom_code.pre_call.blocks", exercised_on=["responses"]) + def test_stream_block_is_sse_with_completed_assistant_message( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + name: Final = self._register_block(client, resources) + model: Final = client.create_backend_model( + resources, prefix="e2e-responses-block", backend="openai/gpt-4.1-mini", api_key="os.environ/OPENAI_API_KEY" + ) + + result: Final = _poll_for_block( + lambda: client.responses_stream_raw(scoped_key, model, "say hi", guardrails=[name]) + ) + + assert result.status_code == 200, f"a pre_call block answers 200, got {result.status_code}: {result.body[:400]}" + assert (result.content_type or "").startswith("text/event-stream"), ( + f"stream=true must answer SSE, got content-type {result.content_type!r}: {result.body[:400]}" + ) + events: Final = tuple(_EVENT_HEAD.validate_json(payload).type for payload in result.stream_events) + completed: Final = tuple( + _CompletedEvent.model_validate_json(payload) + for payload, event_type in zip(result.stream_events, events) + if event_type == "response.completed" + ) + assert len(completed) == 1, ( + f"the denial stream must end in exactly one response.completed event, got events {events!r}" + ) + _assert_blocked_response(completed[0].response) + + @pytest.mark.covers("guardrail.custom_code.pre_call.blocks", exercised_on=["responses"]) + def test_non_stream_block_is_schema_valid_json( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + name: Final = self._register_block(client, resources) + model: Final = client.create_backend_model( + resources, prefix="e2e-responses-block", backend="openai/gpt-4.1-mini", api_key="os.environ/OPENAI_API_KEY" + ) + + result: Final = _poll_for_block(lambda: client.responses(scoped_key, model, "say hi", guardrails=[name])) + + assert result.status_code == 200, f"a pre_call block answers 200, got {result.status_code}: {result.body[:400]}" + assert (result.content_type or "").startswith("application/json"), ( + f"a non-streaming block answers JSON, got content-type {result.content_type!r}" + ) + _assert_blocked_response(_ResponseBody.model_validate_json(result.body)) diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index 4fac42a796d..9f5f3da4302 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -1,15 +1,22 @@ import json +import os +import signal +import socket import uuid +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from typing import Final +import httpx +import psutil import pytest import yaml from integration._support.client import Gateway, eventually, object_value from integration._support.database import read_rows from integration._support.mcp import mcp_peer, register_mcp, tool_names -from integration._support.process import owned_proxy +from integration._support.process import group_members, owned_proxy, owned_proxy_process from integration._support.wire import Reply, Request, wire_server +from openai import AsyncOpenAI, OpenAI @pytest.mark.covers("other.observability.guardrails.rewrite_reaches_correct_anthropic_positions") @@ -208,8 +215,6 @@ def test_guardrail_denial_prevents_provider_and_preserves_allowed_control(gatewa with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: model: Final = scenario.model() key: Final = scenario.key(models=[model]) - import httpx - with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: observed.get("/__observations") denied: Final = candidate.request( @@ -500,3 +505,687 @@ def test_request_selected_mcp_guardrail_blocks_direct_and_virtual_calls(gateway: assert len(calls) == 1 assert calls[0]["body"]["params"]["name"] == tool assert calls[0]["body"]["params"]["arguments"] == arguments + + +_RESPONSES_DENIAL: Final = "This model is not currently available." + + +def _deny_guardrail(name: str, denial: str = _RESPONSES_DENIAL) -> dict[str, object]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "custom_code", + "mode": "pre_call", + "default_on": False, + "custom_code": (f"def apply_guardrail(inputs, request_data, input_type):\n return block({denial!r})\n"), + }, + } + + +def _responses_denial_config(tmp_path: Path, identity: str, denial: str = _RESPONSES_DENIAL) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [_deny_guardrail(identity, denial)] + path: Final = tmp_path / "responses-deny.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _assert_blocked_message_item(item: dict[str, object], response: dict[str, object]) -> None: + assert item["type"] == "message", item + assert item["role"] == "assistant", item + assert item["status"] == "completed", item + assert str(item["id"]).startswith("msg_"), item + assert item["content"] == [{"type": "output_text", "text": _RESPONSES_DENIAL, "annotations": []}], item + assert response["status"] == "completed", response + usage: Final = response["usage"] + assert isinstance(usage, dict), response + assert (usage["input_tokens"], usage["output_tokens"], usage["total_tokens"]) == (0, 0, 0), usage + + +def _response_id(index: int, response: httpx.Response) -> str: + assert response.status_code == 200, (index, response.text) + if index % 3 == 0: + assert response.headers["content-type"].startswith("text/event-stream"), response.text + return str(_blocked_stream_events(response.text)[-1]["response"]["id"]) + if index % 3 == 1: + assert response.headers["content-type"].startswith("text/event-stream"), response.text + blocked: Final = _blocked_stream_events(response.text)[-1]["response"] + _assert_blocked_message_item(blocked["output"][0], blocked) + return str(blocked["id"]) + assert response.headers["content-type"].startswith("application/json"), response.text + body: Final = response.json() + _assert_blocked_message_item(body["output"][0], body) + return str(body["id"]) + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_streams_typed_message") +def test_responses_pre_call_denial_streams_sse_with_typed_message_item(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": "say hi", "stream": True, "guardrails": [identity]}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), ( + response.headers["content-type"], + response.text, + ) + lines: Final = tuple(line for line in response.text.split("\n") if line.startswith("data: ")) + assert lines[-1] == "data: [DONE]", response.text + events: Final = tuple(json.loads(line.removeprefix("data: ")) for line in lines[:-1]) + kinds: Final = tuple(event["type"] for event in events) + assert tuple(kind for kind in kinds if kind != "response.output_text.delta") == ( + "response.created", + "response.in_progress", + "response.output_item.added", + "response.content_part.added", + "response.output_text.done", + "response.content_part.done", + "response.output_item.done", + "response.completed", + ), kinds + assert kinds.index("response.output_text.delta") == kinds.index("response.content_part.added") + 1, kinds + assert "".join(event["delta"] for event in events if event["type"] == "response.output_text.delta") == ( + _RESPONSES_DENIAL + ) + completed: Final = events[-1]["response"] + assert completed["output"] == [events[-2]["item"]], (completed, events[-2]) + _assert_blocked_message_item(completed["output"][0], completed) + assert observed.get("/__observations").json()["requests"] == [] + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_returns_typed_message") +def test_responses_pre_call_denial_returns_json_with_typed_message_item(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "guardrails": [identity]} + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("application/json"), response.headers["content-type"] + body: Final = response.json() + assert body["object"] == "response", body + assert len(body["output"]) == 1, body + _assert_blocked_message_item(body["output"][0], body) + assert observed.get("/__observations").json()["requests"] == [] + + +_RESPONSES_OUTPUT_DENIAL: Final = "Output withheld by policy." +_UPSTREAM_INPUT_TOKENS: Final = 20 +_UPSTREAM_OUTPUT_TOKENS: Final = 20 +_UPSTREAM_TOTAL_TOKENS: Final = 40 + + +def _responses_output_denial_config(tmp_path: Path, identity: str, model: str) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "custom_code", + "mode": "post_call", + "default_on": False, + "custom_code": ( + "def apply_guardrail(inputs, request_data, input_type):\n" + f" return block({_RESPONSES_OUTPUT_DENIAL!r})\n" + ), + }, + } + ] + config["policies"] = { + f"{identity}-pipeline": { + "guardrails": {"add": [identity]}, + "pipeline": { + "mode": "post_call", + "steps": [ + { + "guardrail": identity, + "on_pass": "allow", + "on_fail": "modify_response", + "modify_response_message": _RESPONSES_OUTPUT_DENIAL, + } + ], + }, + } + } + config["policy_attachments"] = [{"policy": f"{identity}-pipeline", "models": [model]}] + path: Final = tmp_path / "responses-output-deny.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _blocked_stream_events(text: str) -> tuple[dict[str, object], ...]: + lines: Final = tuple(line for line in text.split("\n") if line.startswith("data: ")) + assert lines[-1] == "data: [DONE]", text + return tuple(json.loads(line.removeprefix("data: ")) for line in lines[:-1]) + + +def _dead_api_base() -> str: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + port: Final = reserve.getsockname()[1] + return f"http://127.0.0.1:{port}/v1" + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_openai_sdk_streams_typed_message") +def test_responses_pre_call_denial_openai_sdk_streams_typed_message(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + client: Final = OpenAI( + base_url=f"{candidate.client.base_url}/v1", api_key=candidate.key, max_retries=0, timeout=15 + ) + events: Final = tuple( + client.responses.create(model=model, input="say hi", stream=True, extra_body={"guardrails": [identity]}) + ) + assert events[-1].type == "response.completed", [event.type for event in events] + completed: Final = events[-1].response + assert completed is not None and len(completed.output) == 1, completed + item: Final = completed.output[0] + assert item.type == "message", item + assert item.role == "assistant" and item.status == "completed", item + assert item.content[0].type == "output_text" and item.content[0].text == _RESPONSES_DENIAL, item.content + assert completed.usage is not None and completed.usage.total_tokens == 0, completed.usage + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_openai_async_sdk_streams_typed_message") +async def test_responses_pre_call_denial_openai_async_sdk_streams_typed_message( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + client: Final = AsyncOpenAI( + base_url=f"{candidate.client.base_url}/v1", api_key=candidate.key, max_retries=0, timeout=15 + ) + stream: Final = await client.responses.create( + model=model, input="say hi", stream=True, extra_body={"guardrails": [identity]} + ) + kinds: Final = [event.type async for event in stream] + assert kinds[-1] == "response.completed", kinds + assert "response.output_text.delta" in kinds, kinds + assert "response.in_progress" in kinds, kinds + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_openai_sdk_returns_typed_message") +def test_responses_pre_call_denial_openai_sdk_returns_typed_message(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + client: Final = OpenAI( + base_url=f"{candidate.client.base_url}/v1", api_key=candidate.key, max_retries=0, timeout=15 + ) + body: Final = client.responses.create(model=model, input="say hi", extra_body={"guardrails": [identity]}) + assert body.object == "response" and body.status == "completed", body + assert len(body.output) == 1, body.output + item: Final = body.output[0] + assert item.type == "message" and item.role == "assistant", item + assert item.content[0].type == "output_text" and item.content[0].text == _RESPONSES_DENIAL, item.content + assert body.output_text == _RESPONSES_DENIAL, body + assert body.usage is not None and body.usage.total_tokens == 0, body.usage + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_false_returns_json") +def test_responses_pre_call_denial_stream_false_returns_json(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "stream": False, "guardrails": [identity]} + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("application/json"), response.text + body: Final = response.json() + _assert_blocked_message_item(body["output"][0], body) + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_string_true_returns_json") +def test_responses_pre_call_denial_stream_string_true_returns_json(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "stream": "true", "guardrails": [identity]} + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("application/json"), ( + response.headers["content-type"], + response.text, + ) + body: Final = response.json() + _assert_blocked_message_item(body["output"][0], body) + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_event_vocabulary") +def test_responses_pre_call_denial_stream_event_vocabulary(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + second: Final = "guardrail-2-" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + loaded: Final = yaml.safe_load(config.read_text()) + loaded["guardrails"].append(_deny_guardrail(second)) + config.write_text(yaml.safe_dump(loaded)) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": "say hi", "stream": True, "guardrails": [identity, second]}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + events: Final = _blocked_stream_events(response.text) + kinds: Final = {event["type"] for event in events} + assert kinds == { + "response.created", + "response.in_progress", + "response.output_item.added", + "response.content_part.added", + "response.output_text.delta", + "response.output_text.done", + "response.content_part.done", + "response.output_item.done", + "response.completed", + }, kinds + item_done: Final = tuple(event for event in events if event["type"] == "response.output_item.done") + assert len(item_done) == 1, events + assert len(events[-1]["response"]["output"]) == 1, events[-1] + assert observed.get("/__observations").json()["requests"] == [] + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_large_denial_text") +def test_responses_pre_call_denial_stream_large_denial_text(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + denial: Final = ("Denied: " + "mixed ascii and unicode text " * 200 + "fin")[:5000] + config: Final = _responses_denial_config(tmp_path, identity, denial) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": "say hi", "stream": True, "guardrails": [identity]}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + events: Final = _blocked_stream_events(response.text) + assert "".join(event["delta"] for event in events if event["type"] == "response.output_text.delta") == denial + done: Final = next(event for event in events if event["type"] == "response.output_text.done") + assert done["text"] == denial, done + completed: Final = events[-1]["response"] + assert completed["output"][0]["content"][0]["text"] == denial, completed + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_requests_have_distinct_ids") +def test_responses_pre_call_denial_stream_requests_have_distinct_ids(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + responses: Final = tuple( + candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": "say hi", "stream": True, "guardrails": [identity]}, + ) + for _ in range(2) + ) + completed: Final = tuple(_blocked_stream_events(response.text)[-1]["response"] for response in responses) + for response in responses: + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + assert completed[0]["id"] != completed[1]["id"], completed + assert completed[0]["output"][0]["id"] != completed[1]["output"][0]["id"], completed + assert observed.get("/__observations").json()["requests"] == [] + + +def _register_named_model(candidate: Gateway, name: str, api_base: str | None = None, **parameters: object) -> str: + created: Final = candidate.post( + "/model/new", + { + "model_name": name, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "integration-provider-key", + "api_base": api_base or f"{candidate.upstream_url}/v1", + **parameters, + }, + }, + ) + return str(created["model_info"]["id"]) + + +@pytest.mark.covers("other.observability.guardrails.responses_post_call_pipeline_denial_streams_real_usage") +def test_responses_post_call_pipeline_denial_streams_real_usage(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + model: Final = f"integration-{uuid.uuid4().hex}" + config: Final = _responses_output_denial_config(tmp_path, identity, model) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate: + model_id: Final = _register_named_model(candidate, model, use_chat_completions_api=True) + try: + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": f"say hi {uuid.uuid4().hex}", "stream": True} + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + events: Final = _blocked_stream_events(response.text) + assert events[-1]["type"] == "response.completed", events + completed: Final = events[-1]["response"] + item: Final = completed["output"][0] + assert item["type"] == "message" and item["role"] == "assistant", item + assert item["content"][0]["type"] == "output_text", item + assert item["content"][0]["text"] == _RESPONSES_OUTPUT_DENIAL, item + usage: Final = completed["usage"] + assert ( + usage["input_tokens"], + usage["output_tokens"], + usage["total_tokens"], + ) == (_UPSTREAM_INPUT_TOKENS, _UPSTREAM_OUTPUT_TOKENS, _UPSTREAM_TOTAL_TOKENS), usage + finally: + candidate.post("/model/delete", {"id": model_id}) + + +@pytest.mark.covers("other.observability.guardrails.responses_post_call_pipeline_denial_returns_real_usage") +def test_responses_post_call_pipeline_denial_returns_real_usage(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + model: Final = f"integration-{uuid.uuid4().hex}" + config: Final = _responses_output_denial_config(tmp_path, identity, model) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate: + model_id: Final = _register_named_model(candidate, model, use_chat_completions_api=True) + try: + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": f"say hi {uuid.uuid4().hex}"} + ) + assert response.status_code == 200, response.text + body: Final = response.json() + item: Final = body["output"][0] + assert item["type"] == "message" and item["role"] == "assistant", item + assert item["content"][0]["type"] == "output_text", item + assert item["content"][0]["text"] == _RESPONSES_OUTPUT_DENIAL, item + usage: Final = body["usage"] + assert ( + usage["input_tokens"], + usage["output_tokens"], + usage["total_tokens"], + ) == (_UPSTREAM_INPUT_TOKENS, _UPSTREAM_OUTPUT_TOKENS, _UPSTREAM_TOTAL_TOKENS), usage + finally: + candidate.post("/model/delete", {"id": model_id}) + + +@pytest.mark.covers("other.observability.guardrails.responses_denial_requires_authentication") +def test_responses_denial_requires_authentication(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "guardrails": [identity]}, key="sk-invalid" + ) + assert response.status_code == 401, (response.status_code, response.text) + assert response.json()["error"]["type"] == "token_not_found_in_db", response.text + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_does_not_reach_upstream") +def test_responses_pre_call_denial_stream_does_not_reach_upstream(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(api_base=_dead_api_base()) + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": "say hi", "stream": True, "guardrails": [identity]}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + events: Final = _blocked_stream_events(response.text) + completed: Final = events[-1]["response"] + _assert_blocked_message_item(completed["output"][0], completed) + + +@pytest.mark.covers("other.observability.guardrails.responses_unguarded_stream_reaches_upstream") +def test_responses_unguarded_stream_reaches_upstream(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + dead: Final = scenario.model(api_base=_dead_api_base()) + denied: Final = candidate.request( + "POST", + "/v1/responses", + {"model": dead, "input": "say hi", "stream": True, "guardrails": [identity]}, + ) + assert denied.status_code == 200, denied.text + model: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": f"say hi {uuid.uuid4().hex}", "stream": True, "guardrails": []}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + assert "response.completed" in response.text, response.text + requests: Final = eventually( + lambda: observed.get("/__observations").json()["requests"], + lambda values: len(values) >= 1, + seconds=30, + ) + assert len(requests) == 1, requests + assert requests[0]["path"] == "/v1/chat/completions", requests + + +@pytest.mark.covers("other.observability.guardrails.chat_pre_call_denial_streams_content_filter") +def test_chat_pre_call_denial_streams_content_filter(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "say hi"}], + "stream": True, + "guardrails": [identity], + }, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + lines: Final = tuple(line for line in response.text.split("\n") if line.startswith("data: ")) + assert lines[-1] == "data: [DONE]", response.text + chunks: Final = tuple(json.loads(line.removeprefix("data: ")) for line in lines[:-1]) + assert chunks[0]["choices"][0]["delta"]["content"] == _RESPONSES_DENIAL, chunks + assert chunks[-1]["choices"][0]["finish_reason"] == "stop", chunks + + +@pytest.mark.covers("other.observability.guardrails.chat_pre_call_denial_returns_content_filter") +def test_chat_pre_call_denial_returns_content_filter(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "say hi"}], + "guardrails": [identity], + }, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + choice: Final = body["choices"][0] + assert choice["finish_reason"] == "content_filter", body + assert choice["message"]["content"] == _RESPONSES_DENIAL, body + assert ( + body["usage"]["prompt_tokens"], + body["usage"]["completion_tokens"], + body["usage"]["total_tokens"], + ) == (0, 0, 0), body["usage"] + + +@pytest.mark.covers("other.observability.guardrails.messages_pre_call_denial_returns_message") +def test_messages_pre_call_denial_returns_message(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "messages": [{"role": "user", "content": "say hi"}], + "max_tokens": 16, + "guardrails": [identity], + }, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + assert body["type"] == "message" and body["role"] == "assistant", body + assert body["content"] == [{"type": "text", "text": _RESPONSES_DENIAL}], body + assert body["stop_reason"] == "end_turn", body + assert (body["usage"]["input_tokens"], body["usage"]["output_tokens"]) == (0, 0), body + + +@pytest.mark.covers("other.observability.guardrails.messages_pre_call_denial_streams_message") +def test_messages_pre_call_denial_streams_message(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "messages": [{"role": "user", "content": "say hi"}], + "max_tokens": 16, + "stream": True, + "guardrails": [identity], + }, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + lines: Final = tuple(line for line in response.text.split("\n") if line.startswith("data: ")) + assert len(lines) == 1, response.text + body: Final = json.loads(lines[0].removeprefix("data: ")) + assert body["type"] == "message" and body["role"] == "assistant", body + assert body["content"] == [{"type": "text", "text": _RESPONSES_DENIAL}], body + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_writes_zero_spend_row") +def test_responses_pre_call_denial_writes_zero_spend_row(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "guardrails": [identity]} + ) + assert response.status_code == 200, response.text + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, total_tokens FROM "LiteLLM_SpendLogs" WHERE model=%s AND call_type=%s', + (model, "aresponses"), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert float(rows[0]["spend"]) == 0, rows + assert rows[0]["total_tokens"] == 0, rows + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_survives_worker_burst") +def test_responses_pre_call_denial_stream_survives_worker_burst(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(api_base=_dead_api_base()) + healthy: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + + def burst(index: int) -> httpx.Response: + if index % 3 == 0: + return candidate.request( + "POST", + "/v1/responses", + {"model": healthy, "input": f"say hi {uuid.uuid4().hex} {index}", "stream": True}, + ) + stream: Final = index % 3 == 1 + return candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": f"say hi {index}", "stream": stream, "guardrails": [identity]}, + ) + + with ThreadPoolExecutor(max_workers=8) as pool: + responses: Final = tuple(pool.map(burst, range(30))) + response_ids: Final = frozenset(_response_id(index, response) for index, response in enumerate(responses)) + assert len(response_ids) == 30, response_ids + assert len(observed.get("/__observations").json()["requests"]) == 10 + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_survives_worker_kill") +def test_responses_pre_call_denial_stream_survives_worker_kill(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + with candidate.scenario() as scenario: + model: Final = scenario.model(api_base=_dead_api_base()) + members: Final = tuple( + member for member in group_members(owned.process.pid) if member.pid != owned.process.pid + ) + children: Final = tuple(member.pid for member in members) + workers: Final = tuple( + member.pid for member in members if any("spawn_main" in part for part in member.cmdline()) + ) + assert len(workers) >= 2, workers + os.kill(workers[0], signal.SIGKILL) + expected: Final = len(children) + eventually( + lambda: tuple( + member.pid + for member in group_members(owned.process.pid) + if member.pid != owned.process.pid + and member.is_running() + and member.status() != psutil.STATUS_ZOMBIE + ), + lambda pids: len(pids) >= expected and any(pid not in children for pid in pids), + seconds=30, + ) + + def burst(index: int) -> httpx.Response: + return candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": f"say hi {index}", "stream": True, "guardrails": [identity]}, + ) + + with ThreadPoolExecutor(max_workers=5) as pool: + responses: Final = tuple(pool.map(burst, range(10))) + for response in responses: + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text 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 e684aa55b33..656dc33e88c 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -3,6 +3,7 @@ Test for response_api_endpoints/endpoints.py """ import unittest +from collections.abc import Mapping from typing import Any, Final, Literal from unittest.mock import AsyncMock, MagicMock, patch @@ -14,6 +15,7 @@ from httpx import Response import litellm from litellm.proxy.proxy_server import app +from litellm.types.llms.openai import ResponsesAPIResponse @pytest.mark.asyncio @@ -2193,6 +2195,59 @@ class TestCursorGateRecognizesRoutingGroups: assert "reasoning_effort" not in resolved +BLOCK_MESSAGE = "Content flagged by policy, response withheld" + + +def _post_blocked_responses( + original_response: ResponsesAPIResponse | litellm.ModelResponse | None, + payload: Mapping[str, object] | None = None, +) -> httpx.Response: + from litellm.integrations.custom_guardrail import ModifyResponseException + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + exc = ModifyResponseException( + message=BLOCK_MESSAGE, + model="gpt-4o-mini", + request_data={"model": "gpt-4o-mini", "input": "hi"}, + guardrail_name="zero-usage-regression", + original_response=original_response, + ) + mock_proxy_logging = MagicMock() + mock_proxy_logging.post_call_failure_hook = AsyncMock() + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="sk-test", request_route="/v1/responses" + ) + body = {"model": "gpt-4o-mini", "input": "Write a haiku about token accounting"} + if payload: + body.update(payload) + try: + with ( + patch( + "litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", + new=AsyncMock(side_effect=exc), + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging), + ): + client = TestClient(app) + return client.post("/v1/responses", json=body, headers={"Authorization": "Bearer sk-1234"}) + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + +def _assert_blocked_output_item(item: Mapping[str, object], text: str) -> None: + assert item["type"] == "message" + assert item["id"].startswith("msg_") + assert item["role"] == "assistant" + assert item["status"] == "completed" + assert item["content"][0]["type"] == "output_text" + assert item["content"][0]["text"] == text + + +def _sse_data_frames(text: str) -> list[str]: + return [line.removeprefix("data: ").strip() for line in text.splitlines() if line.startswith("data: ")] + + class TestGuardrailBlockedResponsesUsage: """Regression tests for https://github.com/BerriAI/litellm/issues/36880. @@ -2202,38 +2257,7 @@ class TestGuardrailBlockedResponsesUsage: e.original_response, exactly like /v1/chat/completions already does.""" def _post_blocked_responses(self, original_response): - from litellm.integrations.custom_guardrail import ModifyResponseException - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.auth.user_api_key_auth import user_api_key_auth - - exc = ModifyResponseException( - message="Content flagged by policy, response withheld", - model="gpt-4o-mini", - request_data={"model": "gpt-4o-mini", "input": "hi"}, - guardrail_name="zero-usage-regression", - original_response=original_response, - ) - mock_proxy_logging = MagicMock() - mock_proxy_logging.post_call_failure_hook = AsyncMock() - app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( - api_key="sk-test", request_route="/v1/responses" - ) - try: - with ( - patch( - "litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", - new=AsyncMock(side_effect=exc), - ), - patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging), - ): - client = TestClient(app) - return client.post( - "/v1/responses", - json={"model": "gpt-4o-mini", "input": "Write a haiku about token accounting"}, - headers={"Authorization": "Bearer sk-1234"}, - ) - finally: - app.dependency_overrides.pop(user_api_key_auth, None) + return _post_blocked_responses(original_response) def test_post_call_block_reports_real_upstream_usage(self): from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse @@ -2429,6 +2453,66 @@ class TestResponsesInputTokens: assert response.json()["error"]["message"] == "rate limited" +class TestGuardrailBlockedResponsesShape: + """A pre_call block raises ModifyResponseException before any provider call. + + The reply must satisfy the Responses API contract the request selected: + stream=true answers SSE ending in one response.completed whose output[0] is + a completed assistant message item with output_text content, and a plain + POST answers JSON with the same item, both with the usage the blocked call + consumed (zero for pre_call).""" + + def test_non_stream_block_is_a_completed_assistant_message(self): + response = _post_blocked_responses(None) + + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("application/json") + body = response.json() + _assert_blocked_output_item(body["output"][0], BLOCK_MESSAGE) + assert body["usage"]["total_tokens"] == 0 + + def test_stream_block_answers_sse_with_completed_event(self): + response = _post_blocked_responses(None, payload={"stream": True}) + + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream") + frames = _sse_data_frames(response.text) + assert frames[-1] == "[DONE]" + events = [json.loads(frame) for frame in frames[:-1]] + types = [event["type"] for event in events] + assert "response.created" in types + completed = [event for event in events if event["type"] == "response.completed"] + assert len(completed) == 1 + completed_response = completed[0]["response"] + _assert_blocked_output_item(completed_response["output"][0], BLOCK_MESSAGE) + assert completed_response["usage"]["total_tokens"] == 0 + delta_text = "".join(event["delta"] for event in events if event["type"] == "response.output_text.delta") + assert delta_text == BLOCK_MESSAGE + + def test_stream_block_keeps_upstream_usage(self): + from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse + + original = ResponsesAPIResponse( + id="resp_upstream", + created_at=1, + model="gpt-4o-mini", + object="response", + output=[], + status="completed", + usage=ResponseAPIUsage(input_tokens=14, output_tokens=20, total_tokens=34), + ) + + response = _post_blocked_responses(original, payload={"stream": True}) + + assert response.status_code == 200, response.text + frames = _sse_data_frames(response.text) + completed = [json.loads(frame) for frame in frames[:-1] if json.loads(frame)["type"] == "response.completed"] + usage = completed[0]["response"]["usage"] + assert usage["input_tokens"] == 14 + assert usage["output_tokens"] == 20 + assert usage["total_tokens"] == 34 + + def test_responses_routes_document_response_models_in_openapi_schema(): from typing import cast diff --git a/tests/test_litellm/proxy/test_blocked_response_usage.py b/tests/test_litellm/proxy/test_blocked_response_usage.py index 4f20f35e94b..90d861be8e0 100644 --- a/tests/test_litellm/proxy/test_blocked_response_usage.py +++ b/tests/test_litellm/proxy/test_blocked_response_usage.py @@ -4,7 +4,7 @@ proxy endpoints (/v1/chat/completions, /v1/completions, and /v1/responses). A post-call block replaces the LLM response with the violation message, but the upstream call already consumed tokens. `_blocked_response_usage` (and its -Responses API counterpart `_blocked_responses_api_usage`) reports that real +Responses API counterpart `blocked_responses_api_usage`) reports that real usage (carried on `ModifyResponseException.original_response`) rather than zero; a pre-call block never invoked the LLM, so usage is zero. """ @@ -91,8 +91,8 @@ def test_responses_api_blocked_reply_carries_real_usage(): """ import time - from litellm.proxy.response_api_endpoints.endpoints import ( - _blocked_responses_api_usage, + from litellm.llms.base_llm.guardrail_translation.utils import ( + blocked_responses_api_usage, ) original_response = ResponsesAPIResponse( @@ -105,7 +105,7 @@ def test_responses_api_blocked_reply_carries_real_usage(): usage=ResponseAPIUsage(input_tokens=14, output_tokens=20, total_tokens=34), ) - usage = _blocked_responses_api_usage(original_response) + usage = blocked_responses_api_usage(original_response) assert usage.input_tokens == 14 assert usage.output_tokens == 20 @@ -114,11 +114,11 @@ def test_responses_api_blocked_reply_carries_real_usage(): def test_responses_api_blocked_reply_zero_usage_when_no_original_response(): """Pre-call block has no original_response, so usage must be zero.""" - from litellm.proxy.response_api_endpoints.endpoints import ( - _blocked_responses_api_usage, + from litellm.llms.base_llm.guardrail_translation.utils import ( + blocked_responses_api_usage, ) - usage = _blocked_responses_api_usage(None) + usage = blocked_responses_api_usage(None) assert usage.input_tokens == 0 assert usage.output_tokens == 0 @@ -128,14 +128,14 @@ def test_responses_api_blocked_reply_zero_usage_when_no_original_response(): def test_responses_api_blocked_reply_maps_bridged_chat_usage(): """A chat model bridged through /v1/responses blocks with a ModelResponse whose Usage fields must map prompt_tokens -> input_tokens and completion_tokens -> output_tokens.""" - from litellm.proxy.response_api_endpoints.endpoints import ( - _blocked_responses_api_usage, + from litellm.llms.base_llm.guardrail_translation.utils import ( + blocked_responses_api_usage, ) resp = litellm.ModelResponse() resp.usage = litellm.Usage(prompt_tokens=14, completion_tokens=18, total_tokens=32) - usage = _blocked_responses_api_usage(resp) + usage = blocked_responses_api_usage(resp) assert usage.input_tokens == 14 assert usage.output_tokens == 18 From b396b0b72499e7c79b280d8821a44a8b84835a7a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 18:05:55 -0700 Subject: [PATCH 05/65] feat(e2e): read management routes back from the control plane replicas (#43373) * feat(e2e): read management routes back from the control plane replicas The suite's management read-backs (/key/info, /team/info and friends) polled the same replica list as the data plane. On a componentized stack whose LITELLM_PROXY_REPLICA_URLS names the gateway pods directly, that list answers those routes 404, since a gateway pod trims the management routes at startup. A new LITELLM_CONTROL_PLANE_REPLICA_URLS names the addresses a management read-back polls instead: an exported list wins, and when it is unset the old rule stands, the data-plane replicas while the control plane shares the suite's base URL and the control-plane base alone once it is split. build_proxy_client takes the list as control_replica_urls and read_back_everywhere picks its replicas per path, the way the rest of the client already does. * fix(e2e): derive the control replicas of a client built for another proxy --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- tests/e2e/CONTRIBUTING.md | 2 +- tests/e2e/claude_code/_env.py | 1 + tests/e2e/claude_code/conftest.py | 1 + tests/e2e/e2e_config.py | 61 +++++++++++++++- tests/e2e/mcp/oauth_gateway.py | 1 + tests/e2e/proxy_client.py | 65 +++++++++++------ tests/e2e/test_proxy_client.py | 114 ++++++++++++++++++++++++++++-- 7 files changed, 215 insertions(+), 30 deletions(-) diff --git a/tests/e2e/CONTRIBUTING.md b/tests/e2e/CONTRIBUTING.md index 7e1f516422e..8e221b2da5e 100644 --- a/tests/e2e/CONTRIBUTING.md +++ b/tests/e2e/CONTRIBUTING.md @@ -105,7 +105,7 @@ A couple of logging destinations are configured on the proxy rather than by the ### The pull request check -Every same-repository PR that adds, modifies, or renames a `tests/e2e/**/test_*.py` file runs those changed files three times. A change to the harness itself, meaning a root-level `tests/e2e/*.py` file or `pytest.ini`, `tests/e2e/gateway/`, `.github/e2e-stack/`, or the workflow, also runs the `access_control` suite and both JWT suites as canaries, because those files have no test of their own that exercises the stack. `.github/e2e-stack/select_tests.py` applies both rules. The stack config at `tests/e2e/gateway/stage_mirror_ci_config.yml` must declare every model the selected suites use; a missing one shows up as a failed test id in the public log. The suite's own single rerun for network errors and 5xx responses (see `pytest.ini`) applies on every pass, so a transport blip does not fail the check while a race inside a test still does. The stage-mirror stack has a control-plane backend, two gateways behind nginx, Postgres, Keycloak, Jaeger, and TLS cluster-mode Valkey. Realm-only edits also trigger these canaries. The stack exports every gateway address in `LITELLM_PROXY_REPLICA_URLS`, so model registration waits until each gateway lists the new model rather than whichever one the load balancer answered from. Documentation, deleted-file, and application-only changes do not start the stack or request environment approval. The `ui/`, `claude_code/`, `load/`, and `secret_manager/` directories, `batches/test_managed_files_enforcement_e2e.py`, `llm_translation/realtime/test_realtime_pipecat_audio_e2e.py`, and `guardrails/test_presidio_masking_e2e.py` remain outside this check because they use separate tooling or need a differently configured stack: the pipecat audio suite skips itself at import time unless the NLTK `punkt_tab` data is installed, and the presidio suite fails without the analyzer and anonymizer services this stack does not start. `logging/test_otel_v2_langfuse_generation_output_e2e.py` is marked `otel_v2` and deselects itself unless `E2E_OTEL_V2` is set, because it needs a gateway booted with `LITELLM_OTEL_V2=true` and Langfuse credentials, neither of which this stack provides, so run it with `E2E_OTEL_V2=1` against a local OTel v2 proxy. The Redis chaos test under `load/` needs a proxy it can pause the Redis of on the same host (`gateway/redis_chaos_ci_config.yml`), which `.github/workflows/test-e2e-redis-chaos.yml` boots, and which the Buildkite `e2e-redis-chaos` step in project-releaser runs co-located with Postgres and Valkey in one pod; it is deselected unless `E2E_REDIS_CHAOS` is set. The `secret_manager/` lanes each need a proxy configured against their own secret manager (see Secret manager lanes below) +Every same-repository PR that adds, modifies, or renames a `tests/e2e/**/test_*.py` file runs those changed files three times. A change to the harness itself, meaning a root-level `tests/e2e/*.py` file or `pytest.ini`, `tests/e2e/gateway/`, `.github/e2e-stack/`, or the workflow, also runs the `access_control` suite and both JWT suites as canaries, because those files have no test of their own that exercises the stack. `.github/e2e-stack/select_tests.py` applies both rules. The stack config at `tests/e2e/gateway/stage_mirror_ci_config.yml` must declare every model the selected suites use; a missing one shows up as a failed test id in the public log. The suite's own single rerun for network errors and 5xx responses (see `pytest.ini`) applies on every pass, so a transport blip does not fail the check while a race inside a test still does. The stage-mirror stack has a control-plane backend, two gateways behind nginx, Postgres, Keycloak, Jaeger, and TLS cluster-mode Valkey. Realm-only edits also trigger these canaries. The stack exports every gateway address in `LITELLM_PROXY_REPLICA_URLS`, so model registration waits until each gateway lists the new model rather than whichever one the load balancer answered from. The Buildkite PR stack exports its two gateway pods the same way and, because those pods sit behind one router base that also fronts the backend, names that base in `LITELLM_CONTROL_PLANE_REPLICA_URLS` so management read-backs poll the plane that serves them instead of the gateway pods, which trim management routes at startup and answer them 404. Documentation, deleted-file, and application-only changes do not start the stack or request environment approval. The `ui/`, `claude_code/`, `load/`, and `secret_manager/` directories, `batches/test_managed_files_enforcement_e2e.py`, `llm_translation/realtime/test_realtime_pipecat_audio_e2e.py`, and `guardrails/test_presidio_masking_e2e.py` remain outside this check because they use separate tooling or need a differently configured stack: the pipecat audio suite skips itself at import time unless the NLTK `punkt_tab` data is installed, and the presidio suite fails without the analyzer and anonymizer services this stack does not start. `logging/test_otel_v2_langfuse_generation_output_e2e.py` is marked `otel_v2` and deselects itself unless `E2E_OTEL_V2` is set, because it needs a gateway booted with `LITELLM_OTEL_V2=true` and Langfuse credentials, neither of which this stack provides, so run it with `E2E_OTEL_V2=1` against a local OTel v2 proxy. The Redis chaos test under `load/` needs a proxy it can pause the Redis of on the same host (`gateway/redis_chaos_ci_config.yml`), which `.github/workflows/test-e2e-redis-chaos.yml` boots, and which the Buildkite `e2e-redis-chaos` step in project-releaser runs co-located with Postgres and Valkey in one pod; it is deselected unless `E2E_REDIS_CHAOS` is set. The `secret_manager/` lanes each need a proxy configured against their own secret manager (see Secret manager lanes below) Every selected file must execute at least one passing test in each pass, and any test failure, collection error, or entirely skipped or deselected file fails the check. A file whose tests are all marked skip therefore cannot pass this check, so unskip at least one of them, or add the file to `UNSUPPORTED` in `select_tests.py` with the reason, before changing one. A failed pass stops the run. The public log prints pytest's one-line summary for each pass, including the rerun count, and names each failed or errored test as `classname::name`, so a retried network error or a failing test is visible without the raw output. The final `e2e-changed-tests` job succeeds only when no supported test files changed or the approved run completed all three passes. Fork PRs with selected tests fail this gate until a maintainer brings the reviewed change onto a same-repository branch diff --git a/tests/e2e/claude_code/_env.py b/tests/e2e/claude_code/_env.py index 889d8f848dd..431be349a95 100644 --- a/tests/e2e/claude_code/_env.py +++ b/tests/e2e/claude_code/_env.py @@ -100,5 +100,6 @@ def require_proxy_client( master_key=cfg.api_key, control_plane_base_url=cfg.base_url, replica_urls=(cfg.base_url,), + control_replica_urls=(cfg.base_url,), ) return ProxyClientConfig(client=client, api_key=cfg.api_key) diff --git a/tests/e2e/claude_code/conftest.py b/tests/e2e/claude_code/conftest.py index bf226161267..c69f2dae462 100644 --- a/tests/e2e/claude_code/conftest.py +++ b/tests/e2e/claude_code/conftest.py @@ -595,6 +595,7 @@ def _build_control_plane_client(proxy_config: ProxyConfig): master_key=proxy_config.api_key, control_plane_base_url=proxy_config.base_url, replica_urls=(proxy_config.base_url,), + control_replica_urls=(proxy_config.base_url,), ) diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index ca3b74281ae..b2682c04841 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -7,6 +7,7 @@ environment so the same tests run against localhost or a deployed proxy. from __future__ import annotations import os +from dataclasses import dataclass import time import uuid from pathlib import Path @@ -33,12 +34,68 @@ CONTROL_PLANE_BASE_URL = os.environ.get( ).rstrip("/") +def split_replica_urls(raw: str) -> tuple[str, ...]: + return tuple(dict.fromkeys(url.strip().rstrip("/") for url in raw.split(",") if url.strip())) + + def parse_replica_urls(raw: str, fallback: str) -> tuple[str, ...]: - urls: Final = tuple(dict.fromkeys(url.strip().rstrip("/") for url in raw.split(",") if url.strip())) - return urls or (fallback,) + return split_replica_urls(raw) or (fallback,) + + +def parse_control_plane_replica_urls( + raw: str, *, control_plane_base_url: str, base_url: str, replica_urls: tuple[str, ...] +) -> tuple[str, ...]: + """The replicas a management read-back polls. LITELLM_CONTROL_PLANE_REPLICA_URLS + names them outright; unset, they follow the two base URLs: every data-plane + replica when the planes share a base (a monolith serves every route from every + replica) and the control-plane base alone when they differ. A stack sets it when + LITELLM_PROXY_REPLICA_URLS names gateway pods behind a shared router base, since + a gateway trims the management routes at startup and answers them 404.""" + explicit: Final = split_replica_urls(raw) + if explicit: + return explicit + return replica_urls if control_plane_base_url == base_url else (control_plane_base_url,) PROXY_REPLICA_URLS: Final = parse_replica_urls(os.environ.get("LITELLM_PROXY_REPLICA_URLS", ""), PROXY_BASE_URL) +CONTROL_PLANE_REPLICA_URLS: Final = parse_control_plane_replica_urls( + os.environ.get("LITELLM_CONTROL_PLANE_REPLICA_URLS", ""), + control_plane_base_url=CONTROL_PLANE_BASE_URL, + base_url=PROXY_BASE_URL, + replica_urls=PROXY_REPLICA_URLS, +) + + +@dataclass(frozen=True, slots=True) +class StackEndpoints: + base_url: str + control_plane_base_url: str + replica_urls: tuple[str, ...] + control_replica_urls: tuple[str, ...] + + def control_replica_urls_for( + self, *, base_url: str, control_plane_base_url: str, replica_urls: tuple[str, ...] + ) -> tuple[str, ...]: + """The control replicas a client built for these endpoints polls when its caller names none: + this stack's own list for this stack's endpoints, since an exported list describes one stack only, + and the base-URL rule for any other proxy.""" + if (base_url, control_plane_base_url, replica_urls) == ( + self.base_url, + self.control_plane_base_url, + self.replica_urls, + ): + return self.control_replica_urls + return parse_control_plane_replica_urls( + "", control_plane_base_url=control_plane_base_url, base_url=base_url, replica_urls=replica_urls + ) + + +ENV_STACK: Final = StackEndpoints( + base_url=PROXY_BASE_URL, + control_plane_base_url=CONTROL_PLANE_BASE_URL, + replica_urls=PROXY_REPLICA_URLS, + control_replica_urls=CONTROL_PLANE_REPLICA_URLS, +) UI_USERNAME = os.environ.get("E2E_UI_USERNAME", "admin") UI_PASSWORD = os.environ.get("E2E_UI_PASSWORD", MASTER_KEY) diff --git a/tests/e2e/mcp/oauth_gateway.py b/tests/e2e/mcp/oauth_gateway.py index 82bb5f7ba0b..b328c81687b 100644 --- a/tests/e2e/mcp/oauth_gateway.py +++ b/tests/e2e/mcp/oauth_gateway.py @@ -187,6 +187,7 @@ def owned_gateway(idp: Keycloak, directory: Path, cleanup: ExitStack) -> OAuthGa base_url=base_url, control_plane_base_url=base_url, replica_urls=(base_url,), + control_replica_urls=(base_url,), master_key=os.environ["LITELLM_MASTER_KEY"], ), _environment=environment, diff --git a/tests/e2e/proxy_client.py b/tests/e2e/proxy_client.py index 51ea9fbe7bd..4ea83e4b0d3 100644 --- a/tests/e2e/proxy_client.py +++ b/tests/e2e/proxy_client.py @@ -20,6 +20,7 @@ from typing import Final, Literal from e2e_config import ( CONTROL_PLANE_BASE_URL, + ENV_STACK, MASTER_KEY, POLL_INTERVAL, POLL_TIMEOUT, @@ -530,16 +531,16 @@ class ProxyClient: response_type: type[R], converged: Callable[[Result[R]], bool], ) -> Mapping[str, Result[R]]: - """GET `path` under the master key on every replica in PROXY_REPLICA_URLS (the - data-plane URL alone when the stack exports no per-gateway addresses), polling - each to poll_timeout until its read satisfies `converged`. Returns that read per - replica, or fails naming the first replica that never converged and its last - read. Behind a load balancer the single address proves one replica converged, - not all of them; only per-gateway addresses make this a fleet-wide proof.""" + """GET `path` under the master key on every replica that serves it (see + replicas_for), polling each to poll_timeout until its read satisfies + `converged`. Returns that read per replica, or fails naming the first replica + that never converged and its last read. Behind a load balancer the single + address proves one replica converged, not all of them; only per-replica + addresses make this a fleet-wide proof.""" outcomes: Final = await_converged_everywhere( { url: self._body_poller(transport, path, params, response_type) - for url, transport in self.replicas.items() + for url, transport in self.replicas_for(path).items() }, converged=converged, timeout=self.poll_timeout, @@ -765,13 +766,11 @@ class ProxyClient: def replicas_for(self, path: str) -> Mapping[str, Transport]: """The replicas that serve `path`: every data-plane replica for an LLM route, and for a management route the control-plane replicas, since the data-plane - replicas trim management routes and answer them 404. A monolith serves both - from every replica, so a management read-back polls all of them; a split - deployment exposes one control-plane address (there is one backend process - behind it on the stack these suites run against), so it polls that. A - control plane fronting several backends would need its own replica list to - prove each one converged, the way PROXY_REPLICA_URLS does for the gateways. - Never empty: a read-back against no replica would assert nothing and pass.""" + replicas trim management routes and answer them 404. CONTROL_PLANE_REPLICA_URLS + names those (see e2e_config): every data-plane replica for a monolith, the + control plane's own address for a split deployment, and the stack's own list + when its gateway pods sit behind a shared router base. Never empty: a + read-back against no replica would assert nothing and pass.""" replicas: Final = self.control_replicas if is_control_plane_path(path) else self.replicas assert replicas, f"no replica is configured to serve {path}, so a read-back there would prove nothing" return replicas @@ -1132,6 +1131,7 @@ def build_proxy_client( master_key: str = MASTER_KEY, control_plane_base_url: str = CONTROL_PLANE_BASE_URL, replica_urls: tuple[str, ...] = PROXY_REPLICA_URLS, + control_replica_urls: tuple[str, ...] | None = None, ) -> ProxyClient: """The ProxyClient every suite's client is built from: a SplitTransport that routes LLM calls to the data plane (PROXY_BASE_URL) and management/admin calls to the @@ -1139,15 +1139,24 @@ def build_proxy_client( base URLs are the same for a monolithic proxy, so routing is then a no-op. ``replica_urls`` (PROXY_REPLICA_URLS) names every data-plane replica the model barrier polls directly; it is the data-plane URL itself unless the stack - exports each gateway's own address. Management read-backs poll those same - replicas when the two planes share a base URL (a monolith, where every replica - serves every route) and the control plane alone when they differ (a split - deployment, where the data-plane replicas do not serve management routes). + exports each gateway's own address. ``control_replica_urls`` + (CONTROL_PLANE_REPLICA_URLS) names the replicas a management read-back polls: + those same replicas when the two planes share a base URL (a monolith, where + every replica serves every route), the control plane alone when they differ (a + split deployment, where the data-plane replicas do not serve management + routes), or the list the stack exports when its gateway pods sit behind a + shared router base, since a gateway pod trims management routes and its + address cannot stand in for the control plane. The endpoints are injectable for callers that resolve the proxy some other - way than ``e2e_config``'s env names (see ``claude_code/_env.py``); they must - pass all four together, since a caller that overrides only the data plane - would leave management calls and the replica poll pointed at the env defaults. + way than ``e2e_config``'s env names (see ``claude_code/_env.py``); they pass + the three URL parameters together, since a caller that overrides only the + data plane would leave management calls and the replica polls pointed at the + env defaults. An omitted ``control_replica_urls`` is derived from those three + (``ENV_STACK.control_replica_urls_for``): the env stack's own endpoints take + its exported list, any other proxy follows the base-URL rule above, so a + client built for a local test server never reads management state back from + the env proxy. Test-to-proxy traffic always goes over the wire, in every E2E_FIXTURE_MODE: record and replay scope to the proxy's provider-bound calls via the @@ -1170,8 +1179,18 @@ def build_proxy_client( for url in replica_urls } ) - control_replicas: Final = ( - replicas if control_plane_base_url == base_url else MappingProxyType({control_plane_base_url: split.control}) + control_replica_urls_named: Final = ( + control_replica_urls + if control_replica_urls is not None + else ENV_STACK.control_replica_urls_for( + base_url=base_url, control_plane_base_url=control_plane_base_url, replica_urls=replica_urls + ) + ) + control_replicas: Final = MappingProxyType( + { + url: HttpTransport(base_url=url, master_key=master_key, request_timeout=REQUEST_TIMEOUT) + for url in control_replica_urls_named + } ) return ProxyClient( transport=split, diff --git a/tests/e2e/test_proxy_client.py b/tests/e2e/test_proxy_client.py index a8f07ed6dd7..e1615ec5f65 100644 --- a/tests/e2e/test_proxy_client.py +++ b/tests/e2e/test_proxy_client.py @@ -24,7 +24,7 @@ from types import MappingProxyType from typing import Final, cast import pytest -from e2e_config import parse_replica_urls +from e2e_config import StackEndpoints, parse_control_plane_replica_urls, parse_replica_urls from e2e_http import NoBody, Result, Success, without_retries from idp import Keycloak from lifecycle import ResourceManager @@ -110,7 +110,7 @@ def caller_boundary( thread.start() url: Final = f"http://127.0.0.1:{server.server_port}" proxy: Final = build_proxy_client( - base_url=url, control_plane_base_url=url, replica_urls=(url,), master_key="bootstrap" + base_url=url, control_plane_base_url=url, replica_urls=(url,), control_replica_urls=(url,), master_key="bootstrap" ) try: yield ManagementClient(proxy=proxy, master_key="bootstrap"), received @@ -343,6 +343,50 @@ class TestParseReplicaUrls: assert parse_replica_urls(raw, "http://lb") == ("http://127.0.0.1:4010", "http://127.0.0.1:4011") +class TestParseControlPlaneReplicaUrls: + def test_an_exported_list_wins_over_the_base_url_rule(self) -> None: + assert parse_control_plane_replica_urls( + " http://router/, http://router ", + control_plane_base_url="http://router", + base_url="http://router", + replica_urls=("http://10.0.0.1:4000", "http://10.0.0.2:4000"), + ) == ("http://router",) + + def test_unset_with_one_shared_base_follows_the_data_plane_replicas(self) -> None: + assert parse_control_plane_replica_urls( + "", control_plane_base_url="http://lb", base_url="http://lb", replica_urls=("http://pod-1", "http://pod-2") + ) == ("http://pod-1", "http://pod-2") + + def test_unset_with_a_split_control_plane_polls_its_base_alone(self) -> None: + assert parse_control_plane_replica_urls( + "", control_plane_base_url="http://backend", base_url="http://lb", replica_urls=("http://gateway-1",) + ) == ("http://backend",) + + +class TestStackEndpointsControlReplicas: + STACK: Final = StackEndpoints( + base_url="http://router", + control_plane_base_url="http://router", + replica_urls=("http://10.0.0.1:4000", "http://10.0.0.2:4000"), + control_replica_urls=("http://router",), + ) + + def test_the_stacks_own_endpoints_take_its_exported_control_list(self) -> None: + assert self.STACK.control_replica_urls_for( + base_url="http://router", + control_plane_base_url="http://router", + replica_urls=("http://10.0.0.1:4000", "http://10.0.0.2:4000"), + ) == ("http://router",) + + def test_any_other_endpoints_follow_the_base_url_rule(self) -> None: + assert self.STACK.control_replica_urls_for( + base_url="http://router", control_plane_base_url="http://router", replica_urls=("http://10.0.0.1:4000",) + ) == ("http://10.0.0.1:4000",) + assert self.STACK.control_replica_urls_for( + base_url="http://lb", control_plane_base_url="http://backend", replica_urls=("http://gateway-1",) + ) == ("http://backend",) + + def _answers(answers: Iterable[str]) -> ReplicaRead[str]: it: Final = iter(answers) return lambda _timeout: next(it) @@ -390,6 +434,7 @@ class TestReplicasFor: base_url="http://lb", control_plane_base_url="http://backend", replica_urls=("http://gateway-1", "http://gateway-2"), + control_replica_urls=("http://backend",), ) assert set(client.replicas_for("/key/info")) == {"http://backend"} assert set(client.replicas_for("/project/info")) == {"http://backend"} @@ -400,9 +445,63 @@ class TestReplicasFor: base_url="http://lb", control_plane_base_url="http://lb", replica_urls=("http://pod-1", "http://pod-2"), + control_replica_urls=("http://pod-1", "http://pod-2"), ) assert set(client.replicas_for("/key/info")) == {"http://pod-1", "http://pod-2"} + def test_gateway_pods_behind_one_router_read_management_routes_back_from_the_router(self) -> None: + """The Buildkite PR stack names each gateway pod in PROXY_REPLICA_URLS while + both planes share the router base, so a management read-back polls the + router (CONTROL_PLANE_REPLICA_URLS) rather than the pods, which trim + management routes, while a data-plane read-back still polls every pod.""" + client: Final = build_proxy_client( + base_url="http://router", + control_plane_base_url="http://router", + replica_urls=("http://10.0.0.1:4000", "http://10.0.0.2:4000"), + control_replica_urls=("http://router",), + ) + assert set(client.replicas_for("/key/info")) == {"http://router"} + assert set(client.replicas_for("/v1/models")) == {"http://10.0.0.1:4000", "http://10.0.0.2:4000"} + + def test_a_client_built_for_another_proxy_reads_management_routes_back_from_that_proxy(self) -> None: + """A caller that points the client at its own server (test_provider_cache.py) + names no control list, so the derived one has to follow that server rather + than the env proxy, on a shared base and on split ones alike.""" + local: Final = build_proxy_client( + base_url="http://local", control_plane_base_url="http://local", replica_urls=("http://local",) + ) + assert set(local.replicas_for("/key/info")) == {"http://local"} + assert set(local.replicas_for("/v1/models")) == {"http://local"} + split: Final = build_proxy_client( + base_url="http://lb", control_plane_base_url="http://backend", replica_urls=("http://gateway-1",) + ) + assert set(split.replicas_for("/key/info")) == {"http://backend"} + assert set(split.replicas_for("/v1/models")) == {"http://gateway-1"} + + def test_management_read_backs_poll_the_control_replicas_only(self) -> None: + """A gateway pod answers /key/info 404 even after the write landed on the + control plane, so a read-back that polled the data-plane replicas for it + would never converge there.""" + with caller_boundary(status=404) as (pod, pod_headers), caller_boundary() as (router, router_headers): + pod_url: Final = next(iter(pod.proxy.replicas)) + router_url: Final = next(iter(router.proxy.replicas)) + proxy: Final = build_proxy_client( + base_url=router_url, + control_plane_base_url=router_url, + replica_urls=(pod_url,), + control_replica_urls=(router_url,), + master_key="bootstrap", + ) + read: Final = proxy.read_back_everywhere( + "/key/info", + params=NoBody(), + response_type=KeyInfoResponse, + converged=lambda result: isinstance(result, Success), + ) + assert set(read) == {router_url} + assert router_headers.get_nowait() == "Bearer bootstrap" + assert router_headers.empty() and pod_headers.empty() + def test_mcp_admin_routes_read_back_from_every_data_plane_replica(self) -> None: """/v1/mcp/* is a lazily mounted feature, so a data-plane replica serves it too and answers from its own in-memory registry. Routing it to the control @@ -412,6 +511,7 @@ class TestReplicasFor: base_url="http://lb", control_plane_base_url="http://backend", replica_urls=("http://gateway-1", "http://gateway-2"), + control_replica_urls=("http://backend",), ) assert set(client.replicas_for("/v1/mcp/server/abc")) == {"http://gateway-1", "http://gateway-2"} assert set(client.replicas_for("/v1/mcp/toolset/abc")) == {"http://gateway-1", "http://gateway-2"} @@ -531,6 +631,7 @@ class TestSplitCallerPropagation: base_url=data_url, control_plane_base_url=control_url, replica_urls=(data_url,), + control_replica_urls=(control_url,), master_key="bootstrap", ).with_caller(Caller(credential="tenant-token", kind="direct_jwt", role="team_member")) proxy.key_info("owned") @@ -543,8 +644,13 @@ class TestSplitCallerPropagation: response_type=KeyInfoResponse, converged=lambda result: isinstance(result, Success), ) - assert control_headers.get_nowait() == "Bearer tenant-token" - assert control_headers.get_nowait() == "Bearer tenant-token" + proxy.read_back_everywhere( + "/v1/models", + params=NoBody(), + response_type=ModelsListResponse, + converged=lambda result: isinstance(result, Success), + ) + assert tuple(control_headers.get_nowait() for _ in range(3)) == ("Bearer tenant-token",) * 3 assert data_headers.get_nowait() == "Bearer tenant-token" assert control_headers.empty() and data_headers.empty() From 501ef23f4aae19c2bdfee50e7d860107a90e9e7e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 18:07:09 -0700 Subject: [PATCH 06/65] feat(rust): add the openai_like chat config foundation (#43379) * feat(rust): add the openai_like chat config foundation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): let max_completion_tokens outrank max_tokens and decline refusal responses Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../core/src/chat_completions/common_utils.rs | 2 + litellm-rust/crates/llms/src/lib.rs | 1 + .../crates/llms/src/openai_like/chat/mod.rs | 1 + .../src/openai_like/chat/transformation.rs | 270 ++++++++++++++ .../llms/src/openai_like/common_utils.rs | 58 +++ .../crates/llms/src/openai_like/mod.rs | 2 + .../tests/openai_like_chat_transformation.rs | 343 ++++++++++++++++++ 7 files changed, 677 insertions(+) create mode 100644 litellm-rust/crates/llms/src/openai_like/chat/mod.rs create mode 100644 litellm-rust/crates/llms/src/openai_like/chat/transformation.rs create mode 100644 litellm-rust/crates/llms/src/openai_like/common_utils.rs create mode 100644 litellm-rust/crates/llms/src/openai_like/mod.rs create mode 100644 litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs diff --git a/litellm-rust/crates/core/src/chat_completions/common_utils.rs b/litellm-rust/crates/core/src/chat_completions/common_utils.rs index 4ed39a90366..1c7875c3c33 100644 --- a/litellm-rust/crates/core/src/chat_completions/common_utils.rs +++ b/litellm-rust/crates/core/src/chat_completions/common_utils.rs @@ -3,6 +3,7 @@ use litellm_llms::{ anthropic::chat::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG, base_llm::chat::transformation::BaseConfig, bedrock::chat::converse_transformation::BEDROCK_CHAT_COMPLETIONS_CONFIG, + openai_like::chat::transformation::OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG, }; use serde_json::{Map, Value}; @@ -14,6 +15,7 @@ pub(super) fn chat_completions_provider_config(provider: &str) -> Option<&'stati match provider { "anthropic" => Some(&ANTHROPIC_CHAT_COMPLETIONS_CONFIG), "bedrock" => Some(&BEDROCK_CHAT_COMPLETIONS_CONFIG), + "openai_like" => Some(&OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG), _ => None, } } diff --git a/litellm-rust/crates/llms/src/lib.rs b/litellm-rust/crates/llms/src/lib.rs index e71a9466c0c..a25b822d0c7 100644 --- a/litellm-rust/crates/llms/src/lib.rs +++ b/litellm-rust/crates/llms/src/lib.rs @@ -7,6 +7,7 @@ pub mod cohere; mod error; pub mod mistral; pub mod openai; +pub mod openai_like; pub mod reducto; pub mod vertex_ai; diff --git a/litellm-rust/crates/llms/src/openai_like/chat/mod.rs b/litellm-rust/crates/llms/src/openai_like/chat/mod.rs new file mode 100644 index 00000000000..f239b6921fa --- /dev/null +++ b/litellm-rust/crates/llms/src/openai_like/chat/mod.rs @@ -0,0 +1 @@ +pub mod transformation; diff --git a/litellm-rust/crates/llms/src/openai_like/chat/transformation.rs b/litellm-rust/crates/llms/src/openai_like/chat/transformation.rs new file mode 100644 index 00000000000..4e483ddf8cc --- /dev/null +++ b/litellm-rust/crates/llms/src/openai_like/chat/transformation.rs @@ -0,0 +1,270 @@ +//! `litellm/llms/openai_like/chat/transformation.py`: the chat config every +//! OpenAI-compatible endpoint shares. The body is already OpenAI-shaped, so +//! parameters pass through verbatim; the port keeps Python's two deviations, +//! the `max_completion_tokens` -> `max_tokens` rename and the usage +//! `*_tokens` null-to-zero sanitize. + +use litellm_auth::{CredentialPlacement, SecretValue}; +use litellm_core_utils::core_helpers::unix_now; +use litellm_types::{ + llms::openai::ChatMessage, + utils::{ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse}, +}; +use serde_json::{Map, Value, json}; + +use crate::{ + Error, + base_llm::{ + auth::AuthScheme, + chat::transformation::{ + BaseConfig, Headers, ProviderChatRequestData, ProviderChatResponseData, + ValidatedEnvironment, + }, + }, + openai_like::common_utils::{complete_openai_like_url, openai_compatible_provider_info}, +}; + +/// OpenAI parameter names the Rust path can place verbatim in the request body. +/// Tool parameters are absent on purpose: the message gate already declines +/// tool-call content, and a `tools` request that did get through would produce +/// a tool-call response this port cannot normalize yet, so it declines before +/// the call instead of after it. +const SUPPORTED_PARAMS: &[(&str, &str)] = &[ + ("frequency_penalty", "frequency_penalty"), + ("logit_bias", "logit_bias"), + ("logprobs", "logprobs"), + ("top_logprobs", "top_logprobs"), + ("max_tokens", "max_tokens"), + ("max_completion_tokens", "max_completion_tokens"), + ("modalities", "modalities"), + ("prediction", "prediction"), + ("n", "n"), + ("presence_penalty", "presence_penalty"), + ("seed", "seed"), + ("stop", "stop"), + ("stream_options", "stream_options"), + ("temperature", "temperature"), + ("top_p", "top_p"), + ("audio", "audio"), + ("web_search_options", "web_search_options"), + ("service_tier", "service_tier"), + ("safety_identifier", "safety_identifier"), + ("prompt_cache_key", "prompt_cache_key"), + ("prompt_cache_retention", "prompt_cache_retention"), + ("store", "store"), + ("response_format", "response_format"), +]; + +/// Call configuration the caller may pass that never enters the request body. +const CONFIG_PARAMS: &[&str] = &["custom_endpoint", "extra_headers", "max_retries"]; + +pub struct OpenAILikeChatConfig; + +pub const OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG: OpenAILikeChatConfig = OpenAILikeChatConfig; + +impl BaseConfig for OpenAILikeChatConfig { + fn supported_openai_param_mappings(&self) -> &'static [(&'static str, &'static str)] { + SUPPORTED_PARAMS + } + + fn get_complete_url( + &self, + api_base: Option<&str>, + _model: &str, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + let custom_endpoint = optional_params + .get("custom_endpoint") + .and_then(Value::as_bool) + .unwrap_or(false); + complete_openai_like_url(api_base, custom_endpoint, env_lookup) + } + + fn transform_request( + &self, + model: &str, + messages: Vec, + optional_params: Map, + ) -> Result { + let mut params = Map::from_iter( + optional_params + .into_iter() + .filter(|(key, _)| !CONFIG_PARAMS.contains(&key.as_str())), + ); + // Most OpenAI-compatible endpoints take `max_tokens`, not + // `max_completion_tokens`, so Python's `map_openai_params` renames it + // and lets it overwrite a `max_tokens` the caller also sent. + if let Some(limit) = params.remove("max_completion_tokens") { + params.insert("max_tokens".to_string(), limit); + } + let body = Map::from_iter( + [ + ("model".to_string(), json!(model)), + ("messages".to_string(), json!(messages)), + ] + .into_iter() + .chain(params), + ); + Ok(ProviderChatRequestData { + body: Value::Object(body), + stream_shape: Default::default(), + }) + } + + fn transform_response( + &self, + model: &str, + response: ProviderChatResponseData, + ) -> Result { + let mut body = response.body; + sanitize_usage(&mut body); + let body = body + .as_object() + .ok_or_else(|| Error::InvalidResponse("chat response is not an object".into()))?; + + let choices = body + .get("choices") + .and_then(Value::as_array) + .ok_or(Error::MissingField("choices"))? + .iter() + .enumerate() + .map(|(position, choice)| normalize_choice(position, choice)) + .collect::, _>>()?; + + let usage = body.get("usage").and_then(Value::as_object); + let field = |name: &str| { + usage + .and_then(|usage| usage.get(name)) + .and_then(Value::as_u64) + .unwrap_or(0) + }; + let details = usage.and_then(|usage| usage.get("prompt_tokens_details")); + + Ok(ChatCompletionsResponse { + created: body + .get("created") + .and_then(Value::as_u64) + .unwrap_or_else(unix_now), + model: body + .get("model") + .and_then(Value::as_str) + .unwrap_or(model) + .to_string(), + choices, + usage: litellm_types::utils::ChatCompletionsUsage { + prompt_tokens: field("prompt_tokens"), + completion_tokens: field("completion_tokens"), + total_tokens: field("total_tokens"), + prompt_tokens_details: litellm_types::utils::PromptTokensDetails { + cached_tokens: details + .and_then(|d| d.get("cached_tokens")) + .and_then(Value::as_u64) + .unwrap_or(0), + cache_creation_tokens: details + .and_then(|d| d.get("cache_creation_tokens")) + .and_then(Value::as_u64) + .unwrap_or(0), + text_tokens: details + .and_then(|d| d.get("text_tokens")) + .and_then(Value::as_u64) + .unwrap_or(0), + }, + }, + }) + } + + /// `OpenAILikeBase._validate_environment`: a forwarded `authorization` is + /// the whole credential, and any other call authenticates with the + /// resolved key as a bearer. The key resolves to `""` when neither the + /// deployment nor `OPENAI_LIKE_API_KEY` sets one, because vllm-compatible + /// endpoints take no key; Python still sends `Bearer ` in that case. + fn validate_environment( + &self, + headers: Headers, + api_key: Option<&str>, + _model: &str, + _optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + if headers + .iter() + .any(|(name, _)| name.eq_ignore_ascii_case("authorization")) + { + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Forwarded, + }); + } + let (_, key) = openai_compatible_provider_info(None, api_key, env_lookup); + Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret: SecretValue::new(key.unwrap_or_default()), + }, + }) + } + + fn config_params(&self) -> &'static [&'static str] { + CONFIG_PARAMS + } +} + +/// `OpenAILikeChatConfig._sanitize_usage_obj`: a provider that reports a null +/// `*_tokens` entry breaks OpenAI clients, so nulls become 0. Python scrubs +/// every top-level usage key ending in `_tokens`. +fn sanitize_usage(body: &mut Value) { + if let Some(usage) = body.get_mut("usage").and_then(Value::as_object_mut) { + for (key, value) in usage.iter_mut() { + if key.ends_with("_tokens") && value.is_null() { + *value = json!(0); + } + } + } +} + +fn normalize_choice(position: usize, choice: &Value) -> Result { + let message = choice + .get("message") + .and_then(Value::as_object) + .ok_or(Error::MissingField("message"))?; + if message + .get("tool_calls") + .and_then(Value::as_array) + .is_some_and(|calls| !calls.is_empty()) + { + // Python rewrites the lone tool call into content only under + // `json_mode`, a request flag `transform_response` cannot see, and the + // normalized type cannot carry tool calls at all. Declining is + // terminal at this point, but passing back an empty assistant turn + // would fabricate the reply. + return Err(Error::Unsupported("tool call response")); + } + if message.get("refusal").is_some_and(|value| !value.is_null()) { + return Err(Error::Unsupported("refusal response")); + } + let content = message.get("content"); + if content.is_some_and(|value| !value.is_null() && !value.is_string()) { + return Err(Error::Unsupported("non-text response content")); + } + Ok(ChatCompletionsChoice { + index: choice + .get("index") + .and_then(Value::as_u64) + .unwrap_or(position as u64), + message: ChatCompletionsChoiceMessage { + role: message + .get("role") + .and_then(Value::as_str) + .unwrap_or("assistant") + .to_string(), + content: content.and_then(Value::as_str).map(str::to_string), + }, + finish_reason: choice + .get("finish_reason") + .and_then(Value::as_str) + .unwrap_or("") + .to_string(), + }) +} diff --git a/litellm-rust/crates/llms/src/openai_like/common_utils.rs b/litellm-rust/crates/llms/src/openai_like/common_utils.rs new file mode 100644 index 00000000000..b855e6dc812 --- /dev/null +++ b/litellm-rust/crates/llms/src/openai_like/common_utils.rs @@ -0,0 +1,58 @@ +//! Shared OpenAI-like credential and endpoint resolution, mirroring +//! `litellm/llms/openai_like/common_utils.py`. + +use crate::Error; + +/// `OpenAILikeChatConfig._get_openai_compatible_provider_info`: the deployment's +/// `api_base` wins over `OPENAI_LIKE_API_BASE`, and the deployment key over +/// `OPENAI_LIKE_API_KEY`, with an empty key allowed because vllm-compatible +/// endpoints do not require one. +pub fn openai_compatible_provider_info( + api_base: Option<&str>, + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> (Option, Option) { + let api_base = api_base + .map(str::to_string) + .or_else(|| env_lookup("OPENAI_LIKE_API_BASE")); + let api_key = api_key + .map(str::to_string) + .or_else(|| env_lookup("OPENAI_LIKE_API_KEY")) + .or(Some(String::new())); + (api_base, api_key) +} + +/// `OpenAILikeBase._validate_environment` requires an api base and, when the +/// caller gave no `custom_endpoint`, appends the route suffix. A caller-supplied +/// `custom_endpoint` base is used as is. +pub fn complete_openai_like_url( + api_base: Option<&str>, + custom_endpoint: bool, + env_lookup: &dyn Fn(&str) -> Option, +) -> Result { + let (api_base, _) = openai_compatible_provider_info(api_base, None, env_lookup); + let api_base = api_base.ok_or_else(|| { + Error::InvalidRequest( + "Missing API Base - A call is being made to LLM Provider but no api base is set either in the environment variables ({LLM_PROVIDER}_API_KEY) or via params" + .to_string(), + ) + })?; + if custom_endpoint { + return Ok(api_base); + } + Ok(format!( + "{}/chat/completions", + api_base.trim_end_matches('/') + )) +} + +/// The api key the call resolves to. `None` means neither the deployment nor the +/// environment supplied one, which is valid for endpoints that take no key. +pub fn resolve_openai_like_api_key( + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> Option { + openai_compatible_provider_info(None, api_key, env_lookup) + .1 + .filter(|key| !key.is_empty()) +} diff --git a/litellm-rust/crates/llms/src/openai_like/mod.rs b/litellm-rust/crates/llms/src/openai_like/mod.rs new file mode 100644 index 00000000000..df0cc73a5b0 --- /dev/null +++ b/litellm-rust/crates/llms/src/openai_like/mod.rs @@ -0,0 +1,2 @@ +pub mod chat; +pub mod common_utils; diff --git a/litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs b/litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs new file mode 100644 index 00000000000..8ee0654bddb --- /dev/null +++ b/litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs @@ -0,0 +1,343 @@ +use litellm_llms::{ + Error, + base_llm::{ + auth::AuthScheme, + chat::transformation::{BaseConfig, ProviderChatResponseData, Unsupported}, + }, + openai_like::chat::transformation::OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG, +}; +use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse}; +use rstest::rstest; +use serde_json::{Map, Value, json}; + +fn messages(value: Value) -> Vec { + serde_json::from_value(value).expect("valid messages") +} + +fn params(value: Value) -> Map { + match value { + Value::Object(map) => map, + other => panic!("params must be an object, got {other}"), + } +} + +fn no_env(_: &str) -> Option { + None +} + +fn env_with<'a>(name: &'a str, value: &'a str) -> impl Fn(&str) -> Option + 'a { + move |key| (key == name).then(|| value.to_string()) +} + +fn transform(model: &str, msgs: Value, opts: Value) -> Value { + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .transform_request(model, messages(msgs), params(opts)) + .expect("request transforms") + .body +} + +fn transform_response(body: Value) -> Result { + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .transform_response("some-model", ProviderChatResponseData { body }) +} + +fn reason(msgs: Value, opts: Value) -> Option { + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), ¶ms(opts)) +} + +#[rstest] +fn builds_the_openai_shaped_body() { + let body = transform( + "my-model", + json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"}, + ]), + json!({"temperature": 0.5, "max_tokens": 8}), + ); + assert_eq!(body["model"], json!("my-model")); + assert_eq!( + body["messages"], + json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"}, + ]) + ); + assert_eq!(body["temperature"], json!(0.5)); + assert_eq!(body["max_tokens"], json!(8)); +} + +#[rstest] +fn renames_max_completion_tokens_to_max_tokens() { + // `OpenAILikeChatConfig.map_openai_params`: most OpenAI-compatible providers + // support `max_tokens`, not `max_completion_tokens`. + let body = transform( + "my-model", + json!([{"role": "user", "content": "hi"}]), + json!({"max_completion_tokens": 12}), + ); + assert_eq!(body["max_tokens"], json!(12)); + assert!(body.get("max_completion_tokens").is_none()); +} + +#[rstest] +fn max_completion_tokens_wins_when_both_limits_are_sent() { + // Python assigns `max_tokens = max_completion_tokens` after copying the + // params, so the renamed value outranks a caller-supplied `max_tokens`. + let body = transform( + "my-model", + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 8, "max_completion_tokens": 12}), + ); + assert_eq!(body["max_tokens"], json!(12)); + assert!(body.get("max_completion_tokens").is_none()); +} + +#[rstest] +fn call_configuration_never_enters_the_body() { + let body = transform( + "my-model", + json!([{"role": "user", "content": "hi"}]), + json!({"custom_endpoint": true, "extra_headers": {"x": "y"}, "max_retries": 2}), + ); + assert_eq!( + body.as_object().unwrap().keys().collect::>(), + vec!["model", "messages"] + ); +} + +#[rstest] +#[case::appends_the_chat_completions_suffix("https://vllm.example.com/v1", json!({}), "https://vllm.example.com/v1/chat/completions")] +#[case::trims_a_trailing_slash("https://vllm.example.com/v1/", json!({}), "https://vllm.example.com/v1/chat/completions")] +#[case::a_custom_endpoint_is_used_as_is("https://vllm.example.com/v1/chat/completions", json!({"custom_endpoint": true}), "https://vllm.example.com/v1/chat/completions")] +fn complete_url(#[case] api_base: &str, #[case] opts: Value, #[case] expected: &str) { + assert_eq!( + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .get_complete_url(Some(api_base), "my-model", ¶ms(opts), &no_env) + .expect("url resolves"), + expected + ); +} + +#[rstest] +fn api_base_falls_back_to_the_environment() { + assert_eq!( + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .get_complete_url( + None, + "my-model", + ¶ms(json!({})), + &env_with("OPENAI_LIKE_API_BASE", "https://env.example.com/v1"), + ) + .expect("url resolves"), + "https://env.example.com/v1/chat/completions" + ); +} + +#[rstest] +fn a_missing_api_base_is_an_error() { + assert!(matches!( + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .get_complete_url(None, "my-model", ¶ms(json!({})), &no_env), + Err(Error::InvalidRequest(message)) if message.starts_with("Missing API Base") + )); +} + +#[rstest] +fn the_resolved_key_authenticates_as_a_bearer() { + let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .validate_environment( + vec![], + Some("sk-test"), + "my-model", + ¶ms(json!({})), + &no_env, + ) + .expect("validates"); + assert!(matches!( + validated.auth, + AuthScheme::Credential { + placement: litellm_auth::CredentialPlacement::Bearer, + .. + } + )); +} + +#[rstest] +fn the_key_falls_back_to_the_environment() { + let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .validate_environment( + vec![], + None, + "my-model", + ¶ms(json!({})), + &env_with("OPENAI_LIKE_API_KEY", "sk-env"), + ) + .expect("validates"); + let AuthScheme::Credential { secret, .. } = validated.auth else { + panic!("expected a bearer credential"); + }; + assert_eq!(secret.expose(), "sk-env"); +} + +#[rstest] +fn a_forwarded_authorization_is_the_whole_credential() { + // Python adds `Bearer ` only when the caller did not already send + // `Authorization`, so the forwarded header wins over the deployment key. + let headers = vec![("Authorization".to_string(), "Bearer caller".to_string())]; + let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .validate_environment( + headers, + Some("sk-test"), + "my-model", + ¶ms(json!({})), + &no_env, + ) + .expect("validates"); + assert!(matches!(validated.auth, AuthScheme::Forwarded)); +} + +#[rstest] +fn keyless_calls_still_validate_for_endpoints_that_take_no_key() { + // vllm-compatible endpoints require no api key; Python resolves `""` and + // sends `Bearer `, so validation must not fail on the missing key. + let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .validate_environment(vec![], None, "my-model", ¶ms(json!({})), &no_env) + .expect("validates"); + let AuthScheme::Credential { secret, .. } = validated.auth else { + panic!("expected a bearer credential"); + }; + assert_eq!(secret.expose(), ""); +} + +#[rstest] +fn normalizes_an_openai_response() { + let response = transform_response(json!({ + "created": 1_700_000_000, + "model": "served-model-name", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "hello"}, + "finish_reason": "stop", + }], + "usage": {"prompt_tokens": 3, "completion_tokens": 5, "total_tokens": 8}, + })) + .expect("response normalizes"); + assert_eq!(response.created, 1_700_000_000); + assert_eq!(response.model, "served-model-name"); + assert_eq!( + response.choices[0].message.content.as_deref(), + Some("hello") + ); + assert_eq!(response.choices[0].finish_reason, "stop"); + assert_eq!(response.usage.prompt_tokens, 3); + assert_eq!(response.usage.completion_tokens, 5); + assert_eq!(response.usage.total_tokens, 8); +} + +#[rstest] +fn null_token_fields_in_usage_become_zero() { + // `_sanitize_usage_obj`: providers that return null token values break + // OpenAI clients, so the response is scrubbed at the source. + let response = transform_response(json!({ + "model": "m", + "choices": [{"message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 3, "completion_tokens": null, "total_tokens": null}, + })) + .expect("response normalizes"); + assert_eq!(response.usage.completion_tokens, 0); + assert_eq!(response.usage.total_tokens, 0); + assert_eq!(response.usage.prompt_tokens, 3); +} + +#[rstest] +fn a_tool_call_response_declines_instead_of_dropping_the_calls() { + // The `json_mode` rewrite needs a request flag the route does not carry, so + // a tool-call answer falls back to Python rather than losing the calls. + assert_eq!( + transform_response(json!({ + "model": "m", + "choices": [{ + "message": { + "role": "assistant", + "content": null, + "tool_calls": [{ + "id": "call_1", + "type": "function", + "function": {"name": "f", "arguments": "{}"}, + }], + }, + "finish_reason": "tool_calls", + }], + })), + Err(Error::Unsupported("tool call response")) + ); +} + +#[rstest] +fn a_refusal_declines_instead_of_returning_an_empty_reply() { + assert_eq!( + transform_response(json!({ + "model": "m", + "choices": [{ + "message": {"role": "assistant", "content": null, "refusal": "cannot help"}, + "finish_reason": "stop", + }], + })), + Err(Error::Unsupported("refusal response")) + ); +} + +#[rstest] +fn a_non_text_response_content_declines() { + assert_eq!( + transform_response(json!({ + "model": "m", + "choices": [{ + "message": {"role": "assistant", "content": [{"type": "text", "text": "hi"}]}, + "finish_reason": "stop", + }], + })), + Err(Error::Unsupported("non-text response content")) + ); +} + +#[rstest] +#[case::streaming(json!({"stream": true}), "streaming")] +#[case::unrecognized_param(json!({"some_provider_knob": 1}), "unrecognized request parameter")] +fn declines(#[case] opts: Value, #[case] expected: &'static str) { + assert_eq!( + reason(json!([{"role": "user", "content": "hi"}]), opts), + Some(Unsupported(expected)) + ); +} + +#[rstest] +fn accepts_standard_openai_params() { + assert_eq!( + reason( + json!([{"role": "user", "content": "hi"}]), + json!({ + "temperature": 0.2, + "top_p": 0.9, + "max_tokens": 16, + "response_format": {"type": "json_object"}, + "custom_endpoint": true, + }), + ), + None + ); +} + +#[rstest] +fn tool_parameters_decline_before_the_call() { + // A `tools` request would come back with tool calls this port cannot + // normalize, so it declines at the gate instead of after the call. + assert_eq!( + reason( + json!([{"role": "user", "content": "hi"}]), + json!({"tools": [{"type": "function", "function": {"name": "f"}}]}), + ), + Some(Unsupported("unrecognized request parameter")) + ); +} From de06c937670f8298f9b304e9fbaf4034cd4007fd Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 26 Sep 2026 18:13:13 -0700 Subject: [PATCH 07/65] feat(router): opt in to prompt-cache cost routing (#43232) * feat(router): opt in to prompt-cache cost routing * fix(router): address prompt-cache routing review * ci: include cache-routing regressions in coverage --- .circleci/scripts/unit_selection.sh | 1 + litellm/llms/anthropic/cache_aware_routing.py | 236 +++++++ .../proxy/common_utils/cache_aware_routing.py | 343 ++++++++++ .../common_utils/prompt_cache_prediction.py | 20 + .../common_utils/prompt_cache_pricing.py | 25 +- .../prompt_cache_prediction.py | 113 +--- .../complexity_router/README.md | 41 +- .../complexity_router/complexity_router.py | 54 +- .../complexity_router/config.py | 19 + litellm/types/utils.py | 1 + .../common_utils/test_cache_aware_routing.py | 588 ++++++++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 20 +- 12 files changed, 1337 insertions(+), 124 deletions(-) create mode 100644 litellm/llms/anthropic/cache_aware_routing.py create mode 100644 litellm/proxy/common_utils/cache_aware_routing.py create mode 100644 litellm/proxy/common_utils/prompt_cache_prediction.py create mode 100644 tests/unit/proxy/common_utils/test_cache_aware_routing.py diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh index 3f4f5620176..6510b3fd4b5 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -46,6 +46,7 @@ legacy_paths() { echo tests/unit/google_genai echo tests/unit/router_strategy echo tests/unit/router_utils + echo tests/unit/proxy/common_utils/test_cache_aware_routing.py echo tests/unit/enterprise/enterprise_callbacks/send_emails echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py diff --git a/litellm/llms/anthropic/cache_aware_routing.py b/litellm/llms/anthropic/cache_aware_routing.py new file mode 100644 index 00000000000..e1a50781ace --- /dev/null +++ b/litellm/llms/anthropic/cache_aware_routing.py @@ -0,0 +1,236 @@ +from __future__ import annotations + +import time +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from types import MappingProxyType +from typing import Final +from urllib.parse import urlparse + +from pydantic import BaseModel, JsonValue, TypeAdapter + +import litellm +from litellm._internal_context import current_billing_time, pinned_billing_time +from litellm.caching.dual_cache import DualCache +from litellm.llms.anthropic.prompt_cache_prediction import ( + NativePredictionTarget, + PromptPrefix, + TokenCounter, + UnsupportedPredictionTarget, + cache_scope, + count_prompt_tokens, + parse_prompt, + resolve_prediction_target, + supported_prediction_headers, +) +from litellm.proxy.common_utils.prompt_cache_pricing import price_cache_tokens +from litellm.proxy.hooks.prompt_cache_prediction import lookup +from litellm.types.management_endpoints.prompt_cache_prediction import ( + CacheCostScenario, + CacheEvidence, + CachePredictionArm, + CacheTokenBuckets, +) +from litellm.types.router import Deployment +from litellm.utils import get_prompt_cache_min_tokens + +__all__: Final = ("AnthropicCacheRouting", "TokenCounter", "predict_arm") + +_JSON: Final = TypeAdapter(Mapping[str, JsonValue]) +_NATIVE_OPTIONS: Final = frozenset( + ( + "max_tokens", + "system", + "tools", + "tool_choice", + "thinking", + "output_config", + "cache_control", + "speed", + "service_tier", + "temperature", + "top_p", + "top_k", + "stop_sequences", + "stream", + ) +) + + +class _ModelLimits(BaseModel): + max_input_tokens: int | None = None + max_output_tokens: int | None = None + + +@dataclass(frozen=True, slots=True) +class AnthropicCacheRouting: + body: Mapping[str, JsonValue] + prefix: PromptPrefix + requested_output_limit: int + + @staticmethod + def request_body( + url: str, + headers: Mapping[str, str], + body: Mapping[str, JsonValue], + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, + ) -> Mapping[str, JsonValue] | None: + if not urlparse(url).path.endswith("/v1/messages") or not supported_prediction_headers(headers): + return None + return _JSON.validate_python( + MappingProxyType( + { + **body, + **MappingProxyType({key: request_kwargs[key] for key in _NATIVE_OPTIONS if key in request_kwargs}), + "messages": messages, + } + ) + ) + + @classmethod + def from_body(cls, body: Mapping[str, JsonValue]) -> AnthropicCacheRouting | None: + prefix: Final = parse_prompt(body) + limit: Final = body.get("max_tokens") + if prefix is None or not isinstance(limit, int) or isinstance(limit, bool) or limit <= 0: + return None + return cls(body, prefix, limit) + + @staticmethod + def supports(deployment: Deployment) -> bool: + return isinstance(resolve_prediction_target(deployment.litellm_params), NativePredictionTarget) + + async def is_warm(self, deployment: Deployment, caller: str, cache: DualCache, now: float) -> bool: + target: Final = resolve_prediction_target(deployment.litellm_params) + if not isinstance(target, NativePredictionTarget): + return False + scope: Final = cache_scope(caller, deployment.model_info.id or "", target.api_key, target.model) + observation: Final = await lookup(cache, scope, self.prefix, now=now) + return observation is not None and observation.expires_at > now + + @staticmethod + def fits(deployment: Deployment, input_tokens: int, output_tokens: int) -> bool: + target: Final = resolve_prediction_target(deployment.litellm_params) + if not isinstance(target, NativePredictionTarget): + return False + limits: Final = _ModelLimits.model_validate( + MappingProxyType( + { + **litellm.get_model_info(target.model, custom_llm_provider="anthropic"), + **deployment.model_info.model_dump(exclude_none=True), + } + ) + ) + return ( + limits.max_input_tokens is not None + and input_tokens + output_tokens <= limits.max_input_tokens + and limits.max_output_tokens is not None + and output_tokens <= limits.max_output_tokens + ) + + async def predict( + self, + deployment: Deployment, + caller: str, + cache: DualCache, + counter: TokenCounter, + now: float | None, + ) -> CachePredictionArm: + return await predict_arm(deployment, self.body, self.prefix, caller, cache, counter, now=now) + + @staticmethod + def cost(arm: CachePredictionArm, output_tokens: int) -> float | None: + return ( + price_cache_tokens(arm.model or "", arm.deployment_id, arm.estimate.tokens, output_tokens) + if arm.estimate is not None + else None + ) + + @staticmethod + async def count_tokens(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + return await count_prompt_tokens(model, api_key, body) + + +def _buckets(prefix_tokens: int, suffix_tokens: int, read_tokens: int, ttl_seconds: int) -> CacheTokenBuckets: + return CacheTokenBuckets( + uncached_input_tokens=suffix_tokens, + cache_read_input_tokens=read_tokens, + cache_creation_5m_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 300 else 0, + cache_creation_1h_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 3600 else 0, + ) + + +def _scenario(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> CacheCostScenario | None: + cost: Final = price_cache_tokens(model=model, deployment_id=deployment_id, tokens=tokens) + return CacheCostScenario(tokens=tokens, input_cost=cost) if cost is not None else None + + +async def predict_arm( + deployment: Deployment, + body: Mapping[str, JsonValue], + prefix: PromptPrefix, + caller_key_hash: str, + cache: DualCache, + token_counter: TokenCounter, + now: float | None = None, +) -> CachePredictionArm: + deployment_id: Final = deployment.model_info.id or "" + params: Final = deployment.litellm_params + unknown: Final = CachePredictionArm(deployment_id=deployment_id, model=params.model) + if deployment.model_info.blocked: + return unknown.model_copy(update=MappingProxyType({"reason": "unsupported_deployment_configuration"})) + target: Final = resolve_prediction_target(params) + if isinstance(target, UnsupportedPredictionTarget): + return unknown.model_copy(update=MappingProxyType({"reason": target.reason})) + model: Final = target.model + api_key: Final = target.api_key + total_count: Final = await token_counter(model, api_key, body) + prefix_count: Final = await token_counter(model, api_key, prefix.prefix_body) + if total_count is None or prefix_count is None or total_count < prefix_count: + return unknown.model_copy(update=MappingProxyType({"reason": "token_count_unavailable"})) + scope: Final = cache_scope(caller_key_hash, deployment_id, api_key, model) + checked_at: Final = time.time() if now is None else now + observation: Final = await lookup(cache, scope, prefix, now=checked_at) + exact: Final = observation is not None and observation.fingerprint == prefix.fingerprint + cacheable: Final = observation.cached_tokens if exact and observation is not None else prefix_count + if cacheable > total_count or (observation is not None and observation.cached_tokens > cacheable): + return unknown.model_copy(update=MappingProxyType({"reason": "inconsistent_prefix_token_count"})) + suffix: Final = total_count - cacheable + evidence: Final = ( + CacheEvidence(observed_at=observation.observed_at, expires_at=observation.expires_at) + if observation is not None + else None + ) + if cacheable < get_prompt_cache_min_tokens(params.model): + disabled: Final = _scenario(model, deployment_id, CacheTokenBuckets(uncached_input_tokens=total_count)) + if disabled is None: + return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"})) + return CachePredictionArm( + deployment_id=deployment_id, + model=model, + cache_state="disabled", + reason="below_cache_minimum", + estimate=disabled, + cold=disabled, + warm=disabled, + token_count_source="anthropic_count_tokens", + ) + fresh: Final = observation is not None and observation.expires_at > checked_at + read: Final = observation.cached_tokens if fresh and observation is not None else 0 + with pinned_billing_time(current_billing_time()): + cold: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, 0, prefix.ttl_seconds)) + warm: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, cacheable, prefix.ttl_seconds)) + estimate: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, read, prefix.ttl_seconds)) + if cold is None or warm is None or estimate is None: + return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"})) + return CachePredictionArm( + deployment_id=deployment_id, + model=model, + cache_state="warm" if fresh and exact else "partial" if fresh else "stale" if observation else "unknown", + reason=None if fresh else "observation_expired" if observation else "no_compatible_observation", + estimate=estimate, + cold=cold, + warm=warm, + evidence=evidence, + token_count_source="anthropic_count_tokens", + ) diff --git a/litellm/proxy/common_utils/cache_aware_routing.py b/litellm/proxy/common_utils/cache_aware_routing.py new file mode 100644 index 00000000000..4ae2dce2440 --- /dev/null +++ b/litellm/proxy/common_utils/cache_aware_routing.py @@ -0,0 +1,343 @@ +from __future__ import annotations + +import asyncio +import time +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from itertools import chain +from types import MappingProxyType +from typing import TYPE_CHECKING, Final + +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError + +from litellm._logging import verbose_router_logger +from litellm.caching.dual_cache import DualCache +from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs +from litellm.llms.anthropic.cache_aware_routing import AnthropicCacheRouting, TokenCounter +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import can_key_call_resolved_model +from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig +from litellm.types.router import Deployment, PreRoutingHookResponse +from litellm.types.utils import StandardLoggingRoutingDecision + +if TYPE_CHECKING: + from litellm.router import Router + +_MESSAGES: Final = TypeAdapter(list[Mapping[str, object]]) +_MAPPING: Final = TypeAdapter(Mapping[str, object]) +_DEPLOYMENTS: Final[TypeAdapter[tuple[Deployment, ...] | Deployment]] = TypeAdapter(tuple[Deployment, ...] | Deployment) +_MARKER_OPTIONS: Final = frozenset( + ("model", "complexity_router_config", "rpm", "tpm", "tags", "timeout", "stream_timeout", "num_retries") +) +_CLASSIFIED_CAUSES: Final = frozenset( + { + "heuristic_scorer", + "heuristic_v2", + "reasoning_override", + "llm_classifier", + "llm_v2_classifier", + "jev_classifier", + "capability_classifier", + "heuristic_first_short_circuit", + "hybrid_short_circuit", + "classifier_plugin", + } +) + + +class _ProxyRequest(BaseModel): + model_config = ConfigDict(strict=True) + url: str + body: Mapping[str, JsonValue] + headers: Mapping[str, str] + + +class _CallerSettings(BaseModel): + config: Mapping[str, object] | None = None + + +@dataclass(frozen=True, slots=True) +class CacheAwareChoice: + model: str + tier: str + deployment_id: str + original_cost: float + estimated_cost: float + + +@dataclass(frozen=True, slots=True) +class _Candidate: + model: str + tier: str + deployment: Deployment + + +def eligible_models( + config: ComplexityRouterConfig, decision: StandardLoggingRoutingDecision +) -> tuple[tuple[str, str], ...]: + tier: Final = decision.get("tier") + order: Final = config.tier_names() + tier_entries: Final = chain.from_iterable(config.tier_model_configs.values()) + if ( + tier is None + or tier not in order + or decision.get("cause") not in _CLASSIFIED_CAUSES + or config.has_custom_tiers + or config.plugins + or config.adaptive + or config.session_affinity + or config.classification_mode != "every_request" + or any(entry.litellm_params for entry in tier_entries) + or any(not isinstance(model, str) for model in config.tiers.values()) + ): + return () + floor: Final = order.index(tier) + eligible: Final = tuple( + (name, model) for name, model in config.tiers.items() if isinstance(model, str) and name in order[floor:] + ) + return tuple(entry for index, entry in enumerate(eligible) if entry[1] not in tuple(m for _, m in eligible[:index])) + + +def _candidate(router: Router, tier: str, model: str, request_kwargs: Mapping[str, object]) -> _Candidate | None: + deployments: Final = router.deployments_for_request(model, request_kwargs) + if len(deployments) != 1: + return None + deployment: Final = Deployment.model_validate(deployments[0]) + if deployment.model_info.blocked or not deployment.model_info.id or not AnthropicCacheRouting.supports(deployment): + return None + return _Candidate(model, tier, deployment) + + +async def _available( + candidate: _Candidate, + router: Router, + caller: UserAPIKeyAuth, + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, +) -> bool: + try: + await can_key_call_resolved_model( + model=candidate.model, llm_model_list=router.get_model_list(), valid_token=caller, llm_router=router + ) + healthy: Final = _DEPLOYMENTS.validate_python( + await router.async_get_healthy_deployments( # pyright: ignore[reportUnknownMemberType] # legacy router results are validated at this boundary + model=candidate.model, + messages=_MESSAGES.validate_python(messages) if messages else None, # pyright: ignore[reportArgumentType] # router annotations predate structured native messages + request_kwargs=dict(request_kwargs), # mutable-ok: Router's filtering API accepts a request dictionary + ) + ) + except Exception: # noqa: BLE001 # an unavailable optional candidate must not fail the originally selected route + return False + available: Final = (healthy,) if isinstance(healthy, Deployment) else healthy + return any(entry.model_info.id == candidate.deployment.model_info.id for entry in available) + + +def supported_router_marker(router: Router, alias: str, request_kwargs: Mapping[str, object]) -> bool: + markers: Final = tuple( + Deployment.model_validate(entry) for entry in router.deployments_for_request(alias, request_kwargs) + ) + return bool(markers) and all( + marker.litellm_params.model == "auto_router/complexity_router" + and not frozenset(marker.litellm_params.model_dump(exclude_defaults=True, exclude_none=True)) - _MARKER_OPTIONS + for marker in markers + ) + + +async def select_cached_model( + *, + router: Router, + config: ComplexityRouterConfig, + params_for_model: Callable[[str, str], Mapping[str, object]], + response: PreRoutingHookResponse, + body: Mapping[str, JsonValue], + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, + caller: UserAPIKeyAuth, + cache: DualCache, + counter_for_model: Callable[[str], TokenCounter], + now: float | None = None, +) -> CacheAwareChoice | None: + checked_at: Final = time.time() if now is None else now + decision: Final = response.routing_decision + provider: Final = AnthropicCacheRouting.from_body(body) + if not config.cache_aware_routing or decision is None or provider is None or not caller.api_key: + return None + names: Final = eligible_models(config, decision) + if not names or response.model not in tuple(model for _, model in names): + return None + candidates: Final = tuple( + candidate for tier, model in names if (candidate := _candidate(router, tier, model, request_kwargs)) is not None + ) + original: Final = next((candidate for candidate in candidates if candidate.model == response.model), None) + if original is None: + return None + alternatives: Final = tuple(candidate for candidate in candidates if candidate.model != original.model) + warm_flags: Final = await asyncio.gather( + *(provider.is_warm(candidate.deployment, caller.api_key, cache, checked_at) for candidate in alternatives) + ) + warm: Final = tuple(candidate for candidate, fresh in zip(alternatives, warm_flags) if fresh) + if not warm: + return None + considered: Final = (original, *warm) + availability: Final = await asyncio.gather( + *(_available(candidate, router, caller, request_kwargs, messages) for candidate in considered) + ) + authorized: Final = tuple(candidate for candidate, available in zip(warm, availability[1:]) if available) + if not availability[0] or not authorized: + return None + compared: Final = (original, *authorized) + output_limits: Final = tuple( + params_for_model(candidate.tier, candidate.model).get("max_tokens", provider.requested_output_limit) + for candidate in compared + ) + if any(not isinstance(limit, int) or isinstance(limit, bool) or limit <= 0 for limit in output_limits): + return None + limits: Final = tuple(limit for limit in output_limits if isinstance(limit, int)) + arms: Final = await asyncio.gather( + *( + provider.predict( + candidate.deployment, + caller.api_key, + cache, + counter_for_model(candidate.model), + now=now, + ) + for candidate in compared + ) + ) + costs: Final = tuple( + provider.cost(arm, min(config.cache_aware_routing_output_tokens, limit)) for arm, limit in zip(arms, limits) + ) + original_cost: Final = costs[0] + if original_cost is None: + return None + finished_at: Final = time.time() if now is None else now + qualifying: Final = tuple( + CacheAwareChoice(candidate.model, candidate.tier, arm.deployment_id, original_cost, cost) + for candidate, arm, cost, limit in zip(authorized, arms[1:], costs[1:], limits[1:]) + if cost is not None + and cost < original_cost + and arm.cache_state in ("warm", "partial") + and arm.evidence is not None + and arm.evidence.expires_at > finished_at + and arm.estimate is not None + and provider.fits(candidate.deployment, arm.estimate.tokens.total_tokens, limit) + ) + return min(qualifying, key=lambda choice: choice.estimated_cost, default=None) + + +async def _choose_cached_model( + *, + router: Router, + config: ComplexityRouterConfig, + params_for_model: Callable[[str, str], Mapping[str, object]], + response: PreRoutingHookResponse | None, + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, +) -> CacheAwareChoice | None: + if not config.cache_aware_routing or response is None or response.routing_decision is None: + return None + if not eligible_models(config, response.routing_decision): + return None + from litellm.proxy import proxy_server + from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner + ) + from litellm.router_strategy.complexity_router.context_compaction import compaction_pending + + if ( + proxy_server.llm_router is not router + or router.routing_plugins + or has_request_transforms() + or compaction_pending(request_kwargs) + or not supported_router_marker(router, response.routing_decision.get("router_model_name") or "", request_kwargs) + ): + return None + metadata: Final = _MAPPING.validate_python( + request_kwargs.get(get_metadata_variable_name_from_kwargs(request_kwargs)) or MappingProxyType({}) + ) + caller: Final = metadata.get("user_api_key_auth") + if not isinstance(caller, UserAPIKeyAuth): + return None + settings: Final = _CallerSettings.model_validate(caller, from_attributes=True) + if settings.config: + return None + try: + incoming: Final = _ProxyRequest.model_validate(request_kwargs.get("proxy_server_request")) + except ValidationError: + return None + if any( + request_kwargs.get(key) + for key in ( + "guardrails", + "cache_control_injection_points", + "api_key", + "api_base", + "extra_headers", + "prompt_id", + "mock_response", + "model_info", + "custom_llm_provider", + ) + ): + return None + body: Final = AnthropicCacheRouting.request_body( + incoming.url, incoming.headers, incoming.body, request_kwargs, messages + ) + if body is None: + return None + limiter: Final = proxy_server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter") + if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + return None + + def counter_for_model(model_name: str) -> TokenCounter: + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + try: + async with limiter.request_capacity(caller, model_name, request_data=request_kwargs): + return await AnthropicCacheRouting.count_tokens(model, api_key, body) + except Exception: # noqa: BLE001 # an optional prediction denied capacity is an unavailable estimate + return None + + return count + + return await select_cached_model( + router=router, + config=config, + params_for_model=params_for_model, + response=response, + body=body, + request_kwargs=request_kwargs, + messages=messages, + caller=caller, + cache=proxy_server.proxy_logging_obj.internal_usage_cache.dual_cache, + counter_for_model=counter_for_model, + ) + + +async def choose_cached_model( + *, + router: Router, + config: ComplexityRouterConfig, + params_for_model: Callable[[str, str], Mapping[str, object]], + response: PreRoutingHookResponse | None, + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, +) -> CacheAwareChoice | None: + if not config.cache_aware_routing: + return None + try: + return await asyncio.wait_for( + _choose_cached_model( + router=router, + config=config, + params_for_model=params_for_model, + response=response, + request_kwargs=request_kwargs, + messages=messages, + ), + timeout=config.cache_aware_routing_timeout_ms / 1000, + ) + except Exception: # noqa: BLE001 # cache prediction is optional and must preserve normal routing on failure + verbose_router_logger.debug("Cache-aware routing unavailable; keeping the classified model") + return None diff --git a/litellm/proxy/common_utils/prompt_cache_prediction.py b/litellm/proxy/common_utils/prompt_cache_prediction.py new file mode 100644 index 00000000000..75a9bd40840 --- /dev/null +++ b/litellm/proxy/common_utils/prompt_cache_prediction.py @@ -0,0 +1,20 @@ +from typing import Final + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.llms.anthropic.cache_aware_routing import predict_arm + +__all__: Final = ("has_request_transforms", "predict_arm") + + +def has_request_transforms() -> bool: + from litellm.proxy.hooks import PROXY_HOOKS + + builtins: Final = frozenset(PROXY_HOOKS.values()) + hooks: Final = ("async_pre_call_hook", "async_pre_request_hook", "async_pre_call_deployment_hook") + callbacks: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=CustomLogger) + return any( + type(callback) not in builtins + and any(getattr(type(callback), hook) is not getattr(CustomLogger, hook) for hook in hooks) + for callback in callbacks + ) diff --git a/litellm/proxy/common_utils/prompt_cache_pricing.py b/litellm/proxy/common_utils/prompt_cache_pricing.py index ff070853b46..1ecc3ac44fe 100644 --- a/litellm/proxy/common_utils/prompt_cache_pricing.py +++ b/litellm/proxy/common_utils/prompt_cache_pricing.py @@ -1,4 +1,5 @@ from collections.abc import Mapping +from datetime import datetime, timezone from math import isfinite from typing import Final @@ -20,12 +21,13 @@ def _valid_price(value: object) -> bool: return isinstance(value, (int, float)) and not isinstance(value, bool) and isfinite(value) and value >= 0 -def _has_required_prices(prices: Mapping[str, object], tokens: CacheTokenBuckets) -> bool: +def _has_required_prices(prices: Mapping[str, object], tokens: CacheTokenBuckets, completion_tokens: int = 0) -> bool: required: Final = ( ("input_cost_per_token", True), ("cache_read_input_token_cost", tokens.cache_read_input_tokens > 0), ("cache_creation_input_token_cost", tokens.cache_creation_5m_input_tokens > 0), ("cache_creation_input_token_cost_above_1hr", tokens.cache_creation_1h_input_tokens > 0), + ("output_cost_per_token", completion_tokens > 0), ) if any(needed and not _valid_price(prices.get(key)) for key, needed in required): return False @@ -36,7 +38,9 @@ def _has_required_prices(prices: Mapping[str, object], tokens: CacheTokenBuckets ) -def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> float | None: +def price_cache_tokens( + model: str, deployment_id: str, tokens: CacheTokenBuckets, completion_tokens: int = 0 +) -> float | None: try: selected_model: Final = _select_model_name_for_cost_calc( model=model, @@ -53,12 +57,12 @@ def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets if price_entry is None: return None prices: Final = _PRICE_ENTRY.validate_python(price_entry) - if not _has_required_prices(prices, tokens): + if completion_tokens < 0 or not _has_required_prices(prices, tokens, completion_tokens): return None usage: Final = Usage( prompt_tokens=tokens.total_tokens, - completion_tokens=0, - total_tokens=tokens.total_tokens, + completion_tokens=completion_tokens, + total_tokens=tokens.total_tokens + completion_tokens, prompt_tokens_details=PromptTokensDetailsWrapper( cached_tokens=tokens.cache_read_input_tokens, cache_creation_tokens=tokens.cache_creation_5m_input_tokens + tokens.cache_creation_1h_input_tokens, @@ -73,7 +77,7 @@ def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets messages=[], # mutable-ok: Logging requires a list stream=False, call_type="completion", - start_time=None, + start_time=datetime.now(timezone.utc), litellm_call_id="prompt-cache-prediction", function_id="prompt-cache-prediction", ) @@ -85,7 +89,12 @@ def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets router_model_id=deployment_id, litellm_logging_obj=logging_obj, ) - cost: Final = logging_obj.cost_breakdown.get("input_cost") if logging_obj.cost_breakdown is not None else None - return cost if cost is not None and _valid_price(cost) else None + breakdown: Final = logging_obj.cost_breakdown + input_cost: Final = breakdown.get("input_cost") if breakdown is not None else None + output_cost: Final = breakdown.get("output_cost") if breakdown is not None else None + if input_cost is None or output_cost is None: + return None + cost: Final = input_cost + output_cost + return cost if _valid_price(cost) else None except Exception: # noqa: BLE001 # the shared pricing owners raise plain Exception for unpriceable models return None diff --git a/litellm/proxy/management_endpoints/prompt_cache_prediction.py b/litellm/proxy/management_endpoints/prompt_cache_prediction.py index 56e844214d6..757880980c9 100644 --- a/litellm/proxy/management_endpoints/prompt_cache_prediction.py +++ b/litellm/proxy/management_endpoints/prompt_cache_prediction.py @@ -1,4 +1,3 @@ -import time from collections.abc import Mapping from types import MappingProxyType from typing import Annotated, Final @@ -6,18 +5,10 @@ from typing import Annotated, Final from fastapi import APIRouter, Depends, HTTPException, Request from pydantic import BaseModel, JsonValue, TypeAdapter -import litellm -from litellm._internal_context import current_billing_time, pinned_billing_time -from litellm.caching.caching import DualCache -from litellm.integrations.custom_logger import CustomLogger from litellm.llms.anthropic.prompt_cache_prediction import ( - PromptPrefix, TokenCounter, - UnsupportedPredictionTarget, - cache_scope, count_prompt_tokens, parse_prompt, - resolve_prediction_target, supported_prediction_headers, ) from litellm.proxy._types import UserAPIKeyAuth @@ -27,22 +18,16 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # canonical parsed-body owner; validate its legacy result at the endpoint boundary ) -from litellm.proxy.common_utils.prompt_cache_pricing import price_cache_tokens +from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms, predict_arm from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner ) -from litellm.proxy.hooks.prompt_cache_prediction import lookup from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.types.management_endpoints.prompt_cache_prediction import ( - CacheCostScenario, - CacheEvidence, CachePredictionArm, CachePredictionRequest, CachePredictionResponse, - CacheTokenBuckets, ) -from litellm.types.router import Deployment -from litellm.utils import get_prompt_cache_min_tokens router: Final = APIRouter() _REQUEST_DATA: Final = TypeAdapter(Mapping[str, object]) @@ -52,33 +37,6 @@ class _CallerSettings(BaseModel): config: Mapping[str, object] | None = None -def has_request_transforms() -> bool: - from litellm.proxy.hooks import PROXY_HOOKS - - builtins: Final = frozenset(PROXY_HOOKS.values()) - hooks: Final = ("async_pre_call_hook", "async_pre_request_hook", "async_pre_call_deployment_hook") - callbacks: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=CustomLogger) - return any( - type(callback) not in builtins - and any(getattr(type(callback), hook) is not getattr(CustomLogger, hook) for hook in hooks) - for callback in callbacks - ) - - -def _buckets(prefix_tokens: int, suffix_tokens: int, read_tokens: int, ttl_seconds: int) -> CacheTokenBuckets: - return CacheTokenBuckets( - uncached_input_tokens=suffix_tokens, - cache_read_input_tokens=read_tokens, - cache_creation_5m_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 300 else 0, - cache_creation_1h_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 3600 else 0, - ) - - -def _scenario(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> CacheCostScenario | None: - cost: Final = price_cache_tokens(model=model, deployment_id=deployment_id, tokens=tokens) - return CacheCostScenario(tokens=tokens, input_cost=cost) if cost is not None else None - - def _capacity_counter( limiter: _PROXY_MaxParallelRequestsHandler_v3, caller: UserAPIKeyAuth, @@ -103,75 +61,6 @@ def _capacity_request_data( return MappingProxyType(data) -async def predict_arm( - deployment: Deployment, - body: Mapping[str, JsonValue], - prefix: PromptPrefix, - caller_key_hash: str, - cache: DualCache, - token_counter: TokenCounter, -) -> CachePredictionArm: - deployment_id: Final = deployment.model_info.id or "" - params: Final = deployment.litellm_params - unknown: Final = CachePredictionArm(deployment_id=deployment_id, model=params.model) - if deployment.model_info.blocked: - return unknown.model_copy(update=MappingProxyType({"reason": "unsupported_deployment_configuration"})) - target: Final = resolve_prediction_target(params) - if isinstance(target, UnsupportedPredictionTarget): - return unknown.model_copy(update=MappingProxyType({"reason": target.reason})) - model: Final = target.model - api_key: Final = target.api_key - total_count: Final = await token_counter(model, api_key, body) - prefix_count: Final = await token_counter(model, api_key, prefix.prefix_body) - if total_count is None or prefix_count is None or total_count < prefix_count: - return unknown.model_copy(update=MappingProxyType({"reason": "token_count_unavailable"})) - scope: Final = cache_scope(caller_key_hash, deployment_id, api_key, model) - observation: Final = await lookup(cache, scope, prefix) - exact: Final = observation is not None and observation.fingerprint == prefix.fingerprint - cacheable: Final = observation.cached_tokens if exact and observation is not None else prefix_count - if cacheable > total_count or (observation is not None and observation.cached_tokens > cacheable): - return unknown.model_copy(update=MappingProxyType({"reason": "inconsistent_prefix_token_count"})) - suffix: Final = total_count - cacheable - evidence: Final = ( - CacheEvidence(observed_at=observation.observed_at, expires_at=observation.expires_at) - if observation is not None - else None - ) - if cacheable < get_prompt_cache_min_tokens(params.model): - disabled: Final = _scenario(model, deployment_id, CacheTokenBuckets(uncached_input_tokens=total_count)) - if disabled is None: - return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"})) - return CachePredictionArm( - deployment_id=deployment_id, - model=model, - cache_state="disabled", - reason="below_cache_minimum", - estimate=disabled, - cold=disabled, - warm=disabled, - token_count_source="anthropic_count_tokens", - ) - fresh: Final = observation is not None and observation.expires_at > time.time() - read: Final = observation.cached_tokens if fresh and observation is not None else 0 - with pinned_billing_time(current_billing_time()): - cold: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, 0, prefix.ttl_seconds)) - warm: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, cacheable, prefix.ttl_seconds)) - estimate: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, read, prefix.ttl_seconds)) - if cold is None or warm is None or estimate is None: - return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"})) - return CachePredictionArm( - deployment_id=deployment_id, - model=model, - cache_state="warm" if fresh and exact else "partial" if fresh else "stale" if observation else "unknown", - reason=None if fresh else "observation_expired" if observation else "no_compatible_observation", - estimate=estimate, - cold=cold, - warm=warm, - evidence=evidence, - token_count_source="anthropic_count_tokens", - ) - - @router.post( "/cost/predict-cache", tags=["Cost Tracking"], # mutable-ok: FastAPI requires a list for OpenAPI tags diff --git a/litellm/router_strategy/complexity_router/README.md b/litellm/router_strategy/complexity_router/README.md index f023d5001d9..f55362b6c41 100644 --- a/litellm/router_strategy/complexity_router/README.md +++ b/litellm/router_strategy/complexity_router/README.md @@ -1,6 +1,6 @@ # Complexity Router -A rule-based routing strategy that classifies requests by complexity and routes them to appropriate models - with zero API calls and sub-millisecond latency. +A routing strategy that classifies requests by complexity and routes them to appropriate models. The default rule-based classifier scores requests locally. Optional classifiers and cache-aware routing can make provider calls ## Overview @@ -68,6 +68,45 @@ still resolve to a deployment in `model_list`; this configuration does not creat - abc ``` +### Opt in to prompt-cache costs + +Set `cache_aware_routing: true` to consider observed prompt-cache savings after classification. This is disabled by default. A warm model in the same or a higher tier can replace the classified model when its estimated input and output cost is strictly lower. Cache savings never lower the required tier + +```yaml +model_list: + - model_name: smart-router + litellm_params: + model: auto_router/complexity_router + complexity_router_config: + cache_aware_routing: true + cache_aware_routing_output_tokens: 1024 + cache_aware_routing_timeout_ms: 2000 + context_compaction: false + tiers: + SIMPLE: haiku + COMPLEX: sonnet + - model_name: haiku + litellm_params: + model: anthropic/claude-haiku-4-5 + api_key: os.environ/ANTHROPIC_API_KEY + - model_name: sonnet + litellm_params: + model: anthropic/claude-sonnet-5 + api_key: os.environ/ANTHROPIC_API_KEY +``` + +This first version supports the proxy's native `POST /v1/messages` endpoint with Anthropic, text and client tools, and one explicit message-content `cache_control` breakpoint. Each tier must name one model group with one deployment. The default v3 rate limiter must be enabled. It uses the same observations and token counting as `/cost/predict-cache`; it does not prewarm caches or enable provider caching on the application's behalf + +The proxy must have observed a successful cache read or write for the candidate's matching prefix, under the same caller key, deployment, provider key and model. A fresh observation allows a cache discount; missing or expired evidence does not. Provider eviction can still turn an expected hit into a miss + +The comparison includes uncached input, cache writes at the requested TTL, cache reads and expected output tokens. Set `cache_aware_routing_output_tokens` to your workload's expected response length; it defaults to 1024 and is capped separately by each model's effective output limit. With `max_tokens_from_tier_model: true` (the default), this is the model's known output ceiling; when disabled or unknown, the caller's `max_tokens` applies. The full effective output limit, together with the counted input, must fit the candidate's known limits. Custom deployment prices are respected + +Prediction makes up to two token-count requests per compared model. These use rate and concurrency capacity and add latency. The default total timeout is two seconds; timeout, missing counts or prices, and prediction failures preserve the classified route. No provider count requests run when there is no warm eligible alternative + +Session affinity, user-turn classification, adaptive routing, routing plugins, custom tier ladders, tier pools and per-tier parameter overrides keep their existing behavior without a cache adjustment. The same applies to unsupported providers or prompt shapes, beta headers, custom provider endpoints, request transforms, and pending context compaction. Disable context compaction as in the example so it cannot rewrite the predicted prompt. Alias markers should contain only routing configuration and rate, timeout or tag settings + +When cache costs change the model, the routing decision reports `cause: prompt_cache_cost`. Its signals include the original model, classification cause and both estimated costs + ### Capability forecasting Set `classifier_type: capability` to use diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 9df6306436b..76ee977bf28 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -4390,10 +4390,13 @@ class ComplexityRouter(CustomLogger): resolved_messages=resolved_messages, context_fit=context_fit, ) + cache_adjusted_response: Final = await self._apply_prompt_cache_routing( + routed_response, messages, request_kwargs, context_fit + ) response: Final = ( await self._gate_response_health( await self._gate_response_modality( - routed_response, messages, resolved_messages, request_kwargs, context_fit + cache_adjusted_response, messages, resolved_messages, request_kwargs, context_fit ), messages, input, @@ -4401,7 +4404,7 @@ class ComplexityRouter(CustomLogger): request_kwargs, context_fit, ) - if routed_response is not None + if cache_adjusted_response is not None else None ) # Sentinel presence, not the plan_mode cause, gates the pin write: a plan-mode turn @@ -4425,6 +4428,53 @@ class ComplexityRouter(CustomLogger): ) return self._with_session_deployment_affinity(response) + async def _apply_prompt_cache_routing( + self, + response: PreRoutingHookResponse | None, + messages: Sequence[Mapping[str, object]] | None, + request_kwargs: Mapping[str, object], + context_fit: _RequestContextFit, + ) -> PreRoutingHookResponse | None: + if not self.config.cache_aware_routing or response is None or response.routing_decision is None: + return response + from litellm.proxy.common_utils.cache_aware_routing import choose_cached_model + + choice: Final = await choose_cached_model( + router=self.litellm_router_instance, + config=self.config, + params_for_model=self._litellm_params_for_model, + response=response, + request_kwargs=request_kwargs, + messages=messages, + ) + if choice is None or not context_fit.accepts(choice.model): + return response + params: Final = self._litellm_params_for_model(choice.tier, choice.model) + decision: Final[StandardLoggingRoutingDecision] = { + **response.routing_decision, + "routed_model": choice.model, + "cause": "prompt_cache_cost", + "tier": choice.tier, + "tier_label": (self.config.tier_labels or {}).get(choice.tier, choice.tier), + "tier_litellm_params": params, + "signals": ( + *(response.routing_decision.get("signals") or ()), + f"cache-aware:classified-model={response.model}", + f"cache-aware:classification-cause={response.routing_decision.get('cause')}", + f"cache-aware:estimated-cost={choice.estimated_cost:.8f};original-cost={choice.original_cost:.8f}", + ), + } + verbose_router_logger.info( + "ComplexityRouter: cache-aware choice model=%s original=%s estimated_cost=%s original_cost=%s", + choice.model, + response.model, + choice.estimated_cost, + choice.original_cost, + ) + return response.model_copy( + update=MappingProxyType({"model": choice.model, "litellm_params": params, "routing_decision": decision}) + ) + async def _classify_and_route( self, model: str, diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index dbc70631298..e0427f89fe3 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -1422,6 +1422,25 @@ class ComplexityRouterConfig(BaseModel): ), ) + cache_aware_routing: bool = Field( + default=False, + description=( + "Opt in to comparing prompt-cache costs after classification. On supported native Anthropic proxy requests, " + "an already warm model in the same or a higher tier may replace the classified model when its estimated " + "input and output cost is lower. Unsupported requests and unavailable estimates keep ordinary routing." + ), + ) + cache_aware_routing_output_tokens: int = Field( + default=1024, + ge=0, + description="Expected output tokens used in cache-aware cost comparisons; capped by each model's effective output limit.", + ) + cache_aware_routing_timeout_ms: int = Field( + default=2000, + gt=0, + description="Total time budget for cache-aware predictions; expiry preserves the original routing decision.", + ) + # Session affinity: pin the first turn's routed model for the rest of the session session_affinity: bool = Field( default=False, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index f8b57139b37..4bda1dd53ce 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2970,6 +2970,7 @@ class StandardLoggingRoutingDecisionTierBoundaries(TypedDict): RoutingDecisionCause = Literal[ + "prompt_cache_cost", "heuristic_scorer", "heuristic_v2", # The scorer found 2+ reasoning markers and forced REASONING regardless of score. diff --git a/tests/unit/proxy/common_utils/test_cache_aware_routing.py b/tests/unit/proxy/common_utils/test_cache_aware_routing.py new file mode 100644 index 00000000000..000d6773d04 --- /dev/null +++ b/tests/unit/proxy/common_utils/test_cache_aware_routing.py @@ -0,0 +1,588 @@ +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from typing import Final + +import pytest +from pydantic import JsonValue + +from litellm import Router +from litellm.caching.dual_cache import DualCache +from litellm.llms.anthropic.prompt_cache_prediction import TokenCounter, cache_scope, parse_prompt +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.cache_aware_routing import ( + CacheAwareChoice, + choose_cached_model, + eligible_models, + select_cached_model, +) +from litellm.proxy.hooks.prompt_cache_prediction import CacheObservation, _cache_key +from litellm.router_strategy.complexity_router.complexity_router import ComplexityRouter +from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig +from litellm.types.router import PreRoutingHookResponse + +_CALLER: Final = "test-cache-aware-caller" +_PROVIDER_KEY: Final = "test-cache-aware-provider" +_NOW: Final = 1000.0 + + +@dataclass(frozen=True, slots=True) +class _Counts: + total: int | None = 51000 + prefix: int | None = 50000 + + async def __call__(self, model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + return self.total if "max_tokens" in body else self.prefix + + +def _counter_for_model(model: str) -> TokenCounter: + return _Counts() + + +def _forbidden_counter(model: str) -> TokenCounter: + raise AssertionError("No provider counts should run without a warm eligible alternative") + + +def _body(text: str = "Stable cached context") -> dict[str, JsonValue]: + return { + "max_tokens": 20000, + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": text, "cache_control": {"type": "ephemeral"}}, + {"type": "text", "text": "What is 2 + 2?"}, + ], + } + ], + } + + +def _router( + strong_output_rate: float = 0.000015, *, free: bool = False, cheap_limit: int = 30000, strong_limit: int = 30000 +) -> Router: + return Router( + model_list=[ + { + "model_name": "cheap", + "litellm_params": { + "model": "anthropic/claude-haiku-4-5", + "api_key": _PROVIDER_KEY, + "input_cost_per_token": 0 if free else 0.000001, + "output_cost_per_token": 0 if free else 0.000004, + "cache_read_input_token_cost": 0 if free else 0.0000001, + "cache_creation_input_token_cost": 0 if free else 0.00000125, + }, + "model_info": {"id": "test-cache-cheap", "max_input_tokens": 100000, "max_output_tokens": cheap_limit}, + }, + { + "model_name": "strong", + "litellm_params": { + "model": "anthropic/claude-sonnet-5", + "api_key": _PROVIDER_KEY, + "input_cost_per_token": 0 if free else 0.000003, + "output_cost_per_token": 0 if free else strong_output_rate, + "cache_read_input_token_cost": 0 if free else 0.0000003, + "cache_creation_input_token_cost": 0 if free else 0.00000375, + }, + "model_info": { + "id": "test-cache-strong", + "max_input_tokens": 100000, + "max_output_tokens": strong_limit, + }, + }, + ] + ) + + +def _config(**overrides: object) -> ComplexityRouterConfig: + return ComplexityRouterConfig.model_validate( + { + "tiers": {"SIMPLE": "cheap", "COMPLEX": "strong"}, + "cache_aware_routing": True, + **overrides, + } + ) + + +def _response(tier: str = "SIMPLE", model: str = "cheap") -> PreRoutingHookResponse: + return PreRoutingHookResponse( + model=model, + messages=None, + routing_decision={ + "router_model_name": "smart", + "router_type": "complexity", + "routed_model": model, + "tier": tier, + "cause": "heuristic_scorer", + }, + ) + + +async def _observed(cache: DualCache, *, caller: str = _CALLER, expires_at: float = 1290.0) -> None: + prefix: Final = parse_prompt(_body()) + assert prefix is not None + scope: Final = cache_scope(caller, "test-cache-strong", _PROVIDER_KEY, "claude-sonnet-5") + observation: Final = CacheObservation( + fingerprint=prefix.fingerprint, + cached_tokens=50000, + observed_at=990.0, + expires_at=expires_at, + ) + await cache.async_set_cache(_cache_key(scope, prefix.fingerprint), observation.model_dump_json(), ttl=3600) + + +async def _select( + *, + router: Router, + config: ComplexityRouterConfig, + response: PreRoutingHookResponse, + body: Mapping[str, JsonValue], + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, + caller: UserAPIKeyAuth, + cache: DualCache, + counter_for_model: Callable[[str], TokenCounter], + now: float, +) -> CacheAwareChoice | None: + complexity: Final = ComplexityRouter("smart", router, config.model_dump()) + return await select_cached_model( + router=router, + config=config, + params_for_model=complexity._litellm_params_for_model, + response=response, + body=body, + request_kwargs=request_kwargs, + messages=messages, + caller=caller, + cache=cache, + counter_for_model=counter_for_model, + now=now, + ) + + +@pytest.mark.asyncio +async def test_warm_stronger_model_wins_after_counting_input_and_output_cost() -> None: + cache: Final = DualCache() + await _observed(cache) + choice: Final = await _select( + router=_router(), + config=_config(), + response=_response(), + body=_body(), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + assert choice is not None + assert (choice.model, choice.tier, choice.deployment_id) == ("strong", "COMPLEX", "test-cache-strong") + assert choice.original_cost == pytest.approx(50000 * 0.00000125 + 1000 * 0.000001 + 1024 * 0.000004) + assert choice.estimated_cost == pytest.approx(50000 * 0.0000003 + 1000 * 0.000003 + 1024 * 0.000015) + + +@pytest.mark.asyncio +async def test_output_price_can_outweigh_the_cache_saving() -> None: + cache: Final = DualCache() + await _observed(cache) + choice: Final = await _select( + router=_router(strong_output_rate=0.001), + config=_config(), + response=_response(), + body=_body(), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + assert choice is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("case", ["missing", "expired", "different_caller", "changed_prefix", "unauthorized"]) +async def test_no_cache_discount_without_fresh_authorized_matching_evidence(case: str) -> None: + cache: Final = DualCache() + if case != "missing": + await _observed( + cache, + caller="someone-else" if case == "different_caller" else _CALLER, + expires_at=999.0 if case == "expired" else 1290.0, + ) + choice: Final = await _select( + router=_router(), + config=_config(), + response=_response(), + body=_body("Changed context" if case == "changed_prefix" else "Stable cached context"), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap"] if case == "unauthorized" else ["cheap", "strong"]), + cache=cache, + counter_for_model=_forbidden_counter, + now=_NOW, + ) + assert choice is None + + +@pytest.mark.asyncio +async def test_disabled_setting_does_not_access_prediction_services() -> None: + config: Final = ComplexityRouterConfig(tiers={"SIMPLE": "cheap"}) + assert config.cache_aware_routing is False + assert ( + await choose_cached_model( + router=_router(), + config=config, + params_for_model=ComplexityRouter("smart", _router(), config.model_dump())._litellm_params_for_model, + response=_response(), + request_kwargs={}, + messages=None, + ) + is None + ) + + +def test_cache_prices_cannot_add_a_model_below_the_classified_tier() -> None: + response: Final = _response("COMPLEX", "strong") + assert response.routing_decision is not None + assert eligible_models(_config(), response.routing_decision) == (("COMPLEX", "strong"),) + + +@pytest.mark.parametrize( + "overrides", [{"adaptive": True}, {"session_affinity": True}, {"classification_mode": "user_turn"}] +) +def test_existing_pinned_or_adaptive_policies_are_preserved(overrides: Mapping[str, object]) -> None: + response: Final = _response() + assert response.routing_decision is not None + assert eligible_models(_config(**overrides), response.routing_decision) == () + + +@pytest.mark.parametrize("total,prefix", [(None, 50000), (51000, None), (1000, 50000)]) +@pytest.mark.asyncio +async def test_unavailable_or_inconsistent_counts_keep_the_classified_model( + total: int | None, prefix: int | None +) -> None: + cache: Final = DualCache() + await _observed(cache) + assert ( + await _select( + router=_router(), + config=_config(), + response=_response(), + body=_body(), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=lambda _: _Counts(total, prefix), + now=_NOW, + ) + is None + ) + + +@pytest.mark.asyncio +async def test_output_estimate_is_capped_by_the_requested_limit() -> None: + cache: Final = DualCache() + await _observed(cache) + choice: Final = await _select( + router=_router(), + config=_config(cache_aware_routing_output_tokens=100000, max_tokens_from_tier_model=False), + response=_response(), + body={**_body(), "max_tokens": 1}, + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + assert choice is not None + assert choice.estimated_cost == pytest.approx(50000 * 0.0000003 + 1000 * 0.000003 + 0.000015) + + +@pytest.mark.asyncio +async def test_warm_model_that_cannot_fit_the_request_is_not_selected() -> None: + cache: Final = DualCache() + await _observed(cache) + assert ( + await _select( + router=_router(), + config=_config(max_tokens_from_tier_model=False), + response=_response(), + body={**_body(), "max_tokens": 100000000}, + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + is None + ) + + +def test_repeated_model_in_multiple_tiers_is_only_considered_once() -> None: + decision: Final = _response().routing_decision + assert decision is not None + assert eligible_models(_config(tiers={"SIMPLE": "cheap", "MEDIUM": "strong", "COMPLEX": "strong"}), decision) == ( + ("SIMPLE", "cheap"), + ("MEDIUM", "strong"), + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "enabled,behavior,expected", + [ + (False, "success", "cheap"), + (True, "success", "strong"), + (True, "tier_cost", "cheap"), + (True, "tier_context", "cheap"), + (True, "error", "cheap"), + (True, "deadline", "cheap"), + (True, "cancel", None), + (True, "transformed", "cheap"), + (True, "unsupported_shape", "cheap"), + (True, "custom_endpoint", "cheap"), + (True, "compaction", "cheap"), + (True, "guardrail", "cheap"), + ], +) +async def test_router_applies_opt_in_and_preserves_failure_semantics( + monkeypatch: pytest.MonkeyPatch, enabled: bool, behavior: str, expected: str | None +) -> None: + import asyncio + import json + + import httpx + + import litellm + from litellm.caching.llm_caching_handler import LLMClientCache + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.proxy import proxy_server + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.utils import ProxyLogging + from litellm.router_strategy.complexity_router.context_compaction import initialize_compaction_state + + config: Final = _config( + cache_aware_routing=enabled, cache_aware_routing_timeout_ms=1 if behavior == "deadline" else 2000 + ) + models: Final = _router( + strong_output_rate=0.000048 if behavior == "tier_cost" else 0.000015, + cheap_limit=50 if behavior == "tier_cost" else 30000, + strong_limit=60000 if behavior == "tier_context" else 30000, + ).get_model_list() + assert models is not None + router: Final = Router( + model_list=[ + *models, + { + "model_name": "smart", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": config.model_dump(), + **({"temperature": 0.1} if behavior == "transformed" else {}), + }, + }, + ] + ) + logging: Final = ProxyLogging(UserApiKeyCache()) + logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3( + logging.internal_usage_cache + ) + await _observed(logging.internal_usage_cache.dual_cache, expires_at=1e100) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging) + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache()) + + requests: Final = asyncio.Queue[httpx.Request]() + + async def count(request: httpx.Request) -> httpx.Response: + requests.put_nowait(request) + assert enabled and behavior not in ("transformed", "unsupported_shape", "custom_endpoint") + if behavior == "error": + return httpx.Response(503, json={"error": "Provider unavailable"}) + if behavior == "cancel": + raise asyncio.CancelledError() + if behavior == "deadline": + await asyncio.Future() + payload: Final = json.loads(request.content) + assert request.url == "https://api.anthropic.com/v1/messages/count_tokens" + assert request.headers["x-api-key"] == _PROVIDER_KEY + return httpx.Response(200, json={"input_tokens": 51000 if "What is 2 + 2?" in str(payload) else 50000}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(count)) as client: + handler: Final = AsyncHTTPHandler() + await handler.client.aclose() + handler.client = client + litellm.in_memory_llm_clients_cache.set_cache("async_httpx_clientanthropic", handler) + body: Final = { + **_body(), + **({"max_tokens": 1000} if behavior in ("tier_cost", "tier_context") else {}), + **({"thinking": {"type": "enabled", "budget_tokens": 10000}} if behavior == "unsupported_shape" else {}), + } + kwargs: Final = { + "litellm_metadata": { + "user_api_key_auth": UserAPIKeyAuth(api_key=_CALLER, models=["smart", "cheap", "strong"]) + }, + "proxy_server_request": {"url": "http://localhost/v1/messages", "body": body, "headers": {}}, + **({"api_base": "https://custom.example"} if behavior == "custom_endpoint" else {}), + **( + {"_context_compaction_state": initialize_compaction_state({}, "messages")} + if behavior == "compaction" + else {} + ), + **({"guardrails": ["test-guardrail"]} if behavior == "guardrail" else {}), + } + if expected is None: + with pytest.raises(asyncio.CancelledError): + await router.async_pre_routing_hook(model="smart", request_kwargs=kwargs, messages=body["messages"]) + return + response: Final = await router.async_pre_routing_hook( + model="smart", request_kwargs=kwargs, messages=body["messages"] + ) + assert response is not None + assert response.model == expected + if not enabled or behavior in ( + "transformed", + "unsupported_shape", + "custom_endpoint", + "compaction", + "guardrail", + ): + assert requests.qsize() == 0 + if enabled and behavior == "success": + assert requests.qsize() == 4 + assert response.routing_decision is not None + assert response.routing_decision["cause"] == ( + "prompt_cache_cost" if expected == "strong" else "heuristic_scorer" + ) + + +@pytest.mark.asyncio +async def test_equal_costs_keep_the_classified_model() -> None: + cache: Final = DualCache() + await _observed(cache) + assert ( + await _select( + router=_router(free=True), + config=_config(), + response=_response(), + body=_body(), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + is None + ) + + +@pytest.mark.parametrize( + "cause", ["llm_v2_classifier", "capability_classifier", "heuristic_first_short_circuit", "hybrid_short_circuit"] +) +def test_successful_classifiers_can_consider_cache_costs(cause: str) -> None: + response: Final = PreRoutingHookResponse.model_validate( + { + "model": "cheap", + "messages": None, + "routing_decision": {"tier": "SIMPLE", "cause": cause}, + } + ) + assert response.routing_decision is not None + assert eligible_models(_config(), response.routing_decision) == (("SIMPLE", "cheap"), ("COMPLEX", "strong")) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "cheap_limit,strong_limit,requested,from_tier,output_rate,expected_limits", + [ + (50, 1000, 1000, True, 0.000048, None), + (30000, 30000, 1, True, 0.000049, None), + (30000, 60000, 1, True, 0.000015, None), + (50, 100, 20000, True, 0.000048, (50, 100)), + (30000, 30000, 1, False, 0.000048, (1, 1)), + ], +) +async def test_each_candidate_uses_its_effective_routed_output_limit( + cheap_limit: int, + strong_limit: int, + requested: int, + from_tier: bool, + output_rate: float, + expected_limits: tuple[int, int] | None, +) -> None: + cache: Final = DualCache() + await _observed(cache) + choice: Final = await _select( + router=_router(strong_output_rate=output_rate, cheap_limit=cheap_limit, strong_limit=strong_limit), + config=_config(max_tokens_from_tier_model=from_tier), + response=_response(), + body={**_body(), "max_tokens": requested}, + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + if expected_limits is None: + assert choice is None + return + assert choice is not None + assert choice.original_cost == pytest.approx(50000 * 0.00000125 + 1000 * 0.000001 + expected_limits[0] * 0.000004) + assert choice.estimated_cost == pytest.approx( + 50000 * 0.0000003 + 1000 * 0.000003 + expected_limits[1] * output_rate + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("warm,authorized", [(False, True), (True, True), (True, False)]) +async def test_authorization_only_runs_for_original_and_warm_alternatives_before_provider_counts( + monkeypatch: pytest.MonkeyPatch, warm: bool, authorized: bool +) -> None: + from unittest.mock import AsyncMock + + from litellm.proxy.common_utils import cache_aware_routing + + cache: Final = DualCache() + if warm: + await _observed(cache) + models: Final = _router().get_model_list() + assert models is not None + router: Final = Router( + model_list=[ + *models, + {**models[0], "model_name": "cold", "model_info": {"id": "test-cache-cold"}}, + ] + ) + authorization: Final = AsyncMock(wraps=cache_aware_routing.can_key_call_resolved_model) + monkeypatch.setattr(cache_aware_routing, "can_key_call_resolved_model", authorization) + + def counter_for_model(model: str) -> TokenCounter: + assert warm and authorized + assert authorization.await_count == 2 + return _Counts() + + choice: Final = await _select( + router=router, + config=_config(tiers={"SIMPLE": "cheap", "MEDIUM": "cold", "COMPLEX": "strong"}), + response=_response(), + body=_body(), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong", "cold"] if authorized else ["cheap", "cold"]), + cache=cache, + counter_for_model=counter_for_model, + now=_NOW, + ) + assert (choice is not None) == (warm and authorized) + assert tuple(call.kwargs["model"] for call in authorization.await_args_list) == ( + ("cheap", "strong") if warm else () + ) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index b5d515bd1fb..6d45f6e691c 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -39483,6 +39483,24 @@ export interface components { adaptive_eligible: "all" | "classified_tier"; /** @description Quality vs cost weights for adaptive selection (used when adaptive=True) */ adaptive_weights?: components["schemas"]["AdaptiveRouterWeights"]; + /** + * Cache Aware Routing + * @description Opt in to comparing prompt-cache costs after classification. On supported native Anthropic proxy requests, an already warm model in the same or a higher tier may replace the classified model when its estimated input and output cost is lower. Unsupported requests and unavailable estimates keep ordinary routing. + * @default false + */ + cache_aware_routing: boolean; + /** + * Cache Aware Routing Output Tokens + * @description Expected output tokens used in cache-aware cost comparisons; capped by each model's effective output limit. + * @default 1024 + */ + cache_aware_routing_output_tokens: number; + /** + * Cache Aware Routing Timeout Ms + * @description Total time budget for cache-aware predictions; expiry preserves the original routing decision. + * @default 2000 + */ + cache_aware_routing_timeout_ms: number; /** @description Probability threshold policy required when classifier_type is 'capability'. The classifier forecasts p_solve for efficient_tier, adjusts base_threshold using the capability-card boundary, and otherwise routes to capable_tier */ capability_classifier_config?: components["schemas"]["CapabilityClassifierConfig"] | null; /** @@ -42605,7 +42623,7 @@ export interface components { * Cause * @enum {string} */ - cause?: "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "capability_classifier" | "jev_classifier" | "llm_v2_classifier" | "llm_v2_fallback" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "capability_classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "health_default_fallback" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit"; + cause?: "prompt_cache_cost" | "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "capability_classifier" | "jev_classifier" | "llm_v2_classifier" | "llm_v2_fallback" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "capability_classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "health_default_fallback" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit"; /** Classifier Calibrated Capable P Solve */ classifier_calibrated_capable_p_solve?: number; /** Classifier Calibrated Efficient P Solve */ From 7aba77197dc53737f8e882bfceab493397a424b0 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 18:15:45 -0700 Subject: [PATCH 08/65] feat(otel): add SigNoz preset for OpenTelemetry v2 (#43296) * feat(otel): add SigNoz preset for OpenTelemetry v2 Adds the signoz callback (OTLP/HTTP exporter, GenAI vocabulary, key and team level dynamic ingestion endpoint and key) as an OpenTelemetry v2 preset, with the preset factory accepting the allow_missing_credentials kwarg the V2 registry always passes so construction no longer falls back silently to legacy OpenTelemetry. Ships the deterministic tests/integration/observability/test_signoz_delivery.py audit suite Absorbs the work from https://github.com/BerriAI/litellm/pull/38206 Co-authored-by: Nagesh Bansal Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(otel): drop explanatory comments from the SigNoz preset Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(types): keep signoz dynamic param lines within ruff format width Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(signoz): assert the missing-endpoint boot path directly instead of in an except block Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(ui): regenerate schema.d.ts for the signoz health service Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(otel): allowlist SigNoz key/team endpoints and route keyless collectors without the operator key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(otel): terminate the SigNoz shutdown cell before the flush and drop test docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(otel): keep the shared tenant routing untouched and require an ingestion key for SigNoz key/team endpoints Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(otel): warn about a keyless SigNoz team endpoint from the header resolver so the shared cache actually reaches it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Nagesh Bansal Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/__init__.py | 1 + litellm/integrations/callback_configs.json | 21 + litellm/integrations/otel/model/config.py | 1 + litellm/integrations/otel/presets/__init__.py | 9 + litellm/integrations/otel/presets/signoz.py | 95 ++ .../custom_logger_registry.py | 1 + .../initialize_dynamic_callback_params.py | 6 + litellm/litellm_core_utils/litellm_logging.py | 32 + .../_experimental/out/assets/logos/signoz.svg | 1 + litellm/proxy/_types.py | 6 + .../health_endpoints/_health_endpoints.py | 2 + litellm/proxy/litellm_pre_call_utils.py | 2 + litellm/types/utils.py | 3 + .../observability/test_signoz_delivery.py | 980 ++++++++++++++++++ .../proxy/test_litellm_pre_call_utils.py | 28 + .../integrations/otel/test_otel_v2_dynamic.py | 66 ++ .../integrations/otel/test_otel_v2_presets.py | 66 ++ .../test_litellm_logging.py | 97 ++ .../public/assets/logos/signoz.svg | 1 + .../src/components/callback_info_helpers.tsx | 12 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +- 21 files changed, 1431 insertions(+), 1 deletion(-) create mode 100644 litellm/integrations/otel/presets/signoz.py create mode 100644 litellm/proxy/_experimental/out/assets/logos/signoz.svg create mode 100644 tests/integration/observability/test_signoz_delivery.py create mode 100644 ui/litellm-dashboard/public/assets/logos/signoz.svg diff --git a/litellm/__init__.py b/litellm/__init__.py index 5a7d6e8125d..5d10737e876 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -172,6 +172,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "levo", "compression_interception", "newrelic", + "signoz", ] cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None diff --git a/litellm/integrations/callback_configs.json b/litellm/integrations/callback_configs.json index 4e72075dc5c..190c283d087 100644 --- a/litellm/integrations/callback_configs.json +++ b/litellm/integrations/callback_configs.json @@ -502,6 +502,27 @@ }, "description": "S3 Bucket (AWS) Logging Integration" }, + { + "id": "signoz", + "displayName": "SigNoz", + "logo": "signoz.svg", + "supports_key_team_logging": true, + "dynamic_params": { + "signoz_ingestion_endpoint": { + "type": "text", + "ui_name": "SigNoz Ingestion Endpoint", + "description": "Ingestion endpoint for this team, e.g. https://ingest.us.signoz.cloud:443 for SigNoz Cloud or your own collector. Leave blank to use the proxy's configured endpoint. Regions: https://signoz.io/docs/ingestion/signoz-cloud/overview/", + "required": false + }, + "signoz_ingestion_key": { + "type": "password", + "ui_name": "SigNoz Ingestion Key (optional)", + "description": "Ingestion key for this team, so its traces land in its own SigNoz account. Not needed for self-hosted SigNoz. Keys: https://signoz.io/docs/ingestion/signoz-cloud/keys/", + "required": false + } + }, + "description": "SigNoz Logging Integration. Setup: https://signoz.io/docs/litellm-observability/" + }, { "id": "sqs", "displayName": "SQS", diff --git a/litellm/integrations/otel/model/config.py b/litellm/integrations/otel/model/config.py index 5447a8ee80a..5a3965862e0 100644 --- a/litellm/integrations/otel/model/config.py +++ b/litellm/integrations/otel/model/config.py @@ -41,6 +41,7 @@ class ExporterOwner(str, Enum): LEVO = "levo" AGENTOPS = "agentops" NEWRELIC = "newrelic" + SIGNOZ = "signoz" class _OTelV2Flag(BaseSettings): diff --git a/litellm/integrations/otel/presets/__init__.py b/litellm/integrations/otel/presets/__init__.py index a0cd5b3fd98..7c891c29409 100644 --- a/litellm/integrations/otel/presets/__init__.py +++ b/litellm/integrations/otel/presets/__init__.py @@ -30,6 +30,11 @@ from litellm.integrations.otel.presets.phoenix import ( phoenix_preset, phoenix_project_headers, ) +from litellm.integrations.otel.presets.signoz import ( + signoz_dynamic_endpoint, + signoz_dynamic_headers, + signoz_preset, +) from litellm.integrations.otel.presets.weave import weave_dynamic_headers, weave_preset from litellm.types.utils import StandardCallbackDynamicParams @@ -44,6 +49,7 @@ PRESET_BY_CALLBACK: Final[Mapping[str, Preset]] = MappingProxyType( "langtrace": langtrace_preset, "levo": levo_preset, "newrelic": newrelic_preset, + "signoz": signoz_preset, "weave_otel": weave_preset, } ) @@ -58,6 +64,7 @@ DYNAMIC_HEADERS_BY_CALLBACK: Final[Mapping[str, Callable[[StandardCallbackDynami "arize": arize_dynamic_headers, "langfuse_otel": langfuse_dynamic_headers, "newrelic": newrelic_dynamic_headers, + "signoz": signoz_dynamic_headers, "weave_otel": weave_dynamic_headers, } ) @@ -71,6 +78,7 @@ DYNAMIC_ENDPOINT_BY_CALLBACK: Final[Mapping[str, Callable[[StandardCallbackDynam MappingProxyType( { "newrelic": newrelic_dynamic_endpoint, + "signoz": signoz_dynamic_endpoint, } ) ) @@ -153,5 +161,6 @@ __all__ = [ "newrelic_preset", "phoenix_preset", "project_routing_headers", + "signoz_preset", "weave_preset", ] diff --git a/litellm/integrations/otel/presets/signoz.py b/litellm/integrations/otel/presets/signoz.py new file mode 100644 index 00000000000..c4d7ed48a38 --- /dev/null +++ b/litellm/integrations/otel/presets/signoz.py @@ -0,0 +1,95 @@ +from functools import lru_cache +from types import MappingProxyType +from typing import Final + +from pydantic import Field +from pydantic_settings import BaseSettings, SettingsConfigDict + +import litellm +from litellm._logging import verbose_logger +from litellm.integrations.otel.model.config import ( + ExporterOwner, + ExporterSpec, + OpenTelemetryV2Config, +) +from litellm.integrations.otel.presets.utils import ensure_mappers +from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host +from litellm.types.utils import StandardCallbackDynamicParams + +SIGNOZ_INGESTION_ENDPOINT_ENV: Final = "SIGNOZ_INGESTION_ENDPOINT" + + +class _SigNozSettings(BaseSettings): + model_config = SettingsConfigDict(case_sensitive=False, extra="ignore") + + endpoint: str | None = Field(default=None, validation_alias=SIGNOZ_INGESTION_ENDPOINT_ENV) + ingestion_key: str | None = Field(default=None, validation_alias="SIGNOZ_INGESTION_KEY") + + +def signoz_preset( + *, + config_overrides: OpenTelemetryV2Config | None = None, + allow_missing_credentials: bool = False, +) -> OpenTelemetryV2Config: + settings: Final = _SigNozSettings() + base: Final = config_overrides or OpenTelemetryV2Config() + key: Final = settings.ingestion_key + spec: Final = ExporterSpec( + kind="otlp_http", + endpoint=settings.endpoint, + headers=(f"signoz-ingestion-key={key}" if key else None), + owner=ExporterOwner.SIGNOZ, + requires_headers=bool(key), + ) + return base.model_copy( + update=MappingProxyType( + { + "exporters": (*base.exporters, spec), + "mapper_names": ensure_mappers(base.mapper_names, "genai"), + } + ) + ) + + +@lru_cache(maxsize=128) +def _warn_host_not_allowlisted(endpoint: str) -> None: + verbose_logger.warning( + "SigNoz: not exporting to key/team endpoint '%s'. Add its host to " + "litellm_settings.provider_url_destination_allowed_hosts to permit it", + endpoint, + ) + + +@lru_cache(maxsize=128) +def _warn_endpoint_without_key(endpoint: str) -> None: + verbose_logger.warning( + "SigNoz: not exporting to key/team endpoint '%s'. Set signoz_ingestion_key alongside it; " + "a keyless collector needs the global callback", + endpoint, + ) + + +def _tenant_endpoint_is_unusable(params: StandardCallbackDynamicParams) -> bool: + return bool(params.get("signoz_ingestion_endpoint")) and signoz_dynamic_endpoint(params) is None + + +def signoz_dynamic_endpoint(params: StandardCallbackDynamicParams) -> str | None: + endpoint: Final = params.get("signoz_ingestion_endpoint") + if not endpoint or not endpoint.startswith(("http://", "https://")): + return None + if not params.get("signoz_ingestion_key"): + _warn_endpoint_without_key(endpoint) + return None + if not is_url_destination_allowed_by_host(endpoint, litellm.provider_url_destination_allowed_hosts): + _warn_host_not_allowlisted(endpoint) + return None + return endpoint + + +def signoz_dynamic_headers( + params: StandardCallbackDynamicParams, +) -> dict[str, str]: # mutable-ok: DYNAMIC_HEADERS_BY_CALLBACK returns a dict + key: Final = params.get("signoz_ingestion_key") + if _tenant_endpoint_is_unusable(params) or not key: + return {} # mutable-ok: same registry contract + return {"signoz-ingestion-key": key} # mutable-ok: same registry contract diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index 7049fdd1f39..1d277995211 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -89,6 +89,7 @@ class CustomLoggerRegistry: "langtrace": OpenTelemetry, "weave_otel": OpenTelemetry, "levo": OpenTelemetry, + "signoz": OpenTelemetry, "mlflow": MlflowLogger, "langfuse": LangfusePromptManagement, "otel": OpenTelemetry, diff --git a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py index 00ab05aba77..3100ca6fba1 100644 --- a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py +++ b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py @@ -113,6 +113,8 @@ _supported_callback_params: Final[tuple[str, ...]] = ( "dd_agent_port", "newrelic_api_key", "newrelic_region", + "signoz_ingestion_endpoint", + "signoz_ingestion_key", "turn_off_message_logging", ) @@ -126,6 +128,8 @@ _request_blocked_callback_params: Final = frozenset( "dd_agent_port", "newrelic_api_key", "newrelic_region", + "signoz_ingestion_endpoint", + "signoz_ingestion_key", } ) @@ -138,6 +142,8 @@ _trusted_overlay_callback_params: Final = frozenset( { "newrelic_api_key", "newrelic_region", + "signoz_ingestion_endpoint", + "signoz_ingestion_key", } ) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 9ee7a7b0a7a..152fd54e55d 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -4931,6 +4931,38 @@ def _init_custom_logger_compatible_class( _in_memory_loggers.append(_otel_logger) return _otel_logger + elif logging_integration == "signoz": + from litellm.integrations.otel.presets.signoz import ( + SIGNOZ_INGESTION_ENDPOINT_ENV, + ) + + _signoz_endpoint: Final = os.getenv(SIGNOZ_INGESTION_ENDPOINT_ENV) + if not _signoz_endpoint: + raise ValueError(f"{SIGNOZ_INGESTION_ENDPOINT_ENV} not found in environment variables") + + _signoz_v2: Final = _maybe_construct_otel_v2("signoz", _in_memory_loggers) + if _signoz_v2 is not None: + return _signoz_v2 + + from litellm.integrations.opentelemetry import ( + OpenTelemetry, + OpenTelemetryConfig, + ) + + _signoz_base: Final = _signoz_endpoint.rstrip("/") + _signoz_key: Final = os.getenv("SIGNOZ_INGESTION_KEY") + _signoz_config: Final = OpenTelemetryConfig( + exporter="otlp_http", + endpoint=(_signoz_base if _signoz_base.endswith("/v1/traces") else f"{_signoz_base}/v1/traces"), + headers=(f"signoz-ingestion-key={_signoz_key}" if _signoz_key else None), + ) + for callback in _in_memory_loggers: + if isinstance(callback, OpenTelemetry) and callback.callback_name == "signoz": + return callback + _signoz_logger: Final = OpenTelemetry(config=_signoz_config, callback_name="signoz") + _in_memory_loggers.append(_signoz_logger) + return _signoz_logger + elif logging_integration == "mlflow": for callback in _in_memory_loggers: if isinstance(callback, MlflowLogger): diff --git a/litellm/proxy/_experimental/out/assets/logos/signoz.svg b/litellm/proxy/_experimental/out/assets/logos/signoz.svg new file mode 100644 index 00000000000..9064cb86bd6 --- /dev/null +++ b/litellm/proxy/_experimental/out/assets/logos/signoz.svg @@ -0,0 +1 @@ + \ No newline at end of file diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 14aa42afefd..34d7fc1e0f0 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4035,6 +4035,12 @@ class AllCallbacks(LiteLLMPydanticObjectBase): ], ) + signoz: CallbackOnUI = CallbackOnUI( + litellm_callback_name="signoz", + ui_callback_name="SigNoz", + litellm_callback_params=("SIGNOZ_INGESTION_ENDPOINT", "SIGNOZ_INGESTION_KEY"), + ) + zerobus: CallbackOnUI = CallbackOnUI( litellm_callback_name="zerobus", ui_callback_name="Databricks Zerobus", diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index fbd4d57bf77..07be73d7573 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -221,6 +221,7 @@ services = ( "galileo", "newrelic", "pointfive", + "signoz", "sqs", ] | str @@ -309,6 +310,7 @@ async def health_services_endpoint( "galileo", "newrelic", "pointfive", + "signoz", "sqs", ]: raise HTTPException( diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 56f647d5acc..866d84ca8f2 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -924,6 +924,8 @@ def convert_key_logging_metadata_to_callback( # must not export to it. if var.startswith("newrelic_") and data.callback_name != "newrelic": continue + if var.startswith("signoz_") and data.callback_name != "signoz": + continue if team_callback_settings_obj.callback_vars is None: team_callback_settings_obj.callback_vars = {} team_callback_settings_obj.callback_vars[var] = str(value) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 4bda1dd53ce..8862df9dc22 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3679,6 +3679,9 @@ class StandardCallbackDynamicParams(TypedDict, total=False): newrelic_api_key: str | None # writable-ok: initialize_standard_callback_dynamic_params assigns into the dict newrelic_region: str | None # writable-ok: initialize_standard_callback_dynamic_params assigns into the dict + signoz_ingestion_endpoint: str | None # writable-ok: initialize_standard_callback_dynamic_params assigns it + signoz_ingestion_key: str | None # writable-ok: initialize_standard_callback_dynamic_params assigns it + # Logging settings turn_off_message_logging: bool | None # when true will not log messages litellm_disabled_callbacks: list[str] | None diff --git a/tests/integration/observability/test_signoz_delivery.py b/tests/integration/observability/test_signoz_delivery.py new file mode 100644 index 00000000000..f3d715fe5cf --- /dev/null +++ b/tests/integration/observability/test_signoz_delivery.py @@ -0,0 +1,980 @@ +import asyncio +import base64 +import json +import os +import re +import signal +import threading +import uuid +from collections import deque +from collections.abc import Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from functools import partial +from pathlib import Path +from typing import Final + +import anthropic +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually, gateway_from_environment, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import OwnedProxy, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest +from opentelemetry.proto.common.v1.common_pb2 import AnyValue +from pydantic import JsonValue, TypeAdapter + +MARKER: Final = re.compile(rb"signoz-[0-9a-f]{32}") +JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +RESPONSE_ID: Final = "gen_ai.response.id" +INGESTION_HEADER: Final = "signoz-ingestion-key" +OPERATOR_KEY: Final = "operator-ingestion-" + uuid.uuid4().hex +TENANT_KEY: Final = "tenant-ingestion-" + uuid.uuid4().hex + + +def _marker() -> str: + return "signoz-" + uuid.uuid4().hex + + +def _chat_reply(identity: str, stream: bool) -> Reply: + if not stream: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "signoz ok"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9}, + } + ).encode() + ) + chunk: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"} + return Reply( + content_type="text/event-stream", + chunks=( + b"data: " + + json.dumps( + {**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "signoz"}}]} + ).encode() + + b"\n\n", + b"data: " + + json.dumps( + {**chunk, "choices": [{"index": 0, "delta": {"content": " ok"}, "finish_reason": "stop"}]} + ).encode() + + b"\n\n", + b"data: " + + json.dumps( + {**chunk, "choices": [], "usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9}} + ).encode() + + b"\n\n", + b"data: [DONE]\n\n", + ), + ) + + +def _responses_reply(identity: str, stream: bool) -> Reply: + response: Final[dict[str, JsonValue]] = { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "id": "msg_" + identity, + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "signoz ok", "annotations": []}], + } + ], + "usage": {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9}, + } + if not stream: + return Reply(body=json.dumps(response).encode()) + events: Final[tuple[dict[str, JsonValue], ...]] = ( + { + "type": "response.created", + "sequence_number": 0, + "response": {**response, "status": "in_progress", "output": []}, + }, + { + "type": "response.output_text.delta", + "sequence_number": 1, + "item_id": "msg_" + identity, + "output_index": 0, + "content_index": 0, + "delta": "signoz ok", + }, + {"type": "response.completed", "sequence_number": 2, "response": response}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _upstream(request: Request) -> Reply: + found: Final = MARKER.search(request.body) + if found is None: + return Reply(status=404, body=b'{"error":"no marker"}') + if request.headers.get("authorization") == "Bearer revoked-provider-key": + return Reply( + status=401, body=b'{"error":{"message":"Incorrect API key provided","type":"invalid_request_error"}}' + ) + marker: Final = found.group(0).decode() + stream: Final = object_value(JSON.validate_json(request.body)).get("stream") is True + if request.target.endswith("/responses"): + return _responses_reply(f"resp_{marker}", stream) + return _chat_reply(f"chatcmpl-{marker}", stream) + + +def _decoded_responses_id(identity: str) -> str: + try: + return base64.b64decode(identity.removeprefix("resp_").encode()).decode() + except (ValueError, UnicodeDecodeError): + return identity + + +def _canonical_id(identity: str) -> str: + return _decoded_responses_id(identity).rpartition("response_id:")[2] + + +def _sse_events(text: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + object_value(JSON.validate_json(line[6:])) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + + +def _text_at(payload: JsonValue, *path: str) -> str: + if not path: + return string_value(payload) + return _text_at(object_value(payload)[path[0]], *path[1:]) + + +def _body_id(response: httpx.Response) -> str: + return _text_at(JSON.validate_json(response.content), "id") + + +@dataclass(frozen=True, slots=True) +class Span: + target: str + ingestion_key: str | None + attributes: Mapping[str, str] + + +def _attribute_text(value: AnyValue) -> str: + match value.WhichOneof("value"): + case "string_value": + return value.string_value + case "int_value": + return str(value.int_value) + case "double_value": + return str(value.double_value) + case "bool_value": + return str(value.bool_value) + case _: + return "" + + +@dataclass(frozen=True, slots=True) +class Collector: + wire: Wire + outage: threading.Event + rejection: threading.Event + missing: threading.Event + slow: threading.Event + release: threading.Event + accepted: Sequence[Request] + refused: Sequence[Request] + guard: threading.Lock + + def refused_batch_carrying(self, response_id: str) -> Request: + def carrying() -> tuple[Request, ...]: + with self.guard: + return tuple(batch for batch in self.refused if response_id.encode() in batch.body) + + return eventually(carrying, lambda found: len(found) >= 1, seconds=30)[0] + + def refused_batches(self) -> tuple[Request, ...]: + def refused() -> tuple[Request, ...]: + with self.guard: + return tuple(self.refused) + + return eventually(refused, lambda found: len(found) >= 1, seconds=30) + + def spans(self) -> tuple[Span, ...]: + with self.guard: + batches: Final = tuple(self.accepted) + return tuple( + Span( + batch.target, + batch.headers.get(INGESTION_HEADER), + {attribute.key: _attribute_text(attribute.value) for attribute in span.attributes}, + ) + for batch in batches + for resource in ExportTraceServiceRequest.FromString(batch.body).resource_spans + for scope in resource.scope_spans + for span in scope.spans + ) + + def spans_for(self, response_id: str) -> tuple[Span, ...]: + return tuple( + span + for span in self.spans() + if RESPONSE_ID in span.attributes + and _canonical_id(span.attributes[RESPONSE_ID]) == _canonical_id(response_id) + ) + + def single_span(self, response_id: str, *, elsewhere: "Collector | None" = None) -> Span: + found: Final = eventually( + lambda: self.spans_for(response_id), lambda spans: len(spans) == 1, seconds=30, return_last_on_timeout=True + ) + assert len(found) == 1, ( + f"{len(found)} spans for {response_id} at this sink; other sink saw " + f"{elsewhere.landed((response_id,)) if elsewhere else 'n/a'}" + ) + return found[0] + + def landed(self, response_ids: Sequence[str]) -> dict[str, int]: + spans: Final = self.spans() + return { + _canonical_id(identity): sum( + 1 + for span in spans + if RESPONSE_ID in span.attributes + and _canonical_id(span.attributes[RESPONSE_ID]) == _canonical_id(identity) + ) + for identity in response_ids + } + + +def _collector() -> Iterator[Collector]: + outage: Final = threading.Event() + rejection: Final = threading.Event() + missing: Final = threading.Event() + slow: Final = threading.Event() + release: Final = threading.Event() + accepted: Final[deque[Request]] = deque() # mutable-ok: the sink thread records each accepted batch as it arrives + refused: Final[deque[Request]] = deque() # mutable-ok: the sink thread records each refused batch as it arrives + guard: Final = threading.Lock() + + def refuse(request: Request, status: int, body: bytes) -> Reply: + with guard: + refused.append(request) + return Reply(status=status, body=body) + + def sink(request: Request) -> Reply: + if slow.is_set(): + release.wait(timeout=30) + if outage.is_set(): + return refuse(request, 503, b'{"error":"sink down"}') + if rejection.is_set(): + return refuse(request, 403, b'{"error":"forbidden"}') + if missing.is_set(): + return refuse(request, 404, b'{"error":"not found"}') + with guard: + accepted.append(request) + return Reply() + + with wire_server(sink) as wire: + yield Collector(wire, outage, rejection, missing, slow, release, accepted, refused, guard) + + +@pytest.fixture(scope="session") +def operator_sink() -> Iterator[Collector]: + yield from _collector() + + +@pytest.fixture(scope="session") +def tenant_sink() -> Iterator[Collector]: + yield from _collector() + + +@pytest.fixture(scope="session") +def provider() -> Iterator[Wire]: + with wire_server(_upstream) as wire: + yield wire + + +@dataclass(frozen=True, slots=True) +class Rig: + proxy: Gateway + process: OwnedProxy + model: str + upstream: Wire + sink: Collector + tenant_sink: Collector + + def openai_client(self) -> openai.OpenAI: + return openai.OpenAI(base_url=str(self.proxy.client.base_url) + "/v1", api_key=self.proxy.key, max_retries=0) + + def async_openai_client(self) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI( + base_url=str(self.proxy.client.base_url) + "/v1", api_key=self.proxy.key, max_retries=0 + ) + + def anthropic_client(self) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=str(self.proxy.client.base_url), api_key=self.proxy.key, max_retries=0) + + def async_anthropic_client(self) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic(base_url=str(self.proxy.client.base_url), api_key=self.proxy.key, max_retries=0) + + def chat( + self, marker: str, *, headers: Mapping[str, str] | None = None, key: str | None = None, **extra: JsonValue + ) -> httpx.Response: + return self.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": self.model, + "messages": [{"role": "user", "content": marker}], + "cache": {"no-cache": True}, + **extra, + }, + headers=headers, + key=key, + ) + + def upstream_bodies(self, marker: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + object_value(JSON.validate_json(request.body)) + for request in self.upstream.drain() + if marker.encode() in request.body + ) + + def spend_rows(self, response_id: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + eventually( + lambda: read_rows( + 'SELECT request_id, spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (response_id,) + ), + lambda values: len(values) == 1, + seconds=70, + ) + ) + + def tenant_logging(self, endpoint: str | None, key: str | None) -> JsonValue: + variables: Final[dict[str, JsonValue]] = { + **({"signoz_ingestion_endpoint": endpoint} if endpoint is not None else {}), + **({"signoz_ingestion_key": key} if key is not None else {}), + } + return [{"callback_name": "signoz", "callback_type": "success", "callback_vars": variables}] + + +@dataclass(frozen=True, slots=True) +class RigFactory: + provider: Wire + sink: Collector + tenant_sink: Collector + directory: Path + otel_v2: bool + workers: int + endpoint: str | None + ingestion_key: str | None = OPERATOR_KEY + + def config_path(self) -> Path: + loaded: Final = object_value( + JSON.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + ) + config: Final = { + **loaded, + "litellm_settings": { + **object_value(loaded["litellm_settings"]), + "callbacks": ["signoz"], + "provider_url_destination_allowed_hosts": [self.tenant_sink.wire.url], + }, + "general_settings": {**object_value(loaded["general_settings"]), "disable_model_info_refresh": True}, + } + path: Final = self.directory / f"signoz-{uuid.uuid4().hex}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + def overrides(self) -> dict[str, str]: + return { + "LITELLM_OTEL_V2": "1" if self.otel_v2 else "0", + "OTEL_BSP_SCHEDULE_DELAY": "300", + **({"SIGNOZ_INGESTION_ENDPOINT": self.endpoint} if self.endpoint is not None else {}), + **({"SIGNOZ_INGESTION_KEY": self.ingestion_key} if self.ingestion_key is not None else {}), + } + + def start(self) -> Iterator[Rig]: + with ( + gateway_from_environment() as gateway, + owned_proxy_process( + gateway, + self.directory, + self.overrides(), + config=self.config_path(), + remove_environment=("SIGNOZ_INGESTION_ENDPOINT", "SIGNOZ_INGESTION_KEY"), + workers=self.workers, + ) as owned, + owned.gateway.scenario() as scenario, + ): + model: Final = scenario.model(api_base=self.provider.url + "/v1") + yield Rig(owned.gateway, owned, model, self.provider, self.sink, self.tenant_sink) + + +@pytest.fixture(scope="session") +def rig( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Rig]: + factory: Final = RigFactory( + provider, operator_sink, tenant_sink, tmp_path_factory.mktemp("signoz"), False, 2, operator_sink.wire.url + ) + yield from factory.start() + + +@pytest.fixture(scope="session") +def v2_rig( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Rig]: + factory: Final = RigFactory( + provider, operator_sink, tenant_sink, tmp_path_factory.mktemp("signoz-v2"), True, 2, operator_sink.wire.url + ) + yield from factory.start() + + +def _assert_operator_span(rig: Rig, response_id: str, marker: str) -> Span: + span: Final = rig.sink.single_span(response_id) + assert span.target == "/v1/traces", span + assert span.ingestion_key == OPERATOR_KEY, span + assert rig.tenant_sink.landed((response_id,)) == {_canonical_id(response_id): 0} + bodies: Final = rig.upstream_bodies(marker) + assert len(bodies) == 1, bodies + assert "signoz" not in json.dumps(bodies[0]).replace(marker, ""), bodies[0] + return span + + +def test_signoz_is_registered_as_an_opentelemetry_callback(rig: Rig) -> None: + listed: Final = rig.proxy.request("GET", "/active/callbacks") + assert listed.status_code == 200, listed.text + assert "OpenTelemetry" in json.dumps(listed.json()), listed.text + log: Final = rig.process.log.read_text() + assert "SIGNOZ_INGESTION_ENDPOINT not found" not in log + + +def test_chat_completion_sdk_span_lands_at_the_operator_sink_with_the_ingestion_key(rig: Rig) -> None: + marker: Final = _marker() + completion: Final = rig.openai_client().chat.completions.create( + model=rig.model, messages=[{"role": "user", "content": marker}] + ) + assert completion.id == f"chatcmpl-{marker}" + span: Final = _assert_operator_span(rig, completion.id, marker) + assert span.attributes.get("gen_ai.request.model") or span.attributes.get("llm.request.model"), span + rows: Final = rig.spend_rows(completion.id) + assert rows[0]["request_id"] == completion.id, rows + + +def test_chat_stream_async_sdk_span_lands_once_after_the_stream_is_consumed(rig: Rig) -> None: + marker: Final = _marker() + + async def consume() -> frozenset[str]: + stream: Final = await rig.async_openai_client().chat.completions.create( + model=rig.model, messages=[{"role": "user", "content": marker}], stream=True + ) + return frozenset([chunk.id async for chunk in stream]) + + identities: Final = asyncio.run(consume()) + assert identities == {f"chatcmpl-{marker}"}, identities + _assert_operator_span(rig, f"chatcmpl-{marker}", marker) + + +def test_messages_sdk_span_lands_at_the_operator_sink(rig: Rig) -> None: + marker: Final = _marker() + message: Final = rig.anthropic_client().messages.create( + model=rig.model, max_tokens=16, messages=[{"role": "user", "content": marker}] + ) + _assert_operator_span(rig, message.id, marker) + + +def test_messages_stream_async_sdk_span_lands_once_after_the_stream_is_consumed(rig: Rig) -> None: + marker: Final = _marker() + + async def consume() -> str: + async with rig.async_anthropic_client().messages.stream( + model=rig.model, max_tokens=16, messages=[{"role": "user", "content": marker}] + ) as stream: + async for _ in stream: + pass + return (await stream.get_final_message()).id + + identity: Final = asyncio.run(consume()) + _assert_operator_span(rig, identity, marker) + + +def test_responses_sdk_span_lands_at_the_operator_sink(rig: Rig) -> None: + marker: Final = _marker() + response: Final = rig.openai_client().responses.create(model=rig.model, input=marker) + _assert_operator_span(rig, response.id, marker) + + +def test_responses_stream_raw_httpx_span_lands_once_after_the_stream_is_consumed(rig: Rig) -> None: + marker: Final = _marker() + response: Final = rig.proxy.request("POST", "/v1/responses", {"model": rig.model, "input": marker, "stream": True}) + assert response.status_code == 200, response.text + _assert_operator_span(rig, _responses_id(response, marker), marker) + + +def test_v2_flag_on_still_delivers_the_operator_span_with_the_ingestion_key(v2_rig: Rig) -> None: + marker: Final = _marker() + response: Final = v2_rig.chat(marker) + assert response.status_code == 200, response.text + _assert_operator_span(v2_rig, _body_id(response), marker) + + +def test_endpoint_already_ending_in_v1_traces_is_not_doubled( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> None: + factory: Final = RigFactory( + provider, + operator_sink, + tenant_sink, + tmp_path_factory.mktemp("signoz-suffixed"), + False, + 2, + operator_sink.wire.url + "/v1/traces", + ) + suffixed: Final = next(started := factory.start()) + marker: Final = _marker() + response: Final = suffixed.chat(marker) + assert response.status_code == 200, response.text + span: Final = suffixed.sink.single_span(_body_id(response)) + assert span.target == "/v1/traces", span + assert tuple(started) == () + + +def test_three_identical_requests_produce_one_span_each(rig: Rig) -> None: + markers: Final = tuple(_marker() for _ in range(3)) + responses: Final = tuple(rig.chat(marker) for marker in markers) + assert all(response.status_code == 200 for response in responses), [response.text for response in responses] + identities: Final = tuple(_body_id(response) for response in responses) + landed: Final = eventually( + lambda: rig.sink.landed(identities), lambda seen: all(count >= 1 for count in seen.values()), seconds=30 + ) + assert landed == {identity: 1 for identity in identities}, landed + assert rig.sink.landed(identities) == landed + + +def test_unauthenticated_request_is_rejected_without_an_upstream_call_and_any_span_records_the_401(rig: Rig) -> None: + marker: Final = _marker() + response: Final = rig.chat(marker, key="sk-not-a-real-key") + assert response.status_code == 401, response.text + later: Final = rig.chat(_marker()) + assert later.status_code == 200, later.text + rig.sink.single_span(_body_id(later)) + assert rig.upstream_bodies(marker) == () + marker_spans: Final = tuple(span for span in rig.sink.spans() if marker in json.dumps(span.attributes)) + assert all(span.attributes.get("error.code") == "401" for span in marker_spans), marker_spans + assert not any(span.attributes.get(RESPONSE_ID, "").startswith("chatcmpl-") for span in marker_spans), marker_spans + + +def test_request_supplied_signoz_variables_are_refused_before_the_upstream_is_called(rig: Rig) -> None: + marker: Final = _marker() + response: Final = rig.chat( + marker, + metadata={"signoz_ingestion_endpoint": rig.tenant_sink.wire.url, "signoz_ingestion_key": TENANT_KEY}, + ) + assert response.status_code == 401, response.text + assert "signoz_ingestion_endpoint is not allowed in request body" in response.text + assert rig.upstream_bodies(marker) == () + assert not any(marker in json.dumps(span.attributes) for span in rig.tenant_sink.spans()) + + +def test_upstream_401_reaches_the_caller_and_unrelated_traffic_keeps_landing(rig: Rig) -> None: + marker: Final = _marker() + with rig.proxy.scenario() as scenario: + broken: Final = scenario.model(api_base=rig.upstream.url + "/v1", api_key="revoked-provider-key") + failed: Final = rig.proxy.request( + "POST", "/v1/chat/completions", {"model": broken, "messages": [{"role": "user", "content": marker}]} + ) + assert failed.status_code == 401, failed.text + assert "Incorrect API key provided" in failed.text + healthy_marker: Final = _marker() + healthy: Final = rig.chat(healthy_marker) + assert healthy.status_code == 200, healthy.text + _assert_operator_span(rig, _body_id(healthy), healthy_marker) + + +def test_health_services_accepts_signoz(rig: Rig) -> None: + response: Final = rig.proxy.request("GET", "/health/services", params={"service": "signoz"}) + assert response.status_code == 200, response.text + + +def test_sink_answering_403_drops_those_spans_and_later_spans_still_land(rig: Rig) -> None: + rig.sink.rejection.set() + try: + rejected: Final = rig.chat(_marker()) + assert rejected.status_code == 200, rejected.text + rig.sink.refused_batch_carrying(_body_id(rejected)) + finally: + rig.sink.rejection.clear() + later: Final = rig.chat(_marker()) + assert later.status_code == 200, later.text + rig.sink.single_span(_body_id(later)) + assert rig.sink.landed((_body_id(rejected),)) == {_body_id(rejected): 0} + + +def test_sink_answering_404_drops_those_spans_and_later_spans_still_land(rig: Rig) -> None: + rig.sink.missing.set() + try: + dropped: Final = rig.chat(_marker()) + assert dropped.status_code == 200, dropped.text + rig.sink.refused_batch_carrying(_body_id(dropped)) + finally: + rig.sink.missing.clear() + later: Final = rig.chat(_marker()) + assert later.status_code == 200, later.text + rig.sink.single_span(_body_id(later)) + assert rig.sink.landed((_body_id(dropped),)) == {_body_id(dropped): 0} + + +def test_key_level_signoz_destination_routes_the_span_to_the_tenant_sink(v2_rig: Rig) -> None: + marker: Final = _marker() + with v2_rig.proxy.scenario() as scenario: + token: Final = scenario.key( + metadata={"logging": v2_rig.tenant_logging(v2_rig.tenant_sink.wire.url, TENANT_KEY)} + ) + response: Final = v2_rig.chat(marker, key=token) + assert response.status_code == 200, response.text + identity: Final = _body_id(response) + span: Final = v2_rig.tenant_sink.single_span(identity, elsewhere=v2_rig.sink) + assert span.ingestion_key == TENANT_KEY, span + assert v2_rig.sink.landed((identity,)) == {identity: 0}, "operator sink also received the tenant span" + + +def test_team_level_signoz_destination_routes_the_span_to_the_tenant_sink(v2_rig: Rig) -> None: + marker: Final = _marker() + with v2_rig.proxy.scenario() as scenario: + team: Final = scenario.team( + metadata={"logging": v2_rig.tenant_logging(v2_rig.tenant_sink.wire.url, TENANT_KEY)} + ) + token: Final = scenario.key(team_id=team) + response: Final = v2_rig.chat(marker, key=token) + assert response.status_code == 200, response.text + identity: Final = _body_id(response) + span: Final = v2_rig.tenant_sink.single_span(identity, elsewhere=v2_rig.sink) + assert span.ingestion_key == TENANT_KEY, span + assert v2_rig.sink.landed((identity,)) == {identity: 0}, "operator sink also received the tenant span" + + +def test_key_level_destination_wins_over_the_team_level_destination(v2_rig: Rig) -> None: + marker: Final = _marker() + team_key: Final = "team-" + TENANT_KEY + with v2_rig.proxy.scenario() as scenario: + team: Final = scenario.team(metadata={"logging": v2_rig.tenant_logging(v2_rig.tenant_sink.wire.url, team_key)}) + token: Final = scenario.key( + team_id=team, metadata={"logging": v2_rig.tenant_logging(v2_rig.tenant_sink.wire.url, TENANT_KEY)} + ) + response: Final = v2_rig.chat(marker, key=token) + assert response.status_code == 200, response.text + span: Final = v2_rig.tenant_sink.single_span(_body_id(response)) + assert span.ingestion_key == TENANT_KEY, span + + +def test_team_endpoint_without_an_ingestion_key_is_ignored_and_the_span_stays_at_the_operator_sink( + v2_rig: Rig, +) -> None: + marker: Final = _marker() + with v2_rig.proxy.scenario() as scenario: + team: Final = scenario.team(metadata={"logging": v2_rig.tenant_logging(v2_rig.tenant_sink.wire.url, None)}) + token: Final = scenario.key(team_id=team) + response: Final = v2_rig.chat(marker, key=token) + assert response.status_code == 200, response.text + _assert_operator_span(v2_rig, _body_id(response), marker) + eventually( + lambda: v2_rig.process.log.read_text(), + lambda text: "Set signoz_ingestion_key alongside it" in text, + seconds=30, + ) + + +def test_team_endpoint_off_the_allowlist_keeps_the_span_at_the_operator_sink(v2_rig: Rig) -> None: + marker: Final = _marker() + with v2_rig.proxy.scenario() as scenario: + team: Final = scenario.team( + metadata={"logging": v2_rig.tenant_logging("http://tenant.invalid:4318/v1/traces", TENANT_KEY)} + ) + token: Final = scenario.key(team_id=team) + response: Final = v2_rig.chat(marker, key=token) + assert response.status_code == 200, response.text + _assert_operator_span(v2_rig, _body_id(response), marker) + eventually( + lambda: v2_rig.process.log.read_text(), + lambda text: "provider_url_destination_allowed_hosts" in text, + seconds=30, + ) + + +def test_legacy_mode_ignores_key_level_signoz_destination_and_keeps_the_operator_sink(rig: Rig) -> None: + marker: Final = _marker() + with rig.proxy.scenario() as scenario: + token: Final = scenario.key(metadata={"logging": rig.tenant_logging(rig.tenant_sink.wire.url, TENANT_KEY)}) + response: Final = rig.chat(marker, key=token) + assert response.status_code == 200, response.text + _assert_operator_span(rig, _body_id(response), marker) + + +@pytest.mark.parametrize( + "endpoint", + ["", "not-a-url", "ftp://tenant.invalid", "x" * 5000, 12345, ["http://tenant.invalid"]], + ids=["empty", "bare", "ftp", "5kb", "int", "list"], +) +def test_hostile_tenant_endpoint_never_breaks_the_request_or_the_operator_sink( + v2_rig: Rig, endpoint: JsonValue +) -> None: + marker: Final = _marker() + created: Final = v2_rig.proxy.request( + "POST", + "/key/generate", + { + "metadata": { + "logging": [ + { + "callback_name": "signoz", + "callback_type": "success", + "callback_vars": {"signoz_ingestion_endpoint": endpoint, "signoz_ingestion_key": TENANT_KEY}, + } + ] + } + }, + ) + assert created.status_code in (200, 400, 422), created.text + if created.status_code != 200: + return + try: + response: Final = v2_rig.chat(marker, key=_text_at(JSON.validate_json(created.content), "key")) + assert response.status_code == 200, response.text + identity: Final = _body_id(response) + eventually( + lambda: v2_rig.sink.landed((identity,))[identity] + v2_rig.tenant_sink.landed((identity,))[identity], + lambda total: total >= 1, + seconds=30, + ) + assert v2_rig.sink.landed((identity,))[identity] + v2_rig.tenant_sink.landed((identity,))[identity] == 1 + assert v2_rig.tenant_sink.landed((identity,)) == {identity: 0}, "unusable endpoint reached the tenant sink" + finally: + v2_rig.proxy.post("/key/delete", {"keys": [created.json()["key"]]}) + + +def test_missing_ingestion_endpoint_fails_loudly_at_boot( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> None: + factory: Final = RigFactory( + provider, operator_sink, tenant_sink, tmp_path_factory.mktemp("signoz-noenv"), False, 2, None + ) + started: Final = factory.start() + broken: Final = next(started) + try: + response: Final = broken.chat(_marker()) + assert response.status_code == 200, response.text + eventually( + lambda: broken.process.log.read_text(), + lambda text: "SIGNOZ_INGESTION_ENDPOINT not found" in text, + seconds=30, + ) + assert operator_sink.landed((_body_id(response),)) == {_body_id(response): 0} + finally: + with pytest.raises(StopIteration): + next(started) + + +def test_empty_ingestion_endpoint_is_treated_as_missing( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> None: + factory: Final = RigFactory( + provider, operator_sink, tenant_sink, tmp_path_factory.mktemp("signoz-empty"), False, 2, "" + ) + started: Final = factory.start() + broken: Final = next(started) + try: + response: Final = broken.chat(_marker()) + assert response.status_code == 200, response.text + eventually( + lambda: broken.process.log.read_text(), + lambda text: "SIGNOZ_INGESTION_ENDPOINT not found" in text, + seconds=30, + ) + finally: + with pytest.raises(StopIteration): + next(started) + + +def _is_event_stream(response: httpx.Response) -> bool: + return "content-type" in response.headers and response.headers["content-type"].startswith("text/event-stream") + + +def _chat_id(response: httpx.Response) -> str: + if not _is_event_stream(response): + return _body_id(response) + identities: Final = frozenset(_text_at(event, "id") for event in _sse_events(response.text)) + assert len(identities) == 1, response.text + return next(iter(identities)) + + +def _responses_id(response: httpx.Response, marker: str) -> str: + if not _is_event_stream(response): + return _body_id(response) + completed: Final = tuple( + _text_at(event, "response", "id") + for event in _sse_events(response.text) + if event.get("type") == "response.completed" + ) + assert len(completed) == 1 and completed[0].startswith("resp_"), response.text + return f"resp_{marker}" + + +def _message_id(response: httpx.Response) -> str: + if not _is_event_stream(response): + return _body_id(response) + starts: Final = tuple( + _text_at(event, "message", "id") for event in _sse_events(response.text) if event.get("type") == "message_start" + ) + assert len(starts) == 1, response.text + return starts[0] + + +def _burst(rig: Rig, count: int) -> tuple[tuple[int, str, str | None], ...]: + markers: Final = tuple(_marker() for _ in range(count)) + + def one(index: int) -> tuple[int, str, str | None]: + marker: Final = markers[index] + headers: Final = {"Authorization": f"Bearer {rig.proxy.key}"} + stream: Final = index % 2 == 0 + path, body, identity_of = ( + ("/v1/chat/completions", {"model": rig.model, "messages": [{"role": "user", "content": marker}]}, _chat_id), + ("/v1/responses", {"model": rig.model, "input": marker}, partial(_responses_id, marker=marker)), + ( + "/v1/messages", + {"model": rig.model, "max_tokens": 16, "messages": [{"role": "user", "content": marker}]}, + _message_id, + ), + )[index % 3] + try: + response: Final = rig.proxy.client.post(path, json={**body, "stream": stream}, headers=headers) + response.read() + except httpx.HTTPError as error: + return index, marker, repr(error) + return (index, marker, response.text) if response.status_code != 200 else (index, identity_of(response), None) + + with ThreadPoolExecutor(max_workers=10) as pool: + return tuple(pool.map(one, range(count))) + + +def _assert_exactly_once(rig: Rig, identities: Sequence[str]) -> None: + landed: Final = eventually( + lambda: rig.sink.landed(identities), lambda seen: all(count >= 1 for count in seen.values()), seconds=80 + ) + assert landed == {_canonical_id(identity): 1 for identity in identities}, landed + charged: Final = frozenset(identity for identity in identities if not identity.startswith("resp_")) + spend: Final = eventually( + lambda: read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s::text[])', + ("{" + ",".join(charged) + "}",), + ), + lambda rows: {str(row["request_id"]) for row in rows} >= charged, + seconds=70, + ) + assert {str(row["request_id"]) for row in spend} == charged, spend + + +def test_sink_outage_during_a_mixed_burst_lands_every_response_exactly_once_after_recovery(rig: Rig) -> None: + rig.sink.outage.set() + try: + health_down: Final = rig.proxy.request("GET", "/health/services", params={"service": "signoz"}) + results: Final = _burst(rig, 30) + assert all(error is None for _, _, error in results), [error for _, _, error in results if error] + rig.sink.refused_batches() + finally: + rig.sink.outage.clear() + assert health_down.status_code == 200, health_down.text + _assert_exactly_once(rig, tuple(identity for _, identity, _ in results)) + + +def test_slow_sink_during_a_burst_lands_every_response_exactly_once(rig: Rig) -> None: + rig.sink.release.clear() + rig.sink.slow.set() + try: + results: Final = _burst(rig, 20) + assert all(error is None for _, _, error in results), [error for _, _, error in results if error] + identities: Final = tuple(identity for _, identity, _ in results) + assert all(count == 0 for count in rig.sink.landed(identities).values()), "sink accepted while held" + finally: + rig.sink.slow.clear() + rig.sink.release.set() + _assert_exactly_once(rig, identities) + + +def test_killing_one_of_two_workers_mid_burst_keeps_serving_and_never_duplicates_a_span(rig: Rig) -> None: + root: Final = psutil.Process(rig.process.process.pid) + workers: Final = eventually( + lambda: tuple(child for child in root.children() if "resource_tracker" not in " ".join(child.cmdline())), + lambda found: len(found) == 2, + seconds=30, + ) + markers: Final = tuple(_marker() for _ in range(24)) + + def one(index: int) -> tuple[str, str | None]: + if index == 8: + os.kill(workers[0].pid, signal.SIGKILL) + try: + response: Final = rig.chat(markers[index]) + return f"chatcmpl-{markers[index]}", None if response.status_code == 200 else response.text + except httpx.HTTPError as error: + return f"chatcmpl-{markers[index]}", repr(error) + + with ThreadPoolExecutor(max_workers=6) as pool: + results: Final = tuple(pool.map(one, range(24))) + assert rig.process.process.poll() is None, "Proxy root exited after a worker was killed" + after: Final = rig.chat(_marker()) + assert after.status_code == 200, after.text + rig.sink.single_span(_body_id(after)) + failures: Final = tuple(error for _, error in results if error) + assert all(error.startswith(("ReadError(", "RemoteProtocolError(", "ConnectError(")) for error in failures), ( + failures + ) + assert len(failures) <= 6, failures + served: Final = tuple(identity for identity, error in results if error is None) + assert len(served) >= 18, results + settled: Final = tuple(identity for index, (identity, error) in enumerate(results) if index > 14 and not error) + _assert_exactly_once(rig, settled) + assert all(count <= 1 for count in rig.sink.landed(served).values()), rig.sink.landed(served) + + +def test_terminating_the_proxy_right_after_a_burst_flushes_every_span_before_exit( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> None: + pytest.skip("BUG: spans still queued in the OTel batch processor at SIGTERM never reach the sink (4 of 10 lost)") + factory: Final = RigFactory( + provider, + operator_sink, + tenant_sink, + tmp_path_factory.mktemp("signoz-shutdown"), + False, + 2, + operator_sink.wire.url, + ) + started: Final = factory.start() + rig: Final = next(started) + responses: Final = tuple(rig.chat(_marker()) for _ in range(10)) + assert all(response.status_code == 200 for response in responses), [response.text for response in responses] + identities: Final = tuple(_body_id(response) for response in responses) + pending_at_signal: Final = rig.sink.landed(identities) + rig.process.process.terminate() + assert rig.process.process.wait(timeout=40) in (0, -signal.SIGTERM) + with pytest.raises(httpx.ConnectError): + next(started) + landed: Final = rig.sink.landed(identities) + assert landed == {identity: 1 for identity in identities}, ( + f"spans at the sink after exit: {landed}, at the moment of SIGTERM: {pending_at_signal}" + ) diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index ee42042bb8e..1a3ffecb0a5 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -8470,3 +8470,31 @@ async def test_mcp_credentials_only_removed_from_logging_copies(path: str, custo for name, value in secrets.items(): assert updated["secret_fields"]["raw_headers"][name.lower()] == value assert request.headers[name] == value + + +def test_signoz_callback_vars_are_scoped_to_the_signoz_callback(): + from litellm.proxy._types import AddTeamCallback + from litellm.proxy.litellm_pre_call_utils import convert_key_logging_metadata_to_callback + + under_signoz = convert_key_logging_metadata_to_callback( + data=AddTeamCallback( + callback_name="signoz", + callback_type="success", + callback_vars={"signoz_ingestion_key": "team-key", "signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443"}, + ), + team_callback_settings_obj=None, + ) + assert under_signoz.callback_vars == { + "signoz_ingestion_key": "team-key", + "signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443", + } + + under_other = convert_key_logging_metadata_to_callback( + data=AddTeamCallback( + callback_name="langfuse", + callback_type="success", + callback_vars={"signoz_ingestion_key": "team-key", "langfuse_host": "https://cloud.langfuse.com"}, + ), + team_callback_settings_obj=None, + ) + assert under_other.callback_vars == {"langfuse_host": "https://cloud.langfuse.com"} diff --git a/tests/unit/integrations/otel/test_otel_v2_dynamic.py b/tests/unit/integrations/otel/test_otel_v2_dynamic.py index 29772eb92c7..8163ba06317 100644 --- a/tests/unit/integrations/otel/test_otel_v2_dynamic.py +++ b/tests/unit/integrations/otel/test_otel_v2_dynamic.py @@ -1,6 +1,7 @@ """Per-request multi-tenant credential routing (V1 parity).""" import base64 +import logging import pytest from opentelemetry.trace import NoOpTracer @@ -677,3 +678,68 @@ def test_newrelic_key_only_team_routes_to_us_not_operator_region(monkeypatch): ) owned = next(e for e in new_cfg.exporters if e.owner == "newrelic") assert owned.endpoint == "https://otlp.nr-data.net" + + +def test_signoz_dynamic_headers_stamp_ingestion_key(): + from litellm.integrations.otel.presets import dynamic_otlp_headers + + assert dynamic_otlp_headers("signoz", {"signoz_ingestion_key": "team-key"}) == {"signoz-ingestion-key": "team-key"} + # No key means no per-request routing; the caller keeps its default tracer. + assert dynamic_otlp_headers("signoz", {}) is None + + +def test_signoz_dynamic_endpoint_comes_from_team_config_when_its_host_is_allowlisted(monkeypatch): + import litellm + from litellm.integrations.otel.presets import dynamic_otlp_endpoint + + monkeypatch.setattr(litellm, "provider_url_destination_allowed_hosts", ["ingest.eu.signoz.cloud"]) + assert ( + dynamic_otlp_endpoint( + "signoz", {"signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443", "signoz_ingestion_key": "k"} + ) + == "https://ingest.eu.signoz.cloud:443" + ) + # A team that saved only a key keeps the operator's configured endpoint. + assert dynamic_otlp_endpoint("signoz", {"signoz_ingestion_key": "k"}) is None + assert dynamic_otlp_endpoint("signoz", {}) is None + + +def test_signoz_team_endpoint_off_the_allowlist_is_dropped_along_with_its_key(monkeypatch): + import litellm + from litellm.integrations.otel.presets import dynamic_otlp_endpoint, dynamic_otlp_headers + + monkeypatch.setattr(litellm, "provider_url_destination_allowed_hosts", []) + params = {"signoz_ingestion_endpoint": "http://169.254.169.254/v1/traces", "signoz_ingestion_key": "k"} + assert dynamic_otlp_endpoint("signoz", params) is None + # The tenant key must not ride to the operator's collector either: the request keeps the default tracer. + assert dynamic_otlp_headers("signoz", params) is None + + +def test_signoz_keyless_team_endpoint_is_ignored_so_the_operator_key_never_reaches_it(monkeypatch, caplog): + import litellm + from litellm.integrations.otel.presets import dynamic_otlp_endpoint, dynamic_otlp_headers + from litellm.integrations.otel.presets.signoz import _warn_endpoint_without_key + + monkeypatch.setattr(litellm, "provider_url_destination_allowed_hosts", ["collector.team.internal"]) + params = {"signoz_ingestion_endpoint": "http://collector.team.internal:4318"} + _warn_endpoint_without_key.cache_clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + assert dynamic_otlp_headers("signoz", params) is None + assert "Set signoz_ingestion_key alongside it" in caplog.text + assert dynamic_otlp_endpoint("signoz", params) is None + cache = _cache( + "signoz", + exporters=[ + ExporterSpec( + kind="otlp_http", + endpoint="https://ingest.us.signoz.cloud:443", + headers="signoz-ingestion-key=OPERATOR", + owner="signoz", + requires_headers=True, + ) + ], + ) + routed = cache._routed_config({}, {}, dynamic_otlp_endpoint("signoz", params), "team-service") + owned = next(e for e in routed.exporters if e.owner == "signoz") + assert owned.endpoint == "https://ingest.us.signoz.cloud:443" + assert owned.headers == "signoz-ingestion-key=OPERATOR" diff --git a/tests/unit/integrations/otel/test_otel_v2_presets.py b/tests/unit/integrations/otel/test_otel_v2_presets.py index 58cfc1ceb3f..a060cdf3648 100644 --- a/tests/unit/integrations/otel/test_otel_v2_presets.py +++ b/tests/unit/integrations/otel/test_otel_v2_presets.py @@ -212,3 +212,69 @@ def test_newrelic_preset_unset_content_knob_keeps_default(monkeypatch): from litellm.integrations.otel.presets.newrelic import newrelic_preset assert newrelic_preset().capture_span_content is False + + +def test_signoz_preset_reads_env_endpoint_and_key(monkeypatch): + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "https://ingest.eu.signoz.cloud:443") + monkeypatch.setenv("SIGNOZ_INGESTION_KEY", "env-ingestion-key") + from litellm.integrations.otel.model.config import ExporterOwner + from litellm.integrations.otel.presets.signoz import signoz_preset + + cfg = signoz_preset() + spec = next(e for e in cfg.exporters if e.owner == ExporterOwner.SIGNOZ) + assert spec.kind == "otlp_http" + assert spec.endpoint == "https://ingest.eu.signoz.cloud:443" + assert spec.headers == "signoz-ingestion-key=env-ingestion-key" + assert spec.requires_headers is True + assert "genai" in cfg.mapper_names + + +def test_signoz_preset_without_key_is_self_hosted(monkeypatch): + # A self-hosted collector accepts unauthenticated OTLP, so requiring headers + # would drop exports that would have succeeded. + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "http://signoz-collector.internal:4318") + monkeypatch.delenv("SIGNOZ_INGESTION_KEY", raising=False) + from litellm.integrations.otel.model.config import ExporterOwner + from litellm.integrations.otel.presets.signoz import signoz_preset + + cfg = signoz_preset() + spec = next(e for e in cfg.exporters if e.owner == ExporterOwner.SIGNOZ) + assert spec.endpoint == "http://signoz-collector.internal:4318" + assert spec.headers is None + assert spec.requires_headers is False + + +def test_signoz_preset_has_no_default_endpoint(monkeypatch): + # No region table and no default host: the preset never invents a destination. + monkeypatch.delenv("SIGNOZ_INGESTION_ENDPOINT", raising=False) + monkeypatch.delenv("SIGNOZ_INGESTION_KEY", raising=False) + from litellm.integrations.otel.model.config import ExporterOwner + from litellm.integrations.otel.presets.signoz import signoz_preset + + cfg = signoz_preset() + spec = next(e for e in cfg.exporters if e.owner == ExporterOwner.SIGNOZ) + assert spec.endpoint is None + + +def test_signoz_preset_endpoint_passed_through_verbatim(monkeypatch): + # The plumbing appends the signal path, so pre-appending would double it. + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "https://ingest.us.signoz.cloud:443/v1/traces") + monkeypatch.delenv("SIGNOZ_INGESTION_KEY", raising=False) + from litellm.integrations.otel.model.config import ExporterOwner + from litellm.integrations.otel.plumbing.providers import _otlp_traces_endpoint + from litellm.integrations.otel.presets.signoz import signoz_preset + + cfg = signoz_preset() + spec = next(e for e in cfg.exporters if e.owner == ExporterOwner.SIGNOZ) + assert spec.endpoint == "https://ingest.us.signoz.cloud:443/v1/traces" + assert _otlp_traces_endpoint(spec.endpoint) == "https://ingest.us.signoz.cloud:443/v1/traces" + + +def test_signoz_preset_accepts_the_factory_call_shape(monkeypatch): + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "http://127.0.0.1:1") + monkeypatch.delenv("SIGNOZ_INGESTION_KEY", raising=False) + from litellm.integrations.otel.model.config import ExporterOwner + from litellm.integrations.otel.presets import PRESET_BY_CALLBACK + + cfg = PRESET_BY_CALLBACK["signoz"](allow_missing_credentials=True) + assert any(e.owner == ExporterOwner.SIGNOZ for e in cfg.exporters) diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index c8b02ebc790..2fc747e1b48 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -8932,3 +8932,100 @@ async def test_prompt_management_with_unchanged_variables_replays_a_byte_identic assert json.dumps(messages_n_plus_one[: len(messages_n)], sort_keys=True) == json.dumps(messages_n, sort_keys=True) assert messages_n[0] == {"role": "system", "content": "You are a pirate. Answer in one sentence."} assert len(messages_n_plus_one) == len(messages_n) + 2 + + +def test_signoz_dispatch_prefers_otel_v2_when_flag_on(monkeypatch): + from litellm.integrations.otel.logger import OpenTelemetryV2 + from litellm.integrations.otel.model.config import ExporterOwner, is_otel_v2_enabled + from litellm.litellm_core_utils import litellm_logging as logging_module + + logging_module._in_memory_loggers.clear() + monkeypatch.setenv("LITELLM_OTEL_V2", "true") + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "https://ingest.eu.signoz.cloud:443") + monkeypatch.setenv("SIGNOZ_INGESTION_KEY", "test-key") + is_otel_v2_enabled.cache_clear() + try: + v2_logger = logging_module._init_custom_logger_compatible_class( + logging_integration="signoz", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert isinstance(v2_logger, OpenTelemetryV2) + assert v2_logger.callback_name == "signoz" + spec = next(e for e in v2_logger.config.exporters if e.owner == ExporterOwner.SIGNOZ) + assert spec.endpoint == "https://ingest.eu.signoz.cloud:443" + assert spec.headers == "signoz-ingestion-key=test-key" + again = logging_module._init_custom_logger_compatible_class( + logging_integration="signoz", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert again is v2_logger + finally: + logging_module._in_memory_loggers.clear() + monkeypatch.delenv("LITELLM_OTEL_V2", raising=False) + is_otel_v2_enabled.cache_clear() + + +def test_signoz_dispatch_keeps_legacy_otel_when_flag_off(monkeypatch): + from litellm.integrations.opentelemetry import OpenTelemetry + from litellm.integrations.otel.model.config import is_otel_v2_enabled + from litellm.litellm_core_utils import litellm_logging as logging_module + + logging_module._in_memory_loggers.clear() + monkeypatch.delenv("LITELLM_OTEL_V2", raising=False) + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "http://signoz-collector.internal:4318") + monkeypatch.setenv("SIGNOZ_INGESTION_KEY", "legacy-key") + monkeypatch.delenv("OTEL_EXPORTER_OTLP_TRACES_HEADERS", raising=False) + is_otel_v2_enabled.cache_clear() + try: + legacy = logging_module._init_custom_logger_compatible_class( + logging_integration="signoz", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert isinstance(legacy, OpenTelemetry) + assert legacy.callback_name == "signoz" + assert legacy.config.endpoint == "http://signoz-collector.internal:4318/v1/traces" + assert legacy.config.headers == "signoz-ingestion-key=legacy-key" + assert "OTEL_EXPORTER_OTLP_TRACES_HEADERS" not in os.environ + # Same name resolves to the same instance, not a second exporter. + again = logging_module._init_custom_logger_compatible_class( + logging_integration="signoz", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert again is legacy + finally: + logging_module._in_memory_loggers.clear() + is_otel_v2_enabled.cache_clear() + + +def test_signoz_dispatch_requires_an_endpoint(monkeypatch): + from litellm.integrations.otel.model.config import is_otel_v2_enabled + from litellm.litellm_core_utils import litellm_logging as logging_module + + logging_module._in_memory_loggers.clear() + monkeypatch.setenv("LITELLM_OTEL_V2", "true") + monkeypatch.delenv("SIGNOZ_INGESTION_ENDPOINT", raising=False) + monkeypatch.delenv("SIGNOZ_INGESTION_KEY", raising=False) + is_otel_v2_enabled.cache_clear() + try: + created = logging_module._init_custom_logger_compatible_class( + logging_integration="signoz", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert created is None + assert not [ + cb for cb in logging_module._in_memory_loggers if getattr(cb, "callback_name", None) == "signoz" + ] + finally: + logging_module._in_memory_loggers.clear() + monkeypatch.delenv("LITELLM_OTEL_V2", raising=False) + is_otel_v2_enabled.cache_clear() diff --git a/ui/litellm-dashboard/public/assets/logos/signoz.svg b/ui/litellm-dashboard/public/assets/logos/signoz.svg new file mode 100644 index 00000000000..9064cb86bd6 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/signoz.svg @@ -0,0 +1 @@ + \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/callback_info_helpers.tsx b/ui/litellm-dashboard/src/components/callback_info_helpers.tsx index bc9889da724..19020d92066 100644 --- a/ui/litellm-dashboard/src/components/callback_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/callback_info_helpers.tsx @@ -10,6 +10,7 @@ import newrelicLogo from "../../public/assets/logos/newrelic.png"; import openmeterLogo from "../../public/assets/logos/openmeter.png"; import otelLogo from "../../public/assets/logos/otel.png"; import pointfiveLogo from "../../public/assets/logos/pointfive.png"; +import signozLogo from "../../public/assets/logos/signoz.svg"; import databricksLogo from "../../public/assets/logos/databricks.svg"; interface CallbackConfig { @@ -209,6 +210,17 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ }, description: "S3 Bucket (AWS) Logging Integration", }, + { + id: "signoz", + displayName: "SigNoz", + logo: signozLogo.src, + supports_key_team_logging: true, + dynamic_params: { + signoz_ingestion_endpoint: "text", + signoz_ingestion_key: "password", + }, + description: "SigNoz Logging Integration. Setup: https://signoz.io/docs/litellm-observability/", + }, { id: "SQS", displayName: "SQS", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 6d45f6e691c..2b8ed9aa58d 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -57081,7 +57081,7 @@ export interface operations { parameters: { query: { /** @description Specify the service being hit. */ - service: ("slack_budget_alerts" | "langfuse" | "langfuse_otel" | "slack" | "ms_teams" | "openmeter" | "webhook" | "email" | "braintrust" | "datadog" | "datadog_llm_observability" | "generic_api" | "arize" | "galileo" | "newrelic" | "pointfive" | "sqs") | string; + service: ("slack_budget_alerts" | "langfuse" | "langfuse_otel" | "slack" | "ms_teams" | "openmeter" | "webhook" | "email" | "braintrust" | "datadog" | "datadog_llm_observability" | "generic_api" | "arize" | "galileo" | "newrelic" | "pointfive" | "signoz" | "sqs") | string; }; header?: never; path?: never; From 013d5fa0150cd010d82eaa84c0007ec6a8dda296 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 26 Sep 2026 18:46:36 -0700 Subject: [PATCH 09/65] feat(cli): reuse saved agent setup and add reconfigure (#43392) --- litellm/proxy/client/cli/README.md | 21 +- .../proxy/client/cli/commands/configure.py | 561 +++++++---------- .../client/cli/commands/configure_profiles.py | 158 +++++ .../client/cli/commands/configure_setup.py | 439 ++++++++++++++ litellm/proxy/client/cli/main.py | 8 +- .../client/cli/test_configure_commands.py | 566 +++++++++++++++++- 6 files changed, 1391 insertions(+), 362 deletions(-) create mode 100644 litellm/proxy/client/cli/commands/configure_profiles.py create mode 100644 litellm/proxy/client/cli/commands/configure_setup.py diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index a02d7cce0d8..8c05264c6e6 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -546,7 +546,20 @@ lite configure --api-key sk-... --gateway-url https://your-proxy.example.com Select Claude Code, Codex, or both, then choose a gateway model for each selected agent. The wizard validates the key and reads the models your key can access before changing settings. Start either configured agent normally with `claude` or `codex`; the gateway connection persists across terminals without a wrapper or exported API key -`--gateway-url` also accepts a deployment path prefix and a trailing `/v1`. `--base-url` is an alias. If omitted, setup uses `lite --base-url`, `LITELLM_PROXY_URL`, or the saved CLI URL; the wizard asks for a URL when none was provided +Your gateway, virtual key and model choice are saved separately from each agent's undo record. The command saves validated choices before applying them; if applying fails, `lite configure` retries those saved choices. Disconnecting keeps that setup so you can reconnect without repeating the wizard: + +```bash +lite unconfigure +lite configure +``` + +`lite configure` reuses all saved setups for the current agent homes. Name an agent to reconnect only that one, such as `lite configure claude` or `lite configure codex`. Saved setup works without a terminal when the key and model are still valid + +Run `lite reconfigure` to edit your choices with the saved values prefilled, or `lite reconfigure codex` to edit one agent. The agent picker selects which setups to edit; unchecked agents keep their settings. For Claude Code, choose its own default in the wizard or use `lite configure claude --default-model` to remove LiteLLM's model pin. Omitting `--model` keeps your saved choice + +`lite unconfigure --forget` undoes settings it still owns and deletes the saved setups, including their saved keys. `lite unconfigure claude --forget` forgets only Claude Code. Both work after an earlier disconnect. Saved setup files have owner-only permissions and follow the same resolved config-file scope as the undo records, including `CLAUDE_CONFIG_DIR` and `CODEX_HOME`. A pending undo record remains available if an original credential could not safely be restored. If the undo record is missing, agent settings are left unchanged and the command asks you to remove any remaining gateway connection and key manually; forgetting the saved profile does not erase unowned agent settings + +`--gateway-url` also accepts a deployment path prefix and a trailing `/v1`. `--base-url` is an alias. Current command-line or environment options override saved setup values. Otherwise an existing setup supplies its own gateway; first setup falls back to the saved CLI URL or prompts for one. Changing the gateway requires a key for that gateway, so an old saved key is never reused for a different destination For a scripted setup, name the agent and model: @@ -569,9 +582,9 @@ lite --base-url https://your-proxy.example.com configure claude --api-key sk-... claude ``` -The key comes from `--api-key` (or `lite --api-key` / `LITELLM_PROXY_API_KEY`) and is written into `env.ANTHROPIC_AUTH_TOKEN`; without one the command refuses, since a `lite login` credential expires within a day and keeping it fresh would mean Claude Code running `lite` through `apiKeyHelper` on every credential refresh. The command checks the key against `GET /v1/models`, then patches `~/.claude/settings.json`: `env.ANTHROPIC_BASE_URL`, the credential, and `env.ENABLE_TOOL_SEARCH` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` when those are missing, so Claude Code's `/model` picker lists the proxy's models (under `claude-router-` for a group whose id contains neither `claude` nor `anthropic`, since Claude Code lists only those) and you pick between them as usual. Claude Code keeps its own default model until you switch, so that id has to exist on the proxy for the first message to go through; `--model` (or the interactive prompt below) sets the model Claude Code starts on instead, as the top-level `model` key and as `env.ANTHROPIC_MODEL`, both of which have to be on `/v1/models` for the key. The second one matters for `claude -c` and `claude --resume`: a resumed session otherwise re-sends the model its transcript recorded, which behind an auto-router with `return_raw_model_name: true` is the tier model that answered, and a key scoped to the router alias gets a 403 for it; `ANTHROPIC_MODEL` outranks the transcript on resume. Nothing forces Claude Code's sub-agent or background tiers onto a proxy model, so those built-in ids need to exist on the proxy too; `lite autoroute start` is the mode that pins every tier to one group. Claude Code treats a name it does not know as an unknown model: it prints a one-line `unrecognized_model` note, assumes a 200k context window (the proxy appends `[1m]` for a group whose configured or known input window reaches 1M) and sends no thinking parameters for it, so name the group like a Claude model id to change that. The other credential slots (`env.ANTHROPIC_API_KEY`, a stale `env.ANTHROPIC_AUTH_TOKEN` or `apiKeyHelper`) are removed so they cannot fight the one written. Every other setting is preserved and the file is written atomically with owner-only permissions; if `settings.json` is a symlink into a dotfiles repository, the key is written through to that target and the command says so, so keep it out of version control +The key comes from `--api-key` (or `lite --api-key` / `LITELLM_PROXY_API_KEY`), or from this agent's saved setup, and is written into `env.ANTHROPIC_AUTH_TOKEN`; without one the command refuses, since a `lite login` credential expires within a day and keeping it fresh would mean Claude Code running `lite` through `apiKeyHelper` on every credential refresh. The command checks the key against `GET /v1/models`, then patches `~/.claude/settings.json`: `env.ANTHROPIC_BASE_URL`, the credential, and `env.ENABLE_TOOL_SEARCH` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` when those are missing, so Claude Code's `/model` picker lists the proxy's models (under `claude-router-` for a group whose id contains neither `claude` nor `anthropic`, since Claude Code lists only those) and you pick between them as usual. On first setup without a model choice, Claude Code keeps its own default model until you switch, so that id has to exist on the proxy for the first message to go through; `--model` (or the interactive prompt below) sets the model Claude Code starts on instead, as the top-level `model` key and as `env.ANTHROPIC_MODEL`, both of which have to be on `/v1/models` for the key. The second one matters for `claude -c` and `claude --resume`: a resumed session otherwise re-sends the model its transcript recorded, which behind an auto-router with `return_raw_model_name: true` is the tier model that answered, and a key scoped to the router alias gets a 403 for it; `ANTHROPIC_MODEL` outranks the transcript on resume. Nothing forces Claude Code's sub-agent or background tiers onto a proxy model, so those built-in ids need to exist on the proxy too; `lite autoroute start` is the mode that pins every tier to one group. Claude Code treats a name it does not know as an unknown model: it prints a one-line `unrecognized_model` note, assumes a 200k context window (the proxy appends `[1m]` for a group whose configured or known input window reaches 1M) and sends no thinking parameters for it, so name the group like a Claude model id to change that. The other credential slots (`env.ANTHROPIC_API_KEY`, a stale `env.ANTHROPIC_AUTH_TOKEN` or `apiKeyHelper`) are removed so they cannot fight the one written. Every other setting is preserved and the file is written atomically with owner-only permissions; if `settings.json` is a symlink into a dotfiles repository, the key is written through to that target and the command says so, so keep it out of version control -Plain `lite configure`, with no agent named, asks which agents to wire and which gateway model each starts on, picked from `/v1/models` with a type-to-filter prompt. All choices and selected config files are checked before the first settings write. If a later filesystem write fails, the output identifies each agent already configured and its undo command +On first use, plain `lite configure` asks which agents to wire and which gateway model each starts on, picked from `/v1/models` with a type-to-filter prompt. Later runs validate and reapply the saved setups. All choices and selected config files are checked before the first settings write. If a later filesystem write fails, the output identifies each agent already configured and its undo command What the command changed is recorded in `~/.litellm/claude_configure_state.json` (previous values plus fingerprints of what was written, never a second copy of the key). `lite unconfigure claude` restores each of those keys only if it still holds what `configure` wrote, so anything you changed since is left alone and named in the output; a `settings.json` or `env` object that only existed because of `configure` is removed again. Ownership moves only by a write: running `configure` again (a re-login is one) refreshes the record only for the keys its merge changed, keeps the original snapshot of a key that still holds what it wrote, and snapshots afresh a key you changed in between, so `unconfigure` brings back whatever the repeat displaced and never adopts your edit as its own. A credential (`env.ANTHROPIC_API_KEY`, `env.ANTHROPIC_AUTH_TOKEN`, `apiKeyHelper`) is put back only when the restored file points at the `ANTHROPIC_BASE_URL` it was captured next to; otherwise it stays removed, the output says which server it belonged to, and the receipt is kept so pointing the URL back and running `unconfigure` again finishes the job. It also undoes `lite login --config-claude`, which writes through the same path. Both refuse to run while a `lite up` or `lite autoroute start` session holds a backup, and that check comes before any request @@ -587,7 +600,7 @@ Claude Opus 5 ██████████████████████ After the first response, the status line uses the latest routed model recorded by `GET /auto_router/session?session_id=...`, so it can show the tier model even when the transcript contains the router alias. If no session record is available, it falls back to Claude Code's transcript. Session records and costs are cached for five seconds under a per-user `$TMPDIR/litellm-statusline-` directory. The gateway records turns asynchronously, so the display can briefly lag a completed turn. Any virtual key may read its own sessions. The baseline is the priciest model in the router's hardest tier, the same counterfactual the auto-router's savings reports use. `lite unconfigure claude` removes the `statusLine` entry only while it still points at that script -After upgrading the CLI, rerun your original `lite configure claude` command with the same gateway, key and model choice to refresh `~/.litellm/statusline.py`. Keep any explicit `--model` value: omitting it removes the earlier model pin. Package upgrades alone do not refresh this installed copy +After upgrading the CLI, run `lite configure claude` to refresh `~/.litellm/statusline.py` using the saved setup. If your setup predates saved profiles, supply the original gateway, key and model once. Package upgrades alone do not refresh this installed copy `lite codex` registers the same script as a Codex `Stop` hook for the launch, so after each turn Codex prints the same block as a system message. Codex asks once to trust the hook; the answer is remembered for later launches. diff --git a/litellm/proxy/client/cli/commands/configure.py b/litellm/proxy/client/cli/commands/configure.py index eca7ba86496..ea24a644084 100644 --- a/litellm/proxy/client/cli/commands/configure.py +++ b/litellm/proxy/client/cli/commands/configure.py @@ -1,284 +1,43 @@ -"""Persistent Claude Code and Codex gateway configuration.""" +"""Commands for saved Claude Code and Codex gateway setup.""" -import os import sys -from collections.abc import Callable, Sequence -from dataclasses import dataclass from pathlib import Path -from types import MappingProxyType from typing import Final import click from InquirerPy import inquirer -from InquirerPy.base.control import Choice from pydantic import BaseModel -from litellm.proxy.common_utils.model_listing_utils import ( - CLAUDE_CODE_CLIENT, - CLAUDE_CODE_PICKER_PATTERN, - GATEWAY_CLIENT_HEADER, -) - -from .agents import codex_config_path from .auth import CliContextObj from .claude_settings import ( - STARTING_MODEL_ROLE, ClaudeSettingsError, - ModelChoice, - StartOn, - StaticToken, UnconfigureOutcome, - UnpinModel, - claude_settings_path, - configure_claude_settings, - configure_state_path, preflight_claude_settings, settings_file_owners, unconfigure_claude_settings, ) -from .codex_settings import ( - CodexSettingsError, - configure_codex_settings, - preflight_codex_settings, - unconfigure_codex_settings, -) +from .codex_settings import CodexSettingsError, unconfigure_codex_settings from .config import normalize_base_url -from .pi import ListedModel, ListingFailure, PiSyncError, fetch_model_listing - -_LISTED_MODELS_SHOWN: Final = 20 -_CLAUDE_TARGET: Final = "claude" -_CODEX_TARGET: Final = "codex" -_TARGETS: Final = ((_CLAUDE_TARGET, "Claude Code (CLI)"), (_CODEX_TARGET, "Codex (CLI)")) -_KEEP_DEFAULT_MODEL: Final = "Keep Claude Code's own default" -_CLAUDE_CODE_VIEW: Final = MappingProxyType( - {"anthropic-version": "2023-06-01", GATEWAY_CLIENT_HEADER: CLAUDE_CODE_CLIENT} +from .configure_profiles import ( + TARGETS, + Target, + forget_saved_setup, + read_saved_setup, + receipt_path_for, + settings_path_for, + setup_locks, + setup_profile_path, ) -_MODEL_OPTION_HELP: Final = ( - f"Proxy model to set as {STARTING_MODEL_ROLE}. Must be listed on /v1/models for the key; without it, " - "Claude Code keeps its own default and a pin an earlier configure made is let go of. Nothing pins Claude " - "Code's sub-agent or background tiers; `lite autoroute start` is the mode that does." +from .configure_setup import ( + MODEL_OPTION_HELP, + ConnectionSettings, + configure_targets, + interactive_configure, + pick_targets, + resolve_credential, ) -def resolve_credential(ctx: click.Context, api_key: str | None) -> StaticToken: - """The long-lived key written into settings.json: --api-key, `lite --api-key` or LITELLM_PROXY_API_KEY. - - A `lite login` credential is never written: it expires within a day, and keeping it fresh would mean - Claude Code running `lite` through `apiKeyHelper` on every credential refresh. - """ - ctx_obj: Final[CliContextObj] = ctx.obj - explicit: Final = api_key or (None if ctx_obj.get("api_key_from_token_file") else ctx_obj.get("api_key")) - if not explicit: - raise ClaudeSettingsError( - "`lite configure` needs a long-lived virtual key: pass --api-key, `lite --api-key`, or set " - "LITELLM_PROXY_API_KEY. Your `lite login` credential expires within a day, so it is not written " - "into agent settings." - ) - if not explicit.strip() or any(ord(char) <= 32 or ord(char) == 127 for char in explicit): - raise ClaudeSettingsError("The virtual key must not be blank or contain whitespace or control characters.") - return StaticToken(explicit) - - -@dataclass(frozen=True, slots=True) -class _Listing: - models: tuple[ListedModel, ...] - - @property - def ids(self) -> tuple[str, ...]: - return tuple(model.id for model in self.models) - - -def _preflight(target: str) -> None: - try: - if target == _CLAUDE_TARGET: - preflight_claude_settings(claude_settings_path(os.environ)) - else: - preflight_codex_settings(codex_config_path(os.environ)) - except (ClaudeSettingsError, CodexSettingsError) as e: - raise click.ClickException(str(e)) from e - - -def _start( - ctx: click.Context, base_url: str, api_key: str | None, target: str = _CLAUDE_TARGET -) -> tuple[StaticToken, _Listing]: - _preflight(target) - try: - credential: Final = resolve_credential(ctx, api_key) - except ClaudeSettingsError as e: - raise click.ClickException(str(e)) - return credential, _listed_models(base_url, credential.token, target) - - -def _listing_error(base_url: str, error: PiSyncError, target: str) -> str: - """The hint that fits how the listing failed: only an unreachable proxy gets the "is it running" question.""" - if error.kind is ListingFailure.REJECTED: - return f"LiteLLM rejected your key (HTTP {error.status}). Pass a valid --api-key." - if error.kind is ListingFailure.UNREACHABLE: - return ( - f"Could not connect. Is the proxy at {base_url} running, and is --base-url (or LITELLM_PROXY_URL) correct?" - ) - if error.kind is ListingFailure.EMPTY: - name: Final = "Claude Code" if target == _CLAUDE_TARGET else "Codex" - return f"{error.message} {name} would have nothing to run; give the key access to at least one model." - return f"The proxy at {base_url} answered, so check that it is a LiteLLM proxy and is healthy." - - -def _listed_models(base_url: str, key: str, target: str = _CLAUDE_TARGET) -> _Listing: - listed: Final = fetch_model_listing( - base_url, key, headers=_CLAUDE_CODE_VIEW if target == _CLAUDE_TARGET else MappingProxyType({}) - ) - if isinstance(listed, PiSyncError): - raise click.ClickException(_listing_error(base_url, listed, target)) - return _Listing(listed) - - -def _starting_model(model: str, listing: _Listing) -> str | None: - source: Final = next((listed.id for listed in listing.models if listed.source_model == model), None) - return source or next((listed.id for listed in listing.models if listed.id == model), None) - - -def _model_choice(model: str | None) -> ModelChoice: - return StartOn(model) if model is not None else UnpinModel() - - -def _validated_model(model: str | None, listing: _Listing, base_url: str) -> str | None: - starting: Final = _starting_model(model, listing) if model is not None else None - if model is not None and starting is None: - shown: Final = ", ".join(listing.ids[:_LISTED_MODELS_SHOWN]) - raise click.ClickException(f"{model!r} is not served by {base_url} for this key. /v1/models lists: {shown}.") - return starting - - -def _apply_claude(base_url: str, credential: StaticToken, listing: _Listing, model: str | None) -> None: - listed: Final = listing.ids - starting: Final = _validated_model(model, listing, base_url) - settings_path: Final = claude_settings_path(os.environ) - try: - configure_claude_settings( - base_url, - credential, - _model_choice(starting), - settings_path, - configure_state_path(settings_path), - settings_file_owners(settings_path), - ) - except ClaudeSettingsError as e: - raise click.ClickException(str(e)) - in_picker: Final = sum(1 for listed_model in listed if CLAUDE_CODE_PICKER_PATTERN.search(listed_model)) - click.echo(f"Configured Claude Code: {settings_path} now routes through {base_url}.") - - click.echo("Credential: your virtual key, stored in the file as ANTHROPIC_AUTH_TOKEN.") - click.echo( - f"Starting model: {starting} ({STARTING_MODEL_ROLE}); switch any time with /model." - if starting is not None - else "Starting model: not pinned (Claude Code's default, or a model you set yourself); switch with /model, or " - "pass --model to start on a proxy model. Without a pin, a resumed session re-sends the model its transcript " - "recorded, which behind a raw-model auto-router is the tier model." - ) - click.echo( - f"/model will list all {len(listed)} of the proxy's models." - if in_picker == len(listed) - else f"/model will list {in_picker} of the proxy's {len(listed)} models: Claude Code shows only ids containing " - "'claude' or 'anthropic', and this proxy does not list the rest under such names." - ) - click.echo("Start `claude` from any terminal. Undo with `lite unconfigure claude`.") - if settings_path.is_symlink(): - click.echo( - f"Note: {settings_path} is a symlink to {settings_path.resolve()}, so your key now lives in " - "that file; keep it out of version control.", - err=True, - ) - - -def _pick_targets() -> tuple[str, ...]: - picked: Final = inquirer.checkbox( - message="Which agents should route through LiteLLM?", - choices=[Choice(value, name=label, enabled=True) for value, label in _TARGETS], - validate=lambda chosen: len(chosen) > 0, - invalid_message="Pick at least one.", - ).execute() - return tuple(str(value) for value in picked) - - -def _pick_model(listed: Sequence[str]) -> str | None: - picked: Final = inquirer.fuzzy( - message="Model Claude Code starts on (type to filter; /model switches any time):", - choices=[_KEEP_DEFAULT_MODEL, *listed], - default=listed[0] if listed else _KEEP_DEFAULT_MODEL, - ).execute() - return None if picked == _KEEP_DEFAULT_MODEL else str(picked) - - -def _pick_codex_model(listed: Sequence[str]) -> str: - choices: Final = list(listed) # mutable-ok: InquirerPy's choices parameter requires a list - return str(inquirer.fuzzy(message="Model Codex starts on (type to filter):", choices=choices).execute()) - - -def _apply_codex(base_url: str, credential: StaticToken, listing: _Listing, model: str) -> None: - _validated_model(model, listing, base_url) - settings_path: Final = codex_config_path(os.environ) - try: - configure_codex_settings(base_url, credential.token, model, settings_path) - except CodexSettingsError as e: - raise click.ClickException(str(e)) from e - click.echo(f"Configured Codex: {settings_path} now routes through {base_url}.") - click.echo(f"Starting model: {model}. Credential: your virtual key, stored in the private provider settings.") - click.echo("Start `codex` from any terminal. Undo with `lite unconfigure codex`.") - if settings_path.is_symlink(): - click.echo(f"Note: your key now lives in {settings_path.resolve()}; keep it out of version control.", err=True) - - -@dataclass(frozen=True, slots=True) -class _Setup: - target: str - listing: _Listing - model: str | None - - -def _choose_setup( - base_url: str, - target: str, - credential: StaticToken, - pick_model: Callable[[Sequence[str]], str | None], - pick_codex_model: Callable[[Sequence[str]], str], -) -> _Setup: - listing: Final = _listed_models(base_url, credential.token, target) - model: Final = ( - pick_model(tuple(item.source_model or item.id for item in listing.models)) - if target == _CLAUDE_TARGET - else pick_codex_model(listing.ids) - ) - _validated_model(model, listing, base_url) - return _Setup(target, listing, model) - - -def interactive_configure( - ctx: click.Context, - pick_targets: Callable[[], tuple[str, ...]] = _pick_targets, - pick_model: Callable[[Sequence[str]], str | None] = _pick_model, - pick_codex_model: Callable[[Sequence[str]], str] = _pick_codex_model, -) -> None: - """`lite configure` with no agent named: ask which agents to wire and which model to pin.""" - targets: Final = pick_targets() - if not targets: - return - for target in targets: - _preflight(target) - try: - credential: Final = resolve_credential(ctx, None) - except ClaudeSettingsError as e: - raise click.ClickException(str(e)) from e - base_url: Final[str] = ctx.obj["base_url"] - setups: Final = tuple( - _choose_setup(base_url, target, credential, pick_model, pick_codex_model) for target in targets - ) - for setup in setups: - if setup.target == _CLAUDE_TARGET: - _apply_claude(base_url, credential, setup.listing, setup.model) - elif setup.model is not None: - _apply_codex(base_url, credential, setup.listing, setup.model) - - class _ConnectionOptions(BaseModel): api_key: str | None = None gateway_url: str | None = None @@ -286,21 +45,20 @@ class _ConnectionOptions(BaseModel): def _connection_settings(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> CliContextObj: """The context object a subcommand runs with: its own --api-key / --gateway-url over the group's, over `lite`'s.""" - ctx_obj: Final[CliContextObj] = ctx.obj + ctx_obj: Final = ConnectionSettings.model_validate(ctx.find_object(object)) group: Final = ( _ConnectionOptions.model_validate(ctx.parent.params) - if ctx.parent is not None and ctx.parent.command.name == "configure" + if ctx.parent is not None and ctx.parent.command.name in ("configure", "reconfigure") else _ConnectionOptions() ) key: Final = api_key if api_key is not None else group.api_key url: Final = gateway_url if gateway_url is not None else group.gateway_url - normalized: Final = normalize_base_url(url if url is not None else ctx_obj["base_url"]) + normalized: Final = normalize_base_url(url if url is not None else ctx_obj.base_url) connection: Final[CliContextObj] = { - **ctx_obj, "base_url": normalized.removesuffix("/v1"), - "base_url_explicit": url is not None or ctx_obj.get("base_url_explicit", False), - "api_key": key if key is not None else ctx_obj.get("api_key"), - "api_key_from_token_file": False if key is not None else ctx_obj.get("api_key_from_token_file", False), + "base_url_explicit": url is not None or ctx_obj.base_url_explicit, + "api_key": key if key is not None else ctx_obj.api_key, + "api_key_from_token_file": False if key is not None else ctx_obj.api_key_from_token_file, } return connection @@ -309,108 +67,216 @@ def _connection_context(ctx: click.Context, settings: CliContextObj) -> click.Co return click.Context(ctx.command, parent=ctx.parent, obj=settings) -@click.group(name="configure", invoke_without_command=True) -@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key to store in the selected agents.") -@click.option( - "--gateway-url", "--base-url", default=None, help="Gateway URL; defaults to `lite --base-url` / LITELLM_PROXY_URL." -) -@click.pass_context -def configure_group(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> None: - """Persistently route a coding agent through your LiteLLM proxy. +def _require_terminal(command: str) -> None: + if sys.stdin.isatty(): + return + raise click.ClickException( + f"`lite {command}` asks questions, so it needs a terminal. Non-interactively, run " + f"`lite {command} claude --api-key --model ` or " + f"`lite {command} codex --api-key --model `" + ) - With no agent named, asks which agents to wire and which proxy model to pin. - """ + +def _configure_group(ctx: click.Context, api_key: str | None, gateway_url: str | None, edit: bool) -> None: if ctx.invoked_subcommand is not None: return settings: Final = _connection_settings(ctx, api_key, gateway_url) connection: Final = _connection_context(ctx, settings) - if not sys.stdin.isatty(): - raise click.ClickException( - "`lite configure` asks questions, so it needs a terminal. Non-interactively, run " - "`lite configure claude --api-key --model ` or " - "`lite configure codex --api-key --model `." + with setup_locks(TARGETS): + saved_targets: Final[tuple[Target, ...]] = tuple( + target for target in TARGETS if read_saved_setup(target) is not None ) - if settings.get("base_url_explicit"): - interactive_configure(connection) - return - prompted: Final = _connection_settings(connection, None, click.prompt("Gateway URL", default=settings["base_url"])) - interactive_configure(_connection_context(connection, prompted)) + if saved_targets and not edit: + configure_targets(connection, saved_targets) + return + _require_terminal("reconfigure" if edit else "configure") + selected: Final = pick_targets(saved_targets or TARGETS, edit=edit) + if not selected: + return + if edit: + configure_targets(connection, selected, interactive=True, edit_connection=True) + return + prompted: Final = ( + settings + if settings.get("base_url_explicit") + else _connection_settings(connection, None, click.prompt("Gateway URL", default=settings["base_url"])) + ) + configure_targets(_connection_context(connection, prompted), selected, interactive=True) -@click.group(name="unconfigure") -def unconfigure_group() -> None: - """Undo `lite configure` for a coding agent.""" +@click.group(name="configure", invoke_without_command=True) +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key to save for the selected agents.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, including any deployment path prefix.") +@click.pass_context +def configure_group(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> None: + """Apply saved setup, or choose agents and models on the first run.""" + _configure_group(ctx, api_key, gateway_url, False) + + +@click.group(name="reconfigure", invoke_without_command=True) +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key to save for the selected agents.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, including any deployment path prefix.") +@click.pass_context +def reconfigure_group(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> None: + """Edit saved gateway, key and model choices, using current choices as defaults.""" + _configure_group(ctx, api_key, gateway_url, True) + + +def _configure_target( + ctx: click.Context, + target: Target, + api_key: str | None, + gateway_url: str | None, + model: str | None, + default_model: bool = False, + *, + edit: bool = False, +) -> None: + settings: Final = _connection_settings(ctx, api_key, gateway_url) + interactive: Final = edit and model is None and not default_model + if interactive: + _require_terminal("reconfigure") + with setup_locks((target,)): + configure_targets( + _connection_context(ctx, settings), + (target,), + model=model, + default_model=default_model, + interactive=interactive, + edit_connection=interactive, + ) @configure_group.command(name="claude") +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key, or reuse the saved key.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, or reuse the saved gateway.") +@click.option("--model", default=None, help=MODEL_OPTION_HELP) @click.option( - "--api-key", - "api_key", - default=None, - help="Long-lived LiteLLM virtual key written into Claude Code's settings. Defaults to the `lite --api-key` / " - "LITELLM_PROXY_API_KEY value; required, since a `lite login` credential expires within a day.", + "--default-model", is_flag=True, help="Stop pinning a starting model; let Claude Code choose its default." ) -@click.option("--model", default=None, help=_MODEL_OPTION_HELP) -@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, including any deployment path prefix.") @click.pass_context -def configure_claude(ctx: click.Context, api_key: str | None, model: str | None, gateway_url: str | None) -> None: - """Route every Claude Code session through your LiteLLM proxy until `lite unconfigure claude`. - - Patches ~/.claude/settings.json in place: the proxy URL, your virtual key as a static token, - and gateway model discovery so /model lists the proxy's models; --model picks the one Claude - Code starts on and resumes with. Every other - setting is kept, and what changed is recorded so `lite unconfigure claude` can put it back. - Assumes the proxy is already running. - """ - settings: Final = _connection_settings(ctx, api_key, gateway_url) - credential, listing = _start(_connection_context(ctx, settings), settings["base_url"], api_key) - _apply_claude(settings["base_url"], credential, listing, model) +def configure_claude( + ctx: click.Context, + api_key: str | None, + gateway_url: str | None, + model: str | None, + default_model: bool, +) -> None: + """Apply Claude Code's saved setup, or save the supplied settings.""" + _configure_target(ctx, "claude", api_key, gateway_url, model, default_model) @configure_group.command(name="codex") -@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key to store in Codex's user config.") -@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, including any deployment path prefix.") -@click.option("--model", required=True, help="Gateway model Codex starts on, as listed by /v1/models for your key.") +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key, or reuse the saved key.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, or reuse the saved gateway.") +@click.option("--model", default=None, help="Gateway model to start on; required only for first-time setup.") @click.pass_context -def configure_codex(ctx: click.Context, api_key: str | None, gateway_url: str | None, model: str) -> None: - """Route plain `codex` through the gateway until `lite unconfigure codex`.""" - settings: Final = _connection_settings(ctx, api_key, gateway_url) - credential, listing = _start(_connection_context(ctx, settings), settings["base_url"], api_key, _CODEX_TARGET) - _apply_codex(settings["base_url"], credential, listing, model) +def configure_codex(ctx: click.Context, api_key: str | None, gateway_url: str | None, model: str | None) -> None: + """Apply Codex's saved setup, or save the supplied settings.""" + _configure_target(ctx, "codex", api_key, gateway_url, model) -@unconfigure_group.command(name="codex") -def unconfigure_codex() -> None: - """Restore only Codex settings still holding what configure wrote.""" - settings_path: Final = codex_config_path(os.environ) +@reconfigure_group.command(name="claude") +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key, or reuse the saved key.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, or reuse the saved gateway.") +@click.option("--model", default=None, help=MODEL_OPTION_HELP) +@click.option( + "--default-model", is_flag=True, help="Stop pinning a starting model; let Claude Code choose its default." +) +@click.pass_context +def reconfigure_claude( + ctx: click.Context, + api_key: str | None, + gateway_url: str | None, + model: str | None, + default_model: bool, +) -> None: + """Edit Claude Code setup, or supply --model / --default-model to apply directly.""" + _configure_target(ctx, "claude", api_key, gateway_url, model, default_model, edit=True) + + +@reconfigure_group.command(name="codex") +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key, or reuse the saved key.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, or reuse the saved gateway.") +@click.option("--model", default=None, help="Gateway model to start on; omit to open the setup wizard.") +@click.pass_context +def reconfigure_codex(ctx: click.Context, api_key: str | None, gateway_url: str | None, model: str | None) -> None: + """Edit Codex setup, or supply --model to apply directly.""" + _configure_target(ctx, "codex", api_key, gateway_url, model, edit=True) + + +def _disconnect(target: Target, forget: bool) -> None: + settings_path: Final = settings_path_for(target) + state_path: Final = receipt_path_for(target, settings_path) + profile: Final = setup_profile_path(target, settings_path) try: - outcome: Final = unconfigure_codex_settings(settings_path) - except CodexSettingsError as e: - raise click.ClickException(str(e)) from e - if outcome.file_removed: - click.echo(f"Removed {settings_path}; it held only settings created by `lite configure codex`.") - elif outcome.restored: - click.echo(f"Restored in {settings_path}: {', '.join(outcome.restored)}.") - else: - click.echo(f"Nothing in {settings_path} was still ours to restore.") - if outcome.kept: - click.echo(f"Left as you changed them since: {', '.join(outcome.kept)}.") + if not state_path.exists(): + click.echo(f"No {target} undo receipt at {state_path}; nothing to undo. Agent settings were not changed.") + if settings_path.exists(): + click.echo( + f"Cannot confirm disconnection. Check {settings_path} and remove any remaining gateway " + "connection and key manually.", + err=True, + ) + elif target == "claude": + preflight_claude_settings(settings_path) + outcome: Final = unconfigure_claude_settings(settings_path, state_path, settings_file_owners(settings_path)) + _report_unconfigure(settings_path, state_path, outcome) + else: + codex_outcome: Final = unconfigure_codex_settings(settings_path) + if codex_outcome.file_removed: + click.echo(f"Removed {settings_path}; it held only settings created by `lite configure codex`.") + elif codex_outcome.restored: + click.echo(f"Restored in {settings_path}: {', '.join(codex_outcome.restored)}.") + else: + click.echo(f"Nothing in {settings_path} was still ours to restore.") + if codex_outcome.kept: + click.echo(f"Left as you changed them since: {', '.join(codex_outcome.kept)}.") + except (ClaudeSettingsError, CodexSettingsError) as error: + raise click.ClickException(str(error)) from error + if forget: + forget_saved_setup(target) + click.echo(f"Forgot saved {target} setup, including its saved key.") + elif profile.exists(): + click.echo(f"Saved setup retained. Run `lite configure {target}` to apply it again.") + + +@click.group(name="unconfigure", invoke_without_command=True) +@click.option("--forget", is_flag=True, help="Also delete saved setups and their keys.") +@click.pass_context +def unconfigure_group(ctx: click.Context, forget: bool) -> None: + """Disconnect agents while retaining saved setup for `lite configure`.""" + if ctx.invoked_subcommand is not None: + return + with setup_locks(TARGETS): + for target in TARGETS: + _disconnect(target, forget) + + +class _UnconfigureOptions(BaseModel): + forget: bool = False + + +def _unconfigure_target(ctx: click.Context, target: Target, forget: bool) -> None: + parent: Final = _UnconfigureOptions.model_validate(ctx.parent.params) if ctx.parent else _UnconfigureOptions() + with setup_locks((target,)): + _disconnect(target, forget or parent.forget) @unconfigure_group.command(name="claude") -def unconfigure_claude() -> None: - """Return Claude Code's settings to what they were before `lite configure claude`. +@click.option("--forget", is_flag=True, help="Also delete the saved Claude Code setup and key.") +@click.pass_context +def unconfigure_claude(ctx: click.Context, forget: bool) -> None: + """Restore Claude Code settings, including those applied by `lite login --config-claude`.""" + _unconfigure_target(ctx, "claude", forget) - Also undoes `lite login --config-claude`. Only keys still holding what configure wrote are - put back; anything you changed since is left as it is and named in the output. - """ - settings_path: Final = claude_settings_path(os.environ) - state_path: Final = configure_state_path(settings_path) - try: - outcome: Final = unconfigure_claude_settings(settings_path, state_path, settings_file_owners(settings_path)) - except ClaudeSettingsError as e: - raise click.ClickException(str(e)) - _report_unconfigure(settings_path, state_path, outcome) + +@unconfigure_group.command(name="codex") +@click.option("--forget", is_flag=True, help="Also delete the saved Codex setup and key.") +@click.pass_context +def unconfigure_codex(ctx: click.Context, forget: bool) -> None: + """Restore Codex settings still holding what configure wrote.""" + _unconfigure_target(ctx, "codex", forget) def _report_unconfigure(settings_path: Path, state_path: Path, outcome: UnconfigureOutcome) -> None: @@ -434,4 +300,11 @@ def _report_unconfigure(settings_path: Path, state_path: Path, outcome: Unconfig ) -__all__ = ("configure_group", "interactive_configure", "resolve_credential", "unconfigure_group") +__all__ = ( + "configure_group", + "inquirer", + "interactive_configure", + "reconfigure_group", + "resolve_credential", + "unconfigure_group", +) diff --git a/litellm/proxy/client/cli/commands/configure_profiles.py b/litellm/proxy/client/cli/commands/configure_profiles.py new file mode 100644 index 00000000000..87c4a05dd22 --- /dev/null +++ b/litellm/proxy/client/cli/commands/configure_profiles.py @@ -0,0 +1,158 @@ +"""Reusable agent setup, separate from the settings writers' undo receipts.""" + +import hashlib +import os +from collections.abc import Generator, Sequence +from contextlib import ExitStack, contextmanager +from pathlib import Path +from typing import Final, Literal, TypeAlias + +import click +from filelock import FileLock, Timeout +from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator + +from litellm.litellm_core_utils.private_json import ( + commit_staged_json, + discard_staged_json, + ensure_private_dir, + stage_private_json, +) + +from .agents import codex_config_path +from .claude_settings import claude_settings_path, configure_state_path +from .codex_settings import codex_configure_state_path +from .config import normalize_base_url + +Target: TypeAlias = Literal["claude", "codex"] +TARGETS: Final[tuple[Target, ...]] = ("claude", "codex") + + +class SavedSetup(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid", strict=True) + + version: Literal[1] = 1 + target: Target + settings_path: str + base_url: str + api_key: str = Field(repr=False) + model: str | None + + @field_validator("base_url") + @classmethod + def normalized_gateway(cls, value: str) -> str: + try: + if normalize_base_url(value).removesuffix("/v1") != value: + raise ValueError("Gateway must be normalized") + except click.UsageError as error: + raise ValueError("Invalid gateway URL") from error + return value + + @field_validator("api_key") + @classmethod + def valid_key(cls, value: str) -> str: + if not value or any(ord(char) <= 32 or ord(char) == 127 for char in value): + raise ValueError("Invalid virtual key") + return value + + @field_validator("model") + @classmethod + def valid_model(cls, value: str | None) -> str | None: + if value is not None and (not value or any(ord(char) < 32 or ord(char) == 127 for char in value)): + raise ValueError("Invalid model choice") + return value + + +def settings_path_for(target: Target) -> Path: + return claude_settings_path(os.environ) if target == "claude" else codex_config_path(os.environ) + + +def receipt_path_for(target: Target, settings_path: Path) -> Path: + return configure_state_path(settings_path) if target == "claude" else codex_configure_state_path(settings_path) + + +def setup_profile_path(target: Target, settings_path: Path) -> Path: + receipt: Final = receipt_path_for(target, settings_path) + return receipt.with_name(f"{receipt.stem}_profile.json") + + +def read_saved_setup(target: Target) -> SavedSetup | None: + settings_path: Final = settings_path_for(target) + path: Final = setup_profile_path(target, settings_path) + try: + payload: Final = path.read_bytes() + except FileNotFoundError: + return None + except OSError as error: + raise click.ClickException( + f"Could not read saved {target} setup at {path}; no settings were changed" + ) from error + try: + saved: Final = SavedSetup.model_validate_json(payload) + if ( + saved.target != target + or saved.settings_path != str(settings_path.resolve()) + or (target == "codex" and saved.model is None) + ): + raise ValueError("Invalid saved setup") + return saved + except (ValidationError, ValueError, click.UsageError) as error: + raise click.ClickException( + f"Saved {target} setup at {path} is invalid or unsupported. " + f"Run `lite unconfigure {target} --forget` to discard it; no settings were changed" + ) from error + + +def _lock_path(target: Target) -> Path: + digest: Final = hashlib.sha256(f"{target}:{settings_path_for(target).resolve()}".encode()).hexdigest() + return Path.home() / ".litellm" / "setup-locks" / f"{digest}.lock" + + +@contextmanager +def setup_locks(targets: Sequence[Target]) -> Generator[None, None, None]: + with ExitStack() as stack: + try: + for path in tuple(_lock_path(target) for target in sorted(frozenset(targets))): + ensure_private_dir(path.parent) + stack.enter_context(FileLock(str(path), timeout=10, mode=0o600)) + except (OSError, Timeout) as error: + raise click.ClickException( + "Could not lock agent setup; retry when other configure commands finish" + ) from error + yield + + +def save_setup(saved: SavedSetup) -> None: + path: Final = setup_profile_path(saved.target, settings_path_for(saved.target)) + try: + ensure_private_dir(path.parent) + staged: Final = stage_private_json( + str(path), + { # mutable-ok: private_json serializes with json.dump, which requires a dict + "version": saved.version, + "target": saved.target, + "settings_path": saved.settings_path, + "base_url": saved.base_url, + "api_key": saved.api_key, + "model": saved.model, + }, + ) + except OSError as error: + raise click.ClickException( + f"Could not save {saved.target} setup; no {saved.target} settings were changed" + ) from error + try: + commit_staged_json(staged, str(path)) + except OSError as error: + raise click.ClickException( + f"Could not save {saved.target} setup; no {saved.target} settings were changed" + ) from error + finally: + discard_staged_json(staged) + + +def forget_saved_setup(target: Target) -> None: + path: Final = setup_profile_path(target, settings_path_for(target)) + try: + path.unlink(missing_ok=True) + except OSError as error: + raise click.ClickException(f"Could not remove saved {target} setup at {path}") from error diff --git a/litellm/proxy/client/cli/commands/configure_setup.py b/litellm/proxy/client/cli/commands/configure_setup.py new file mode 100644 index 00000000000..bd07c19dff4 --- /dev/null +++ b/litellm/proxy/client/cli/commands/configure_setup.py @@ -0,0 +1,439 @@ +"""Persistent Claude Code and Codex gateway configuration.""" + +import os +import sys +from collections.abc import Callable, Sequence +from dataclasses import dataclass +from functools import partial +from types import MappingProxyType +from typing import Final + +import click +import requests +from InquirerPy import inquirer +from InquirerPy.base.control import Choice +from pydantic import BaseModel, TypeAdapter, ValidationError + +from litellm.proxy.common_utils.model_listing_utils import ( + CLAUDE_CODE_CLIENT, + CLAUDE_CODE_PICKER_PATTERN, + GATEWAY_CLIENT_HEADER, +) + +from .agents import codex_config_path +from .claude_settings import ( + STARTING_MODEL_ROLE, + ClaudeSettingsError, + ModelChoice, + StartOn, + StaticToken, + UnpinModel, + claude_settings_path, + configure_claude_settings, + configure_state_path, + preflight_claude_settings, + settings_file_owners, +) +from .codex_settings import ( + CodexSettingsError, + configure_codex_settings, + preflight_codex_settings, +) +from .config import normalize_base_url +from .configure_profiles import ( + TARGETS, + SavedSetup, + Target, + read_saved_setup, + save_setup, + settings_path_for, + setup_locks, +) +from .pi import ListedModel, ListingFailure, PiSyncError, fetch_model_listing + +_LISTED_MODELS_SHOWN: Final = 20 +_CLAUDE_TARGET: Final = "claude" +_CODEX_TARGET: Final = "codex" +_TARGETS: Final = ((_CLAUDE_TARGET, "Claude Code (CLI)"), (_CODEX_TARGET, "Codex (CLI)")) +_KEEP_DEFAULT_MODEL: Final = "Keep Claude Code's own default" +_CLAUDE_CODE_VIEW: Final = MappingProxyType( + {"anthropic-version": "2023-06-01", GATEWAY_CLIENT_HEADER: CLAUDE_CODE_CLIENT} +) +MODEL_OPTION_HELP: Final = ( + f"Proxy model to set as {STARTING_MODEL_ROLE}. Must be listed on /v1/models for the key; omission keeps " + "the saved choice. Use --default-model to stop pinning a model. Nothing pins Claude " + "Code's sub-agent or background tiers; `lite autoroute start` is the mode that does." +) +_TARGET_SELECTION: Final = TypeAdapter(tuple[Target, ...]) +_MODEL_SELECTION: Final = TypeAdapter(str) + + +class ConnectionSettings(BaseModel): + base_url: str + base_url_explicit: bool = False + api_key: str | None = None + api_key_from_token_file: bool = False + + +def resolve_credential(ctx: click.Context, api_key: str | None) -> StaticToken: + """The long-lived key written into settings.json: --api-key, `lite --api-key` or LITELLM_PROXY_API_KEY. + + A `lite login` credential is never written: it expires within a day, and keeping it fresh would mean + Claude Code running `lite` through `apiKeyHelper` on every credential refresh. + """ + ctx_obj: Final = ConnectionSettings.model_validate(ctx.find_object(object)) + explicit: Final = api_key if api_key is not None else (None if ctx_obj.api_key_from_token_file else ctx_obj.api_key) + if explicit is None: + raise ClaudeSettingsError( + "`lite configure` needs a long-lived virtual key: pass --api-key, `lite --api-key`, or set " + "LITELLM_PROXY_API_KEY. Your `lite login` credential expires within a day, so it is not written " + "into agent settings." + ) + if not explicit.strip() or any(ord(char) <= 32 or ord(char) == 127 for char in explicit): + raise ClaudeSettingsError("The virtual key must not be blank or contain whitespace or control characters.") + return StaticToken(explicit) + + +@dataclass(frozen=True, slots=True) +class _Listing: + models: tuple[ListedModel, ...] + + @property + def ids(self) -> tuple[str, ...]: + return tuple(model.id for model in self.models) + + +def _preflight(target: Target) -> None: + try: + if target == _CLAUDE_TARGET: + preflight_claude_settings(claude_settings_path(os.environ)) + else: + preflight_codex_settings(codex_config_path(os.environ)) + except (ClaudeSettingsError, CodexSettingsError) as e: + raise click.ClickException(str(e)) from e + + +def _listing_error(base_url: str, error: PiSyncError, target: str) -> str: + """The hint that fits how the listing failed: only an unreachable proxy gets the "is it running" question.""" + if error.kind is ListingFailure.REJECTED: + return f"LiteLLM rejected your key (HTTP {error.status}). Pass a valid --api-key." + if error.kind is ListingFailure.UNREACHABLE: + return ( + f"Could not connect. Is the proxy at {base_url} running, and is --base-url (or LITELLM_PROXY_URL) correct?" + ) + if error.kind is ListingFailure.EMPTY: + name: Final = "Claude Code" if target == _CLAUDE_TARGET else "Codex" + return f"{error.message} {name} would have nothing to run; give the key access to at least one model." + return f"The proxy at {base_url} answered, so check that it is a LiteLLM proxy and is healthy." + + +def _fetch_models(base_url: str, key: str, target: Target) -> tuple[ListedModel, ...] | PiSyncError: + return fetch_model_listing( + base_url, + key, + get=partial(requests.get, allow_redirects=False), + headers=_CLAUDE_CODE_VIEW if target == _CLAUDE_TARGET else MappingProxyType({}), + ) + + +def _connection_listing( + ctx: click.Context, + base_url: str, + credential: StaticToken, + target: Target, + repair: bool, +) -> tuple[StaticToken, _Listing]: + listed: Final = _fetch_models(base_url, credential.token, target) + if not isinstance(listed, PiSyncError): + return credential, _Listing(listed) + if not repair or listed.kind is not ListingFailure.REJECTED: + raise click.ClickException(_listing_error(base_url, listed, target)) + replacement: Final = click.prompt("Replacement virtual key", hide_input=True, show_default=False) + try: + refreshed: Final = resolve_credential(ctx, replacement) + except ClaudeSettingsError as error: + raise click.ClickException(str(error)) from error + retried: Final = _fetch_models(base_url, refreshed.token, target) + if isinstance(retried, PiSyncError): + raise click.ClickException(_listing_error(base_url, retried, target)) + return refreshed, _Listing(retried) + + +def _starting_model(model: str, listing: _Listing) -> str | None: + source: Final = next((listed.id for listed in listing.models if listed.source_model == model), None) + return source or next((listed.id for listed in listing.models if listed.id == model), None) + + +def _model_choice(model: str | None) -> ModelChoice: + return StartOn(model) if model is not None else UnpinModel() + + +def _validated_model(model: str | None, listing: _Listing, base_url: str) -> str | None: + starting: Final = _starting_model(model, listing) if model is not None else None + if model is not None and starting is None: + shown: Final = ", ".join(listing.ids[:_LISTED_MODELS_SHOWN]) + raise click.ClickException(f"{model!r} is not served by {base_url} for this key. /v1/models lists: {shown}.") + return starting + + +def _apply_claude(base_url: str, credential: StaticToken, listing: _Listing, model: str | None) -> None: + listed: Final = listing.ids + starting: Final = _validated_model(model, listing, base_url) + settings_path: Final = claude_settings_path(os.environ) + try: + configure_claude_settings( + base_url, + credential, + _model_choice(starting), + settings_path, + configure_state_path(settings_path), + settings_file_owners(settings_path), + ) + except ClaudeSettingsError as e: + raise click.ClickException(str(e)) + in_picker: Final = sum(1 for listed_model in listed if CLAUDE_CODE_PICKER_PATTERN.search(listed_model)) + click.echo(f"Configured Claude Code: {settings_path} now routes through {base_url}.") + + click.echo("Credential: your virtual key, stored in the file as ANTHROPIC_AUTH_TOKEN.") + click.echo( + f"Starting model: {starting} ({STARTING_MODEL_ROLE}); switch any time with /model." + if starting is not None + else "Starting model: not pinned (Claude Code's default, or a model you set yourself); switch with /model, or " + "pass --model to start on a proxy model. Without a pin, a resumed session re-sends the model its transcript " + "recorded, which behind a raw-model auto-router is the tier model." + ) + click.echo( + f"/model will list all {len(listed)} of the proxy's models." + if in_picker == len(listed) + else f"/model will list {in_picker} of the proxy's {len(listed)} models: Claude Code shows only ids containing " + "'claude' or 'anthropic', and this proxy does not list the rest under such names." + ) + click.echo("Start `claude` from any terminal. Undo with `lite unconfigure claude`.") + if settings_path.is_symlink(): + click.echo( + f"Note: {settings_path} is a symlink to {settings_path.resolve()}, so your key now lives in " + "that file; keep it out of version control.", + err=True, + ) + + +def _has_targets(chosen: Sequence[object]) -> bool: + return bool(chosen) + + +def pick_targets(defaults: tuple[Target, ...] = ("claude", "codex"), *, edit: bool = False) -> tuple[Target, ...]: + choices: Final = [ # mutable-ok: InquirerPy requires a list + Choice(value, name=label, enabled=value in defaults) for value, label in _TARGETS + ] + picked: Final = _TARGET_SELECTION.validate_python( + inquirer.checkbox( + message="Which agents should be edited? Unselected agents keep their current setup" + if edit + else "Which agents should route through LiteLLM?", + choices=choices, + validate=_has_targets, + invalid_message="Pick at least one.", + ).execute() + ) + return tuple(target for target in TARGETS if target in picked) + + +def _pick_model(listed: Sequence[str], default: str | None = None) -> str | None: + choices: Final = [_KEEP_DEFAULT_MODEL, *listed] # mutable-ok: InquirerPy requires a list + picked: Final = _MODEL_SELECTION.validate_python( + inquirer.fuzzy( + message="Model Claude Code starts on (type to filter; /model switches any time):", + choices=choices, + default=default if default in listed else _KEEP_DEFAULT_MODEL, + ).execute() + ) + return None if picked == _KEEP_DEFAULT_MODEL else picked + + +def _pick_codex_model(listed: Sequence[str], default: str | None = None) -> str: + choices: Final = list(listed) # mutable-ok: InquirerPy's choices parameter requires a list + return _MODEL_SELECTION.validate_python( + inquirer.fuzzy( + message="Model Codex starts on (type to filter):", + choices=choices, + default=default if default in listed else listed[0], + ).execute() + ) + + +def _apply_codex(base_url: str, credential: StaticToken, listing: _Listing, model: str) -> None: + _validated_model(model, listing, base_url) + settings_path: Final = codex_config_path(os.environ) + try: + configure_codex_settings(base_url, credential.token, model, settings_path) + except CodexSettingsError as e: + raise click.ClickException(str(e)) from e + click.echo(f"Configured Codex: {settings_path} now routes through {base_url}.") + click.echo(f"Starting model: {model}. Credential: your virtual key, stored in the private provider settings.") + click.echo("Start `codex` from any terminal. Undo with `lite unconfigure codex`.") + if settings_path.is_symlink(): + click.echo(f"Note: your key now lives in {settings_path.resolve()}; keep it out of version control.", err=True) + + +@dataclass(frozen=True, slots=True) +class PreparedSetup: + saved: SavedSetup + listing: _Listing + + +def _connection(ctx: click.Context, saved: SavedSetup | None) -> tuple[str, StaticToken]: + settings: Final = ConnectionSettings.model_validate(ctx.find_object(object)) + base_url: Final = settings.base_url if saved is None or settings.base_url_explicit else saved.base_url + supplied_key: Final = None if settings.api_key_from_token_file else settings.api_key + reusable_key: Final = saved.api_key if saved is not None and saved.base_url == base_url else None + try: + credential: Final = resolve_credential(ctx, supplied_key if supplied_key is not None else reusable_key) + except ClaudeSettingsError as error: + if saved is not None and saved.base_url != base_url and supplied_key is None: + raise click.ClickException( + "The gateway changed. Pass --api-key for the new gateway; the saved key was not used" + ) from error + raise click.ClickException(str(error)) from error + return base_url, credential + + +def _prompt_connection(ctx: click.Context, target: Target, saved: SavedSetup | None) -> tuple[str, StaticToken]: + settings: Final = ConnectionSettings.model_validate(ctx.find_object(object)) + default_url: Final = settings.base_url if saved is None or settings.base_url_explicit else saved.base_url + base_url: Final = normalize_base_url( + click.prompt(f"{target.capitalize()} gateway URL", default=default_url) + ).removesuffix("/v1") + supplied_key: Final = None if settings.api_key_from_token_file else settings.api_key + kept_key: Final = ( + supplied_key + if supplied_key is not None + else (saved.api_key if saved is not None and saved.base_url == base_url else None) + ) + entered: Final = click.prompt( + "Virtual key (press Enter to keep the current key)" if kept_key is not None else "Virtual key", + default="" if kept_key is not None else None, + show_default=False, + hide_input=True, + ) + key: Final[str | None] = entered or kept_key + try: + return base_url, resolve_credential(ctx, key) + except ClaudeSettingsError as error: + raise click.ClickException(str(error)) from error + + +def _prepare( + ctx: click.Context, + target: Target, + saved: SavedSetup | None, + model: str | None, + default_model: bool, + *, + interactive: bool = False, + edit_connection: bool = False, + pick_model: Callable[[Sequence[str]], str | None] | None = None, + pick_codex_model: Callable[[Sequence[str]], str] | None = None, +) -> PreparedSetup: + default: Final = None if default_model else (model if model is not None else (saved.model if saved else None)) + if target == "codex" and default is None and not interactive: + raise click.UsageError("Missing option '--model'. First-time Codex setup needs a starting model") + base_url, credential = _prompt_connection(ctx, target, saved) if edit_connection else _connection(ctx, saved) + repair: Final = saved is not None and not interactive and sys.stdin.isatty() + active_credential, listing = _connection_listing(ctx, base_url, credential, target, repair) + repair_model: Final = repair and default is not None and _starting_model(default, listing) is None + source_names: Final = tuple(item.source_model or item.id for item in listing.models) + chosen: Final = ( + (pick_model(source_names) if pick_model is not None else _pick_model(source_names, default)) + if (interactive or repair_model) and target == "claude" + else ( + pick_codex_model(listing.ids) if pick_codex_model is not None else _pick_codex_model(listing.ids, default) + ) + if interactive or repair_model + else default + ) + if target == "codex" and chosen is None: + raise click.ClickException("First-time Codex setup needs --model. Run `lite configure` for the model picker") + validated: Final = _validated_model(chosen, listing, base_url) + saved_model: Final = ( + next(item.source_model or item.id for item in listing.models if item.id == validated) + if target == "claude" and validated is not None + else chosen + ) + try: + profile: Final = SavedSetup( + target=target, + settings_path=str(settings_path_for(target).resolve()), + base_url=base_url, + api_key=active_credential.token, + model=saved_model, + ) + except ValidationError as error: + raise click.ClickException("Invalid gateway setup; no settings were changed") from error + return PreparedSetup(profile, listing) + + +def _apply(setup: PreparedSetup) -> None: + saved: Final = setup.saved + save_setup(saved) + try: + if saved.target == "claude": + _apply_claude(saved.base_url, StaticToken(saved.api_key), setup.listing, saved.model) + elif saved.model is not None: + _apply_codex(saved.base_url, StaticToken(saved.api_key), setup.listing, saved.model) + except click.ClickException as error: + raise click.ClickException( + f"{error.format_message()} {saved.target.capitalize()} setup was saved. " + f"Run `lite configure {saved.target}` to retry applying it" + ) from error + click.echo( + f"Setup saved. Edit with `lite reconfigure {saved.target}`; " + f"remove saved settings and key with `lite unconfigure {saved.target} --forget`." + ) + + +def configure_targets( + ctx: click.Context, + targets: tuple[Target, ...], + *, + model: str | None = None, + default_model: bool = False, + interactive: bool = False, + edit_connection: bool = False, + pick_model: Callable[[Sequence[str]], str | None] | None = None, + pick_codex_model: Callable[[Sequence[str]], str] | None = None, +) -> None: + if model is not None and default_model: + raise click.UsageError("--model and --default-model cannot be used together") + for target in targets: + _preflight(target) + setups: Final = tuple( + _prepare( + ctx, + target, + read_saved_setup(target), + model, + default_model, + interactive=interactive, + edit_connection=edit_connection, + pick_model=pick_model, + pick_codex_model=pick_codex_model, + ) + for target in targets + ) + for setup in setups: + _apply(setup) + + +def interactive_configure( + ctx: click.Context, + pick_targets: Callable[[], tuple[str, ...]] = pick_targets, + pick_model: Callable[[Sequence[str]], str | None] | None = None, + pick_codex_model: Callable[[Sequence[str]], str] | None = None, +) -> None: + """Configure selected agents, retaining injectable pickers for embedders.""" + selected: Final = pick_targets() + targets: Final[tuple[Target, ...]] = tuple(target for target in TARGETS if target in selected) + if not targets: + return + with setup_locks(targets): + configure_targets(ctx, targets, interactive=True, pick_model=pick_model, pick_codex_model=pick_codex_model) diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index 6d63acc7479..682dd61ac42 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -21,7 +21,7 @@ from .commands.auth import ( from .commands.autoroute.commands import autoroute_group from .commands.chat import chat from .commands.config import config_commands, get_config_value, hidden_command_names -from .commands.configure import configure_group, unconfigure_group +from .commands.configure import configure_group, reconfigure_group, unconfigure_group from .commands.credentials import credentials from .commands.debug import debug from .commands.encryption import encryption @@ -103,7 +103,8 @@ def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: s # If no API key provided via flag or environment variable, try to load from saved token. # Pass base_url so we only use the stored key when it was issued for this server. - api_key_from_token_file: Final = api_key is None and ctx.invoked_subcommand not in ("configure", "unconfigure") + setup_command: Final = ctx.invoked_subcommand in ("configure", "reconfigure", "unconfigure") + api_key_from_token_file: Final = api_key is None and not setup_command resolved_api_key: Final = ( get_stored_api_key(expected_base_url=base_url, vault=context_secret_vault(ctx)) if api_key_from_token_file @@ -119,7 +120,7 @@ def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: s # "user said localhost:4000 on purpose" so they can fall back to # whatever server the stored token was actually issued for. A base_url # saved via `lite config set` counts as the user saying it. - ctx.obj["base_url_explicit"] = base_url_provided or bool(stored_base_url) + ctx.obj["base_url_explicit"] = base_url_provided or (bool(stored_base_url) and not setup_command) if show_version: print_version(base_url, resolved_api_key) @@ -174,6 +175,7 @@ cli.add_command(autoroute_group, name="autoroute") cli.add_command(config_commands) # Add configure/unconfigure (persistently wire a coding agent to the proxy with a virtual key) cli.add_command(configure_group) +cli.add_command(reconfigure_group) cli.add_command(unconfigure_group) diff --git a/tests/test_litellm/proxy/client/cli/test_configure_commands.py b/tests/test_litellm/proxy/client/cli/test_configure_commands.py index 8f68bb1320b..d8be80ef560 100644 --- a/tests/test_litellm/proxy/client/cli/test_configure_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_configure_commands.py @@ -5,6 +5,7 @@ import stat import time from pathlib import Path from types import SimpleNamespace +from typing import Final, Literal import click import pytest @@ -12,6 +13,8 @@ import requests import responses import tomlkit from click.testing import CliRunner +from InquirerPy.base.control import Choice +from pydantic import JsonValue, TypeAdapter from litellm.proxy.client.cli import cli from litellm.proxy.client.cli.commands import claude_settings as claude_settings_module @@ -393,7 +396,7 @@ class TestConfigureAgents: ) assert (settings_path.read_bytes(), codex_path.read_bytes()) == before assert not state_path.exists() - assert not (codex_path.parent / ".litellm").exists() + assert not tuple((codex_path.parent / ".litellm").glob("*.json")) @responses.activate def test_both_configs_are_preflighted_before_fetching_models_or_writing( @@ -433,7 +436,7 @@ class TestConfigureAgents: assert VALID_KEY not in str(caught.value) assert len(responses.calls) == 0 assert not paths[0].exists() and not paths[1].exists() - assert not codex_path.exists() and not (codex_path.parent / ".litellm").exists() + assert not codex_path.exists() and not tuple((codex_path.parent / ".litellm").glob("*.json")) @responses.activate def test_claude_only_configuration_does_not_require_codex( @@ -509,7 +512,7 @@ class TestConfigureAgents: assert not paths[0].exists() and not codex_path.exists() @responses.activate - def test_configure_and_unconfigure_do_not_read_a_stored_login( + def test_configure_reconfigure_and_unconfigure_do_not_read_a_stored_login( self, runner, paths, codex_path, tmp_path, secret_vault_factory, fake_codex_version ): _mock_agent_models() @@ -528,12 +531,17 @@ class TestConfigureAgents: obj={"secret_vault": vault}, ) assert configured.exit_code == 0, configured.output + reconfigured: Final = runner.invoke( + cli, ["reconfigure", "codex", "--model", "auto"], obj={"secret_vault": vault} + ) + assert reconfigured.exit_code == 0, reconfigured.output fake_codex_version(None, 0) undone = runner.invoke(cli, ["unconfigure", "codex"], obj={"secret_vault": vault}) assert undone.exit_code == 0, undone.output assert vault.reads == 0 and vault.writes == [] and vault.erases == 0 assert not codex_path.exists() and not paths[0].exists() - assert "Removed" in undone.output and "sk-login" not in missing.output + configured.output + undone.output + assert "Removed" in undone.output + assert "sk-login" not in missing.output + configured.output + reconfigured.output + undone.output class TestUnconfigureClaude: @@ -599,9 +607,14 @@ class TestUnconfigureClaude: assert str(state_path) in result.output and state_path.exists() assert "sk-ant" not in result.output - def test_refuses_while_lite_up_holds_a_backup(self, runner, paths, lite_up_backup): - result = runner.invoke(cli, ["unconfigure", "claude"]) - assert result.exit_code != 0 and "lite down" in result.output + def test_disconnected_unconfigure_does_not_touch_lite_up_backup( + self, runner: CliRunner, paths: tuple[Path, Path], lite_up_backup: Path + ) -> None: + result: Final = runner.invoke(cli, ["unconfigure", "claude"]) + assert result.exit_code == 0, result.output + assert "nothing to undo" in result.output + assert "Agent settings were not changed" in result.output + assert lite_up_backup.read_text() == "{}" @responses.activate def test_a_config_dir_is_configured_and_undone_apart_from_the_default_file( @@ -625,12 +638,12 @@ class TestUnconfigureClaude: assert undone.exit_code == 0, undone.output assert json.loads((work_dir / "settings.json").read_text()) == original assert not default_settings.exists() and not default_state.exists() - assert runner.invoke(cli, ["unconfigure", "claude"]).exit_code != 0, "the receipt is gone with the undo" + assert runner.invoke(cli, ["unconfigure", "claude"]).exit_code == 0 - def test_without_a_receipt_it_fails_loudly(self, runner, paths): + def test_without_a_receipt_it_reports_nothing_to_undo(self, runner, paths): result = runner.invoke(cli, ["unconfigure", "claude"]) - assert result.exit_code != 0 - assert "nothing to undo" in result.output + assert result.exit_code == 0, result.output + assert "nothing to undo" in result.output.lower() class TestClaudeCodeView: @@ -702,3 +715,534 @@ class TestClaudeCodeView: result = _configure(runner, "--api-key", VALID_KEY) assert result.exit_code == 0, result.output assert "/model will list 1 of the proxy's 2 models: Claude Code shows only ids containing" in result.output + + +def _saved_profile_path(target: Literal["claude", "codex"], settings_path: Path) -> Path: + from litellm.proxy.client.cli.commands.configure_profiles import setup_profile_path + + return setup_profile_path(target, settings_path) + + +def _configure_saved_agent(runner: CliRunner, target: Literal["claude", "codex"]) -> None: + result: Final = runner.invoke( + cli, + ["configure", "--gateway-url", PROXY, "--api-key", VALID_KEY, target, "--model", "auto"], + ) + assert result.exit_code == 0, result.output + + +def _agent_document(settings_path: Path) -> dict[str, JsonValue]: + adapter: Final = TypeAdapter(dict[str, JsonValue]) + if settings_path.suffix == ".json": + return adapter.validate_json(settings_path.read_text()) + return adapter.validate_python(tomlkit.parse(settings_path.read_text()).unwrap()) + + +def _prompt_answer(answer: str | tuple[str, ...]) -> SimpleNamespace: + def execute() -> str | tuple[str, ...]: + return answer + + return SimpleNamespace(execute=execute) + + +class TestSavedAgentSetup: + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("resume", [("configure",), None], ids=["all", "target"]) + def test_disconnect_then_configure_reuses_connection_and_model_without_prompts( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + target: Literal["claude", "codex"], + resume: tuple[str, ...] | None, + ) -> None: + _mock_agent_models() + settings_path: Final = paths[0] if target == "claude" else codex_path + _configure_saved_agent(runner, target) + configured: Final = _agent_document(settings_path) + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + assert not settings_path.exists() + + resumed: Final = runner.invoke(cli, list(resume or ("configure", target))) + assert resumed.exit_code == 0, resumed.output + assert _agent_document(settings_path) == configured + assert "lite configure" in undone.output and "saved" in undone.output.lower() + repeated: Final = runner.invoke(cli, ["configure", target]) + assert repeated.exit_code == 0, repeated.output + assert _agent_document(settings_path) == configured + restored: Final = runner.invoke(cli, ["unconfigure", target]) + assert restored.exit_code == 0, restored.output + assert not settings_path.exists() + assert VALID_KEY not in resumed.output + repeated.output + restored.output + + @responses.activate + def test_resume_both_agents_captures_the_settings_changed_while_disconnected( + self, runner: CliRunner, paths: tuple[Path, Path], codex_path: Path + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, "claude") + codex_url: Final = "https://codex-gateway.test/prefix" + responses.get( + f"{codex_url}/v1/models", + json={"data": [{"id": "auto"}]}, + match=[responses.matchers.header_matcher({"Authorization": "Bearer sk-codex"})], + ) + codex_setup: Final = runner.invoke( + cli, + ["configure", "codex", "--gateway-url", codex_url, "--api-key", "sk-codex", "--model", "auto"], + ) + assert codex_setup.exit_code == 0, codex_setup.output + undone: Final = runner.invoke(cli, ["unconfigure"]) + assert undone.exit_code == 0, undone.output + paths[0].write_text('{"theme": "light", "model": "personal-claude"}') + codex_path.write_text('model = "personal-codex"\napproval_policy = "on-request"\n') + + resumed: Final = runner.invoke(cli, ["configure"]) + assert resumed.exit_code == 0, resumed.output + assert json.loads(paths[0].read_text())["model"] == "claude-router-6175746f" + assert tomlkit.parse(codex_path.read_text())["model"] == "auto" + assert responses.calls[-1].request.url == f"{codex_url}/v1/models" + restored: Final = runner.invoke(cli, ["unconfigure"]) + assert restored.exit_code == 0, restored.output + assert json.loads(paths[0].read_text()) == {"theme": "light", "model": "personal-claude"} + assert tomlkit.parse(codex_path.read_text()) == { + "model": "personal-codex", "approval_policy": "on-request" + } + + @responses.activate + @pytest.mark.parametrize("disconnected", [False, True], ids=["active", "disconnected"]) + @pytest.mark.parametrize( + "forget, forgotten", + [ + (("unconfigure", "--forget", "claude"), ("claude",)), + (("unconfigure", "codex", "--forget"), ("codex",)), + (("unconfigure", "--forget"), ("claude", "codex")), + ], + ids=["group-option-target", "leaf-option", "all"], + ) + def test_forget_removes_only_selected_saved_setups_even_after_disconnect( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + disconnected: bool, + forget: tuple[str, ...], + forgotten: tuple[str, ...], + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, "claude") + _configure_saved_agent(runner, "codex") + if disconnected: + undone: Final = runner.invoke(cli, ["unconfigure"]) + assert undone.exit_code == 0, undone.output + result: Final = runner.invoke(cli, list(forget)) + assert result.exit_code == 0, result.output + for target, settings_path in (("claude", paths[0]), ("codex", codex_path)): + assert _saved_profile_path(target, settings_path).exists() == (target not in forgotten) + resumed: Final = runner.invoke(cli, ["configure", target]) + assert (resumed.exit_code == 0) == (target not in forgotten), resumed.output + assert settings_path.exists() == (target not in forgotten) + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("source", ["leaf", "global", "environment"]) + def test_saved_key_never_follows_a_gateway_override_without_a_replacement( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + monkeypatch: pytest.MonkeyPatch, + target: Literal["claude", "codex"], + source: str, + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, target) + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + replacement_url: Final = "https://replacement.test/gateway" + if source == "environment": + monkeypatch.setenv("LITELLM_PROXY_URL", replacement_url) + args: Final = ( + ["--base-url", replacement_url, "configure", target] + if source == "global" + else ["configure", target, "--gateway-url", replacement_url] + if source == "leaf" + else ["configure", target] + ) + refused: Final = runner.invoke(cli, args) + assert refused.exit_code != 0, refused.output + assert "--api-key" in refused.output and VALID_KEY not in refused.output + assert len(responses.calls) == 1 + assert not paths[0].exists() and not codex_path.exists() + + responses.get( + f"{replacement_url}/v1/models", + json={"data": [{"id": "auto"}]}, + match=[responses.matchers.header_matcher({"Authorization": "Bearer sk-replacement"})], + ) + replaced: Final = runner.invoke(cli, [*args, "--api-key", "sk-replacement"]) + assert replaced.exit_code == 0, replaced.output + assert len(responses.calls) == 2 + assert responses.calls[-1].request.url == f"{replacement_url}/v1/models" + assert VALID_KEY not in replaced.output and "sk-replacement" not in replaced.output + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + def test_saved_setup_is_private_and_scoped_to_the_resolved_agent_home( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + target: Literal["claude", "codex"], + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, target) + settings_path: Final = paths[0] if target == "claude" else codex_path + profile_path: Final = _saved_profile_path(target, settings_path) + assert stat.S_IMODE(profile_path.stat().st_mode) == 0o600 + assert stat.S_IMODE(profile_path.parent.stat().st_mode) & 0o077 == 0 + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + alternate_home: Final = tmp_path / f"other-{target}" + alternate_settings: Final = alternate_home / settings_path.name + environment: Final = "CLAUDE_CONFIG_DIR" if target == "claude" else "CODEX_HOME" + monkeypatch.setenv(environment, str(alternate_home)) + missing: Final = runner.invoke(cli, ["configure", target]) + assert missing.exit_code != 0, missing.output + assert not alternate_settings.exists() and len(responses.calls) == 1 + assert profile_path.exists() + stored_url: Final = runner.invoke(cli, ["config", "set", "base_url", "https://other-default.test"]) + assert stored_url.exit_code == 0, stored_url.output + monkeypatch.setenv(environment, str(settings_path.parent)) + resumed: Final = runner.invoke(cli, ["configure", target]) + assert resumed.exit_code == 0, resumed.output + assert settings_path.exists() and not alternate_settings.exists() + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("fault", ["json", "version", "target", "path"]) + def test_invalid_saved_setup_fails_without_network_or_secret_output_and_can_be_forgotten( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + target: Literal["claude", "codex"], + fault: str, + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, target) + settings_path: Final = paths[0] if target == "claude" else codex_path + profile_path: Final = _saved_profile_path(target, settings_path) + profile: Final = TypeAdapter(dict[str, JsonValue]).validate_json(profile_path.read_text()) + corrupted: Final = ( + "{ " + VALID_KEY + if fault == "json" + else json.dumps({**profile, "version": 999}) + if fault == "version" + else json.dumps({**profile, "target": "codex" if target == "claude" else "claude"}) + if fault == "target" + else json.dumps({**profile, "settings_path": str(settings_path.parent / "another-file")}) + ) + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + profile_path.write_text(corrupted) + failed: Final = runner.invoke(cli, ["configure", target]) + assert failed.exit_code != 0, failed.output + assert "saved" in failed.output.lower() and "--forget" in failed.output + assert VALID_KEY not in failed.output + assert not settings_path.exists() and len(responses.calls) == 1 + forgotten: Final = runner.invoke(cli, ["unconfigure", "--forget", target]) + assert forgotten.exit_code == 0, forgotten.output + assert not profile_path.exists() and not settings_path.exists() + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + def test_reconfigure_prefills_saved_choices_and_changes_only_the_selected_agent( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + monkeypatch: pytest.MonkeyPatch, + target: Literal["claude", "codex"], + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, "claude") + _configure_saved_agent(runner, "codex") + untouched: Final = codex_path if target == "claude" else paths[0] + before: Final = untouched.read_bytes() + responses.replace( + responses.GET, f"{PROXY}/v1/models", json={"data": [{"id": "auto"}, {"id": "replacement"}]} + ) + + def checkbox(**kwargs: object) -> SimpleNamespace: + choices: Final = kwargs["choices"] + assert isinstance(choices, list) and len(choices) == 2 + for choice in choices: + assert isinstance(choice, Choice) and choice.enabled + return _prompt_answer((target,)) + + def fuzzy(**kwargs: object) -> SimpleNamespace: + assert kwargs["default"] == "auto" + return _prompt_answer("replacement") + + monkeypatch.setattr(configure_module.inquirer, "checkbox", checkbox) + monkeypatch.setattr(configure_module.inquirer, "fuzzy", fuzzy) + changed: Final = runner.invoke(cli, ["reconfigure"], input=_TerminalInput(b"\n\n")) + assert changed.exit_code == 0, changed.output + assert PROXY in changed.output and VALID_KEY not in changed.output + assert untouched.read_bytes() == before + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + resumed: Final = runner.invoke(cli, ["configure", target]) + assert resumed.exit_code == 0, resumed.output + if target == "claude": + assert json.loads(paths[0].read_text())["model"] == "replacement" + else: + assert tomlkit.parse(codex_path.read_text())["model"] == "replacement" + assert untouched.read_bytes() == before + + @responses.activate + def test_reconfigure_cancel_preserves_every_agents_settings_and_saved_choices( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, "claude") + _configure_saved_agent(runner, "codex") + files: Final = ( + paths[0], codex_path, _saved_profile_path("claude", paths[0]), _saved_profile_path("codex", codex_path) + ) + before: Final = tuple(path.read_bytes() for path in files) + + def checkbox(**kwargs: object) -> SimpleNamespace: + return _prompt_answer(("claude", "codex")) + + def fuzzy(**kwargs: object) -> SimpleNamespace: + assert tuple(path.read_bytes() for path in files) == before + if "Codex" in str(kwargs["message"]): + raise KeyboardInterrupt() + return _prompt_answer("Keep Claude Code's own default") + + monkeypatch.setattr(configure_module.inquirer, "checkbox", checkbox) + monkeypatch.setattr(configure_module.inquirer, "fuzzy", fuzzy) + cancelled: Final = runner.invoke(cli, ["reconfigure"], input=_TerminalInput(b"\n\n\n\n")) + assert cancelled.exit_code != 0, cancelled.output + assert "Aborted" in cancelled.output + assert tuple(path.read_bytes() for path in files) == before + + @responses.activate + def test_explicit_default_model_unpins_claude_and_remains_the_saved_choice( + self, runner: CliRunner, paths: tuple[Path, Path] + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, "claude") + changed: Final = runner.invoke(cli, ["reconfigure", "claude", "--default-model"]) + assert changed.exit_code == 0, changed.output + assert "model" not in json.loads(paths[0].read_text()) + undone: Final = runner.invoke(cli, ["unconfigure", "claude"]) + assert undone.exit_code == 0, undone.output + resumed: Final = runner.invoke(cli, ["configure", "claude"]) + assert resumed.exit_code == 0, resumed.output + assert "model" not in json.loads(paths[0].read_text()) + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("fault", ["key", "model"]) + def test_terminal_resume_repairs_only_the_rejected_saved_choice( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + monkeypatch: pytest.MonkeyPatch, + target: Literal["claude", "codex"], + fault: str, + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, target) + settings_path: Final = paths[0] if target == "claude" else codex_path + profile_path: Final = _saved_profile_path(target, settings_path) + before: Final = profile_path.read_bytes() + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + responses.reset() + if fault == "key": + responses.get( + f"{PROXY}/v1/models", status=401, + match=[responses.matchers.header_matcher({"Authorization": f"Bearer {VALID_KEY}"})], + ) + responses.get( + f"{PROXY}/v1/models", + json={"data": [{"id": "auto" if fault == "key" else "replacement"}]}, + match=[responses.matchers.header_matcher({ + "Authorization": "Bearer sk-repaired" if fault == "key" else f"Bearer {VALID_KEY}" + })], + ) + failed: Final = runner.invoke(cli, ["configure", target]) + assert failed.exit_code != 0, failed.output + assert not settings_path.exists() and profile_path.read_bytes() == before + assert len(responses.calls) == 1 + + def checkbox(**kwargs: object) -> SimpleNamespace: + raise AssertionError("Saved resume must not ask which agents to configure") + + def fuzzy(**kwargs: object) -> SimpleNamespace: + assert fault == "model", "A rejected key must not discard the saved model" + return _prompt_answer("replacement") + + monkeypatch.setattr(configure_module.inquirer, "checkbox", checkbox) + monkeypatch.setattr(configure_module.inquirer, "fuzzy", fuzzy) + resumed: Final = runner.invoke( + cli, ["configure"], input=_TerminalInput(b"sk-repaired\n" if fault == "key" else b"") + ) + assert resumed.exit_code == 0, resumed.output + assert "gateway URL" not in resumed.output + assert VALID_KEY not in resumed.output and "sk-repaired" not in resumed.output + assert settings_path.exists() + saved: Final = TypeAdapter(dict[str, JsonValue]).validate_json(profile_path.read_text()) + assert saved["api_key"] == ("sk-repaired" if fault == "key" else VALID_KEY) + assert saved["model"] == ("auto" if fault == "key" else "replacement") + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("lost_receipt", [False, True], ids=["malformed-settings", "lost-receipt"]) + def test_forget_without_receipt_preserves_agent_settings_and_reports_unknown_connection( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + target: Literal["claude", "codex"], + lost_receipt: bool, + ) -> None: + from litellm.proxy.client.cli.commands.configure_profiles import receipt_path_for + + _mock_agent_models() + _configure_saved_agent(runner, target) + settings_path: Final = paths[0] if target == "claude" else codex_path + profile_path: Final = _saved_profile_path(target, settings_path) + if lost_receipt: + receipt_path_for(target, settings_path).unlink() + else: + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + settings_path.write_text("[invalid") + before: Final = settings_path.read_bytes() + forgotten: Final = runner.invoke(cli, ["unconfigure", target, "--forget"]) + assert forgotten.exit_code == 0, forgotten.output + assert not profile_path.exists() + assert settings_path.read_bytes() == before + assert "Cannot confirm disconnection" in forgotten.output + assert "gateway connection and key manually" in forgotten.output + assert str(settings_path) in forgotten.output + assert "already disconnected" not in forgotten.output and VALID_KEY not in forgotten.output + assert len(responses.calls) == 1 + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("failure", ["stage_private_json", "commit_staged_json", "apply"]) + def test_failed_setup_write_preserves_saved_intent_and_plain_configure_retries_it( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + monkeypatch: pytest.MonkeyPatch, + target: Literal["claude", "codex"], + failure: str, + ) -> None: + from litellm.proxy.client.cli.commands import configure_profiles, configure_setup + + _mock_agent_models() + _configure_saved_agent(runner, target) + settings_path: Final = paths[0] if target == "claude" else codex_path + profile_path: Final = _saved_profile_path(target, settings_path) + receipt_path: Final = configure_profiles.receipt_path_for(target, settings_path) + before: Final = (settings_path.read_bytes(), profile_path.read_bytes(), receipt_path.read_bytes()) + original_settings: Final = _agent_document(settings_path) + original_profile: Final = TypeAdapter(dict[str, JsonValue]).validate_json(profile_path.read_text()) + replacement_url: Final = "https://replacement.test/gateway" + replacement_key: Final = "sk-replacement" + responses.get( + f"{replacement_url}/v1/models", + json={"data": [{"id": "replacement"}]}, + match=[responses.matchers.header_matcher({"Authorization": f"Bearer {replacement_key}"})], + ) + + def fail_write(*args: object, **kwargs: object) -> str: + raise OSError(f"simulated disk error {VALID_KEY}") + + def fail_apply(*args: object, **kwargs: object) -> None: + error: Final = ( + configure_setup.ClaudeSettingsError if target == "claude" else configure_setup.CodexSettingsError + ) + raise error("simulated agent settings write failure") + + with monkeypatch.context() as patch: + if failure == "apply": + patch.setattr(configure_setup, f"configure_{target}_settings", fail_apply) + else: + patch.setattr(configure_profiles, failure, fail_write) + failed: Final = runner.invoke( + cli, + [ + "reconfigure", target, "--gateway-url", replacement_url, + "--api-key", replacement_key, "--model", "replacement", + ], + ) + assert failed.exit_code != 0, failed.output + assert VALID_KEY not in failed.output and replacement_key not in failed.output + assert (settings_path.read_bytes(), receipt_path.read_bytes()) == (before[0], before[2]) + saved: Final = TypeAdapter(dict[str, JsonValue]).validate_json(profile_path.read_text()) + if failure == "apply": + assert saved == { + **original_profile, "base_url": replacement_url, "api_key": replacement_key, "model": "replacement" + } + assert "simulated agent settings write failure" in failed.output + assert "setup was saved" in failed.output and f"lite configure {target}" in failed.output + else: + assert "could not save" in failed.output.lower() + assert profile_path.read_bytes() == before[1] + retried: Final = runner.invoke(cli, ["configure", target]) + assert retried.exit_code == 0, retried.output + written: Final = _agent_document(settings_path) + if failure != "apply": + assert written == original_settings + elif target == "claude": + environment: Final = written["env"] + assert isinstance(environment, dict) + assert (environment["ANTHROPIC_BASE_URL"], environment["ANTHROPIC_AUTH_TOKEN"], written["model"]) == ( + replacement_url, replacement_key, "replacement" + ) + else: + providers: Final = written["model_providers"] + assert isinstance(providers, dict) + provider: Final = providers["litellm"] + assert isinstance(provider, dict) + headers: Final = provider["http_headers"] + assert isinstance(headers, dict) + assert (provider["base_url"], headers["Authorization"], written["model"]) == ( + f"{replacement_url}/v1", f"Bearer {replacement_key}", "replacement" + ) + + @responses.activate + def test_contended_setup_lock_blocks_requests_and_agent_writes( + self, runner: CliRunner, paths: tuple[Path, Path] + ) -> None: + from litellm.proxy.client.cli.commands.configure_profiles import setup_locks + + _mock_agent_models() + with setup_locks(("claude",)): + blocked: Final = runner.invoke( + cli, + ["configure", "claude", "--gateway-url", PROXY, "--api-key", VALID_KEY, "--model", "auto"], + ) + assert blocked.exit_code != 0, blocked.output + assert "Could not lock agent setup" in blocked.output + assert len(responses.calls) == 0 + assert not paths[0].exists() and not paths[1].exists() + assert not _saved_profile_path("claude", paths[0]).exists() From be35b22dfc37a9f96852dec3262d78908799d164 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Sat, 26 Sep 2026 18:53:12 -0700 Subject: [PATCH 10/65] fix(streaming): keep litellm Usage on text-completion usage chunks (#43047) * fix(streaming): keep litellm Usage on text-completion usage chunks * fix(streaming): convert provider usage to litellm Usage instead of dropping it --- .../litellm_core_utils/streaming_handler.py | 7 +++- .../test_streaming_handler.py | 41 +++++++++++++++++++ 2 files changed, 46 insertions(+), 2 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index fa687b585f5..fa4650aec4f 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1504,8 +1504,11 @@ class CustomStreamWrapper: self.tool_call = True - if hasattr(chunk, "usage") and chunk.usage is not None: - model_response.usage = chunk.usage + chunk_usage: Final = getattr(chunk, "usage", None) + if isinstance(chunk_usage, Usage): + model_response.usage = chunk_usage + elif isinstance(chunk_usage, BaseModel): + model_response.usage = Usage(**chunk_usage.model_dump()) ## RETURN ARG result: Final = self.return_processed_chunk_logic( diff --git a/tests/unit/litellm_core_utils/test_streaming_handler.py b/tests/unit/litellm_core_utils/test_streaming_handler.py index 6557811b530..62d8b0e203f 100644 --- a/tests/unit/litellm_core_utils/test_streaming_handler.py +++ b/tests/unit/litellm_core_utils/test_streaming_handler.py @@ -2859,6 +2859,47 @@ def test_dispatch_text_completion_openai_with_usage( assert model_response.usage.total_tokens == 8 +@pytest.mark.parametrize("custom_llm_provider", ["text-completion-openai", "azure_text"]) +def test_text_completion_usage_chunk_keeps_provider_usage_as_litellm_usage( + initialized_custom_stream_wrapper: CustomStreamWrapper, + custom_llm_provider: str, +): + from openai.types.completion import Completion + from openai.types.completion_usage import CompletionUsage + + initialized_custom_stream_wrapper.custom_llm_provider = custom_llm_provider + initialized_custom_stream_wrapper.model = "gpt-3.5-turbo-instruct" + initialized_custom_stream_wrapper.send_stream_usage = True + initialized_custom_stream_wrapper.received_finish_reason = "length" + provider_usage: Final = CompletionUsage.model_validate( + { + "prompt_tokens": 7, + "completion_tokens": 4, + "total_tokens": 11, + "completion_tokens_details": {"reasoning_tokens": 3}, + "prompt_tokens_details": {"cached_tokens": 2}, + "cost": 0.0123, + } + ) + chunk: Final = Completion.model_construct( + id="cmpl-usage", + choices=[], + created=1, + model="gpt-3.5-turbo-instruct", + object="text_completion", + usage=provider_usage, + ) + + returned: Final = initialized_custom_stream_wrapper.chunk_creator(chunk=chunk) + + assert isinstance(returned.usage, Usage) + dumped: Final = returned.model_dump()["usage"] + assert (dumped["prompt_tokens"], dumped["completion_tokens"], dumped["total_tokens"]) == (7, 4, 11) + assert dumped["cost"] == provider_usage.model_dump()["cost"] + assert dumped["completion_tokens_details"]["reasoning_tokens"] == 3 + assert dumped["prompt_tokens_details"]["cached_tokens"] == 2 + + @pytest.mark.asyncio async def test_custom_stream_wrapper_anext_does_not_block_event_loop_for_sync_iterators( logging_obj: Logging, From 9ba552d527883d7dad9778e63b73422dccbfbaf4 Mon Sep 17 00:00:00 2001 From: agustin18 Date: Sat, 26 Sep 2026 23:51:45 -0300 Subject: [PATCH 11/65] fix(vertex_ai): consider tools when validating context caching min tokens (#43319) * fix(vertex_ai): consider tools when validating context caching min tokens Pass tools to is_prompt_caching_valid_prompt in both sync and async check_and_create_cache before popping them into the cachedContents request body. This allows agent-shaped requests with heavy tool schemas and small message histories to reach the minimum token threshold and benefit from prompt caching. Fixes #42804 * test(vertex_ai): avoid doubles on internal code and assert tools in cache payload --- .../vertex_ai_context_caching.py | 2 + .../test_vertex_ai_context_caching.py | 123 ++++++++++++++++++ 2 files changed, 125 insertions(+) diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index 75d4ffbed86..2d35dd9b480 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -322,6 +322,7 @@ class ContextCachingEndpoints(VertexBase): if not is_prompt_caching_valid_prompt( model=model, messages=cached_messages, + tools=optional_params.get("tools"), custom_llm_provider=custom_llm_provider, ): verbose_logger.debug( @@ -481,6 +482,7 @@ class ContextCachingEndpoints(VertexBase): if not is_prompt_caching_valid_prompt( model=model, messages=cached_messages, + tools=optional_params.get("tools"), custom_llm_provider=custom_llm_provider, ): verbose_logger.debug( diff --git a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 7913700c8a7..283ed3710d0 100644 --- a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1503,6 +1503,129 @@ class TestContextCachingEndpoints: # Restart the patcher so teardown_method can stop it cleanly self._token_check_patcher.start() + @pytest.mark.parametrize("is_async", [False, True]) + @pytest.mark.parametrize( + "custom_llm_provider", ["gemini", "vertex_ai"] + ) + @pytest.mark.asyncio + async def test_check_and_create_cache_considers_tools_for_min_tokens( + self, custom_llm_provider, is_async + ): + """Test that context caching accounts for tools when validating minimum token count. + + Fixes #42804: When messages alone are below the threshold, but tools push the total + over the minimum token count, context caching must proceed and include tools. + """ + self._token_check_patcher.stop() + + short_cached_messages = [ + { + "role": "system", + "content": "Short system instruction.", + "cache_control": {"type": "ephemeral"}, + } + ] + non_cached_messages = [ + {"role": "user", "content": "Hello world"}, + ] + all_messages = short_cached_messages + non_cached_messages + + large_tools = [ + { + "type": "function", + "function": { + "name": f"synthetic_tool_{i}", + "description": "A very descriptive explanation of a synthetic tool designed to add tokens to the prompt cache prefix " * 8, + "parameters": { + "type": "object", + "properties": { + f"arg_{j}": {"type": "string", "description": "Argument description for caching verification " * 4} + for j in range(10) + }, + "required": [f"arg_{j}" for j in range(5)], + }, + }, + } + for i in range(12) + ] + + optional_params = { + **self.sample_optional_params, + "tools": large_tools, + } + + mock_response = MagicMock() + mock_response.json.return_value = { + "name": "cachedContents/test_cache_id", + "model": "gemini-1.5-pro", + } + mock_response.status_code = 200 + self.mock_client.post.return_value = mock_response + self.mock_async_client.post = AsyncMock(return_value=mock_response) + + with patch.object( + self.context_caching, + "_get_token_and_url_context_caching", + return_value=("fake_token", "https://fake.url/cachedContents"), + ), patch.object( + self.context_caching, + "check_cache", + return_value=None, + ), patch.object( + self.context_caching, + "async_check_cache", + new_callable=AsyncMock, + return_value=None, + ): + if is_async: + result = await self.context_caching.async_check_and_create_cache( + messages=all_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-1.5-pro", + client=self.mock_async_client, + timeout=30.0, + logging_obj=self.mock_logging, + cached_content=None, + custom_llm_provider=custom_llm_provider, + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + else: + result = self.context_caching.check_and_create_cache( + messages=all_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-1.5-pro", + client=self.mock_client, + timeout=30.0, + logging_obj=self.mock_logging, + cached_content=None, + custom_llm_provider=custom_llm_provider, + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + + messages, returned_params, returned_cache = result + assert messages == non_cached_messages + assert returned_cache == "cachedContents/test_cache_id" + assert "tools" not in returned_params + + post_mock = self.mock_async_client.post if is_async else self.mock_client.post + post_mock.assert_called_once() + call_kwargs = post_mock.call_args.kwargs + assert call_kwargs["json"]["tools"] == large_tools + assert call_kwargs["json"]["contents"] == [ + {"role": "user", "parts": [{"text": "Short system instruction."}]} + ] + + self._token_check_patcher.start() + + def _model_turn_final_messages(self, final_cached_role): tool_call = { "id": "call_abc123", From ba2d1c2785255f4143c26f35b18c37693e44d4eb Mon Sep 17 00:00:00 2001 From: Jeremy Schoemaker Date: Sat, 26 Sep 2026 22:17:44 -0500 Subject: [PATCH 12/65] =?UTF-8?q?fix(anthropic):=20drop=20thinking=20block?= =?UTF-8?q?s=20with=20empty=20thinking=20text,=20not=20just=20missing=20si?= =?UTF-8?q?gnature=20=F0=9F=A7=A0=F0=9F=9A=AB=20(#38049)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit _is_unsignable_thinking_block() only checked block["signature"], so a thinking block with a valid-looking signature but empty (or whitespace-only) thinking text sailed through _drop_unsignable_thinking_blocks and into anthropic_messages_pt(). Anthropic rejects that with: 400 messages.N.content.M.thinking: each thinking block must contain thinking This is reachable whenever a thinking_blocks history item gets replayed through this Anthropic-shaped request path (e.g. a non-Anthropic reasoning turn with no summary text), the same class of bug PR #36033 fixed on the Responses adapter's own separate code path. Now the signature check runs first (unsigned blocks are still dropped, same as before), then an additional check drops the block if `thinking` is missing, not a string, or strips to empty. redacted_thinking blocks are untouched since they don't have type == "thinking". --- .../prompt_templates/common_utils.py | 18 +- ...llm_core_utils_prompt_templates_factory.py | 176 ++++++++++++++++++ 2 files changed, 188 insertions(+), 6 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 14d47a15c6d..41563d501a9 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -1992,11 +1992,14 @@ def is_encrypted_reasoning_block(block: object) -> bool: def is_unsignable_thinking_block(block: object) -> bool: """A thinking block Anthropic cannot accept on input. - Anthropic verifies the thinking signature cryptographically, so a block whose - signature is null, empty, or missing (e.g. from an open-source reasoning model) - is rejected with a 400 and must be dropped rather than blanked or repaired, and - so is a block whose signature or data carries another provider's encrypted - reasoning. A `redacted_thinking` block Anthropic minted is always kept. + Anthropic verifies the signature cryptographically, so a block with a null, + empty, or missing signature (e.g. from an open-source reasoning model) is + rejected with a 400, and so is a block whose signature or data carries + another provider's encrypted reasoning. It also rejects a `thinking` block + whose text is empty or whitespace-only ("each thinking block must contain + thinking"), regardless of signature, e.g. when a `thinking_blocks` history + item from a non-Anthropic reasoning provider is replayed through this path. + `redacted_thinking` blocks carry no signature and are always kept. """ if is_encrypted_reasoning_block(block): return True @@ -2006,7 +2009,10 @@ def is_unsignable_thinking_block(block: object) -> bool: if mapping.get("type") != "thinking": return False signature: Final = mapping.get("signature") - return not (isinstance(signature, str) and len(signature) > 0) + if not (isinstance(signature, str) and len(signature) > 0): + return True + thinking_text: Final = mapping.get("thinking") + return not (isinstance(thinking_text, str) and len(thinking_text.strip()) > 0) def strip_encrypted_reasoning_from_messages(messages: object) -> None: diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 26124ac24de..0c08c5dfd85 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -3895,3 +3895,179 @@ def test_anthropic_messages_pt_drops_a_system_message_with_no_text(): result = anthropic_messages_pt(messages=messages, model="claude-opus-4-8", llm_provider="anthropic") assert [m["role"] for m in result] == ["user", "assistant"] + + +def test_anthropic_messages_pt_drops_empty_but_signed_thinking_block(): + """ + Anthropic rejects a `thinking` block whose `thinking` text is empty, even + when it carries a valid-looking signature, with: + 400 messages.N.content.M.thinking: each thinking block must contain thinking + This shape is reachable via cross-provider replay of a `thinking_blocks` + history item (see PR #36033), so `is_unsignable_thinking_block()` must + also check the thinking text, not just the signature. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + anthropic_messages_pt, + ) + + messages = [ + {"role": "user", "content": "What's 2+2?"}, + { + "role": "assistant", + "content": "4", + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "", + "signature": "sig_abc123_looks_valid", + } + ], + }, + ] + + result = anthropic_messages_pt( + messages=messages, + model="claude-sonnet-4-5-20250929", + llm_provider="anthropic", + ) + + assistant_msg = result[1] + assert isinstance(assistant_msg["content"], list) + content_types = [block.get("type") for block in assistant_msg["content"]] + assert "thinking" not in content_types, "empty-text thinking block must be dropped even though it has a signature" + + +def test_anthropic_messages_pt_keeps_non_empty_signed_thinking_block(): + """ + Regression: a real, non-empty, signed thinking block must still pass + through unchanged. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + anthropic_messages_pt, + ) + + messages = [ + {"role": "user", "content": "What's 2+2?"}, + { + "role": "assistant", + "content": "4", + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "Let me add these numbers together.", + "signature": "sig_abc123_looks_valid", + } + ], + }, + ] + + result = anthropic_messages_pt( + messages=messages, + model="claude-sonnet-4-5-20250929", + llm_provider="anthropic", + ) + + assistant_msg = result[1] + assert isinstance(assistant_msg["content"], list) + thinking_block = next((b for b in assistant_msg["content"] if b.get("type") == "thinking"), None) + assert thinking_block is not None, "non-empty signed thinking block must be kept" + assert thinking_block["thinking"] == "Let me add these numbers together." + assert thinking_block["signature"] == "sig_abc123_looks_valid" + + +def test_anthropic_messages_pt_keeps_redacted_thinking_block(): + """ + Regression: `redacted_thinking` blocks carry no signature and no plaintext + `thinking` field by design, and must always be kept regardless of the new + emptiness check (which only applies to `type == "thinking"` blocks). + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + anthropic_messages_pt, + ) + + messages = [ + {"role": "user", "content": "What's 2+2?"}, + { + "role": "assistant", + "content": "4", + "thinking_blocks": [ + { + "type": "redacted_thinking", + "data": "encrypted_opaque_blob", + } + ], + }, + ] + + result = anthropic_messages_pt( + messages=messages, + model="claude-sonnet-4-5-20250929", + llm_provider="anthropic", + ) + + assistant_msg = result[1] + assert isinstance(assistant_msg["content"], list) + content_types = [block.get("type") for block in assistant_msg["content"]] + assert "redacted_thinking" in content_types, "redacted_thinking blocks must always be kept" + + +def test_anthropic_messages_pt_drops_unsigned_thinking_block(): + """ + Regression (pre-existing behaviour): a thinking block with no signature + (or an empty/null one) must still be dropped, independent of whether the + thinking text is populated. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + anthropic_messages_pt, + ) + + messages = [ + {"role": "user", "content": "What's 2+2?"}, + { + "role": "assistant", + "content": "4", + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "Let me add these numbers together.", + "signature": "", + } + ], + }, + ] + + result = anthropic_messages_pt( + messages=messages, + model="claude-sonnet-4-5-20250929", + llm_provider="anthropic", + ) + + assistant_msg = result[1] + assert isinstance(assistant_msg["content"], list) + content_types = [block.get("type") for block in assistant_msg["content"]] + assert "thinking" not in content_types, "unsigned thinking block must still be dropped" + + +def test_is_unsignable_thinking_block_treats_whitespace_only_as_empty(): + """ + Edge case: a `thinking` field that is present but whitespace-only (e.g. + a single trailing newline forwarded from another provider's empty + reasoning summary) is functionally empty and Anthropic's API will still + reject it with "each thinking block must contain thinking". We treat it + the same as a fully empty string and drop the block. + + The check lives in the shared `is_unsignable_thinking_block` helper, which + `_drop_unsignable_thinking_blocks` calls standalone, so the whitespace-aware + test has to hold there rather than only at the factory call site. + """ + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + is_unsignable_thinking_block, + ) + + whitespace_only_block = { + "type": "thinking", + "thinking": " \n\t ", + "signature": "sig_abc123_looks_valid", + } + + assert is_unsignable_thinking_block(whitespace_only_block) is True From 303434d5738b5c1b0bbd226792fb40c1c21b6a16 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Sat, 26 Sep 2026 20:26:23 -0700 Subject: [PATCH 13/65] test(e2e): report batch cleanup leftovers as a plain UserWarning (#43405) The leftover warning used a class defined in a test-directory module. The xdist controller cannot import it, so an uncaught leftover warning crashed the whole e2e run. Same change as #43391 on rc/1.103.0 --- tests/e2e/batches/COVERAGE.md | 2 +- tests/e2e/batches/batch_cleanup.py | 8 ++------ tests/e2e/batches/test_batch_cleanup.py | 5 ++--- 3 files changed, 5 insertions(+), 10 deletions(-) diff --git a/tests/e2e/batches/COVERAGE.md b/tests/e2e/batches/COVERAGE.md index cd0fb35165e..2ba49a492ff 100644 --- a/tests/e2e/batches/COVERAGE.md +++ b/tests/e2e/batches/COVERAGE.md @@ -134,7 +134,7 @@ reporting failures as test errors. Already deleted files and batches that are terminal are safe to clean up again. Managed batch cancellation polls for up to two minutes before input deletion. A managed batch still `cancelling` after that is left for the provider to finish, and its input file is left in place because LiteLLM refuses to delete a file a non-terminal -batch references. Both are reported as `BatchCleanupLeftover` warnings naming their ids rather than +batch references. Both are reported as `UserWarning`s naming their ids rather than failing the test. Any other status or error still fails Accepted cancellation may still report validating or in_progress while the provider updates its state. Raw and model-encoded batches are polled until cancelling or diff --git a/tests/e2e/batches/batch_cleanup.py b/tests/e2e/batches/batch_cleanup.py index 5b3baaa624c..f1142a60782 100644 --- a/tests/e2e/batches/batch_cleanup.py +++ b/tests/e2e/batches/batch_cleanup.py @@ -28,10 +28,6 @@ class BatchCleanupClient(Protocol): def cancel_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]: ... -class BatchCleanupLeftover(UserWarning): - pass - - def cleanup_result[R: BaseModel]( action: Callable[[], Result[R]], *, wait: Callable[[float], None] = sleep ) -> Result[R]: @@ -68,7 +64,7 @@ def cleanup_file(client: BatchCleanupClient, file_id: str, *, key: str, provider if isinstance(result, UnknownApiError) and result.status_code == 400 and FILE_IN_USE_REFUSAL in result.body: warnings.warn( f"Left file {file_id} in place: LiteLLM refused to delete it while a batch still references it", - BatchCleanupLeftover, + UserWarning, stacklevel=2, ) return @@ -140,7 +136,7 @@ def cleanup_batch( ) warnings.warn( f"Left batch {batch_id} cancelling after {BATCH_CANCEL_TIMEOUT_SECONDS}s for the provider to finish", - BatchCleanupLeftover, + UserWarning, stacklevel=2, ) return diff --git a/tests/e2e/batches/test_batch_cleanup.py b/tests/e2e/batches/test_batch_cleanup.py index 5e2ac12d300..218e79f37ac 100644 --- a/tests/e2e/batches/test_batch_cleanup.py +++ b/tests/e2e/batches/test_batch_cleanup.py @@ -7,7 +7,6 @@ import pytest from batch_cleanup import ( BATCH_CANCEL_TIMEOUT_SECONDS, CLEANUP_DELAYS, - BatchCleanupLeftover, cleanup_batch, cleanup_file, cleanup_result, @@ -141,7 +140,7 @@ class TestFileCleanup: calls=ExpectedCalls((f"delete None {MANAGED_FILE_ID}",)), files=(UnknownApiError(status_code=400, body=IN_USE_REFUSAL),), ) - with pytest.warns(BatchCleanupLeftover, match=MANAGED_FILE_ID): + with pytest.warns(UserWarning, match=MANAGED_FILE_ID): cleanup_file(client, MANAGED_FILE_ID, key="test-key") client.calls.assert_done() @@ -241,7 +240,7 @@ class TestBatchCancellation: key: Final = manager.key() manager.defer(lambda: cleanup_file(client, MANAGED_FILE_ID, key=key)) manager.defer(lambda: cleanup_batch(client, MANAGED_BATCH_ID, key=key, clock=ticks)) - with pytest.warns(BatchCleanupLeftover) as leftovers: + with pytest.warns(UserWarning, match="^Left ") as leftovers: manager.teardown() client.calls.assert_done() messages: Final = tuple(str(warning.message) for warning in leftovers) From 21055e3fd840375f11a9f90d8b8e1c176284c79b Mon Sep 17 00:00:00 2001 From: Techboy bebop <142545999+kumarpriyanshu09@users.noreply.github.com> Date: Sun, 27 Sep 2026 00:15:12 -0400 Subject: [PATCH 14/65] fix(tools): salvage concatenated JSON tool call arguments (#43260) * fix(tools): salvage concatenated JSON tool call arguments * fix(tools): harden concatenated tool-call salvage for review findings Skip non-dict JSON during split so salvage cannot emit empty tool calls. Collapse srvtoolu_ expansions to the first object so server results stay paired. Allocate __concat_n ids that cannot collide with sibling tool call ids. Propagate cache_control onto every expanded Anthropic tool_use block. Rename the XML invoke loop variable so the key-leak gate no longer flags {args} * test(tools): cover concat id bump and srvtoolu array keep Only collapse srvtoolu_ when concatenated salvage expanded; a valid JSON array argument stays one server tool input * revert(anthropic): drop concat expansion from pass-through adapter Co-authored-by: Techboy bebop * revert(tools): keep concat salvage out of request-side tool converters Co-authored-by: Techboy bebop * fix(tools): expand strictly salvaged concatenated tool arguments in normalized tool calls Co-authored-by: Techboy bebop * fix(tools): retain at most the salvage cap while validating concatenated arguments Co-authored-by: Techboy bebop * test(tools): assert concat sibling ids unique after sanitization A sibling id that only collides after colon-to-underscore sanitization must force the next concat suffix Co-authored-by: Techboy bebop * refactor(tools): drop unused strict mode from split_concatenated_json_objects Strict mode had no production caller. Rejection cases now sit on salvage, and split matches upstream main Co-authored-by: Techboy bebop --------- Co-authored-by: Techboy bebop --- .../prompt_templates/common_utils.py | 55 ++ .../prompt_templates/factory.py | 210 +++++-- ...ore_utils_prompt_templates_common_utils.py | 40 +- ...llm_core_utils_prompt_templates_factory.py | 540 +++++++++--------- 4 files changed, 527 insertions(+), 318 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 41563d501a9..e555d7e8ec0 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -2503,6 +2503,61 @@ def split_concatenated_json_objects(raw: str) -> list[dict[str, object]]: return results +MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS: Final = 8 + + +def salvage_concatenated_tool_arguments(raw: str) -> tuple[dict[str, object], ...]: + """Return complete concatenated JSON objects that are safe to expand. + + Identical objects collapse to the first one and are not capped. More than + ``MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS`` objects that are not all identical + returns an empty tuple. Anything that is not a full concatenation of JSON + objects returns an empty tuple. Repeated copies of the first object are not + retained, and once the cap is passed the rest of the string is only checked. + """ + stripped: Final = raw.strip() + if not stripped: + return () + decoder: Final = json.JSONDecoder() + length: Final = len(stripped) + idx = 0 # rebind-ok: cursor walks the concatenated JSON string + count = 0 # rebind-ok: counts complete objects without retaining duplicates + kept = () # rebind-ok: holds at most one object past the salvage cap + exceeded = False # rebind-ok: cap already passed, the tail is only validated + while idx < length: + while idx < length and stripped[idx] in " \t\n\r": + idx += 1 + if idx >= length: + break + try: + obj, end_idx = decoder.raw_decode(stripped, idx) + except json.JSONDecodeError: + return () + if not isinstance(obj, dict): + return () + idx = end_idx + if exceeded: + continue + count += 1 + if not kept: + kept = (obj,) + continue + if obj == kept[0] and len(kept) == 1: + continue + if len(kept) == 1 and count > 2 and count - 1 > MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS: + exceeded = True + continue + if len(kept) == 1 and count > 2: + kept = (kept[0],) * (count - 1) + if len(kept) >= MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS: + exceeded = True + continue + kept = (*kept, obj) + if exceeded: + return () + return kept + + def text_completion_prompt_to_messages(prompt: object) -> tuple[AllMessageValues, ...]: """ Wrap an OpenAI ``/v1/completions`` ``prompt`` into Chat Completion messages. diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 7b12d1e939f..c4e242fd360 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1,12 +1,14 @@ import base64 import copy import hashlib +import itertools import json import mimetypes import re import xml.etree.ElementTree as ET from collections.abc import Iterator, Mapping, Sequence from enum import Enum +from types import MappingProxyType from typing import Any, Final, TypeAlias, TypedDict, cast, overload from jinja2.sandbox import ImmutableSandboxedEnvironment @@ -52,6 +54,7 @@ from .common_utils import ( is_non_content_values_set, is_unsignable_thinking_block, parse_tool_call_arguments, + salvage_concatenated_tool_arguments, ) from .image_handling import convert_url_to_base64 @@ -5381,80 +5384,167 @@ class NormalizedToolCall(TypedDict): arguments: dict[str, object] -def _parse_tool_call_arguments(raw: object, tool_name: str | None, context: str) -> dict[str, object]: +_ArgumentObjects: TypeAlias = tuple[dict[str, object], ...] +_ParsedToolCall: TypeAlias = tuple[str | None, str | None, _ArgumentObjects] + + +def _optional_call_id(value: object) -> str | None: + if isinstance(value, str) and value: + return value + return None + + +def _optional_tool_name(value: object) -> str | None: + if isinstance(value, str): + return value + return None + + +def _split_tool_call_ids(calls: Sequence[tuple[str | None, int]]) -> tuple[tuple[str | None, ...], ...]: + taken: Final = frozenset(_sanitize_anthropic_tool_use_id(call_id) for call_id, _ in calls if call_id) + + def fresh(call_id: str) -> Iterator[str]: + return filter( + lambda candidate: _sanitize_anthropic_tool_use_id(candidate) not in taken, + (f"{call_id}__concat_{n}" for n in itertools.count(1)), + ) + + suffixes: Final = MappingProxyType( + {_sanitize_anthropic_tool_use_id(call_id): fresh(call_id) for call_id, count in calls if call_id and count > 1} + ) + return tuple( + ( + call_id, + *(next(suffixes[_sanitize_anthropic_tool_use_id(call_id)]) for _ in range(count - 1)), + ) + if call_id + else (None,) * count + for call_id, count in calls + ) + + +def _parse_tool_call_arguments(raw: object, tool_name: str | None, context: str) -> _ArgumentObjects: # Anthropic's tool_use blocks already carry a parsed dict in "input"; # chat completions and the Responses API carry a JSON string that may be # truncated by the model, so route those through the repair-aware parser. if isinstance(raw, dict): - return raw + return (raw,) if not isinstance(raw, str): - return {} + return ({},) normalized_raw: Final = "{}" if raw == REDACTED_BY_LITELLM else raw - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - parse_tool_call_arguments, - ) - try: parsed: Final = parse_tool_call_arguments(normalized_raw, tool_name=tool_name, context=context) except ValueError as e: + salvaged: Final = salvage_concatenated_tool_arguments(normalized_raw) + if salvaged: + verbose_logger.warning( + "Recovered %d tool call(s) from concatenated JSON arguments for tool '%s' (%s)", + len(salvaged), + tool_name or "", + context, + ) + return salvaged verbose_logger.warning("Failed to parse tool call arguments: %s", e) - return {} - return parsed if isinstance(parsed, dict) else {} + return ({},) + return (parsed,) if isinstance(parsed, dict) else ({},) + + +def _choice_tool_calls(choice: object) -> tuple[object, ...]: + message: Final = get_attribute_or_key(choice, "message", None) + tool_calls: Final = get_attribute_or_key(message, "tool_calls", None) if message is not None else None + if isinstance(tool_calls, list): + return tuple(tool_calls) + return () + + +def _selected_choices(response: object, include_all_choices: bool) -> tuple[object, ...]: + choices: Final = get_attribute_or_key(response, "choices", None) + if not isinstance(choices, list) or not choices: + return () + if include_all_choices: + return tuple(choices) + return (choices[0],) + + +def _parsed_chat_tool_call(tool_call: object) -> _ParsedToolCall | None: + function: Final = get_attribute_or_key(tool_call, "function", None) + if function is None: + return None + name: Final = _optional_tool_name(get_attribute_or_key(function, "name")) + return ( + _optional_call_id(get_attribute_or_key(tool_call, "id")), + name, + _parse_tool_call_arguments( + get_attribute_or_key(function, "arguments", "{}"), + tool_name=name, + context="chat completions", + ), + ) + + +def _parsed_calls_in_choice(choice: object) -> tuple[_ParsedToolCall, ...]: + return tuple( + parsed for tool_call in _choice_tool_calls(choice) if (parsed := _parsed_chat_tool_call(tool_call)) is not None + ) + + +def _parsed_chat_tool_calls(response: object, include_all_choices: bool) -> tuple[_ParsedToolCall, ...]: + grouped: Final = tuple( + _parsed_calls_in_choice(choice) for choice in _selected_choices(response, include_all_choices) + ) + return tuple(itertools.chain.from_iterable(grouped)) + + +def _normalized_tool_calls_for_parse( + name: str | None, + call_ids: tuple[str | None, ...], + arguments: _ArgumentObjects, +) -> tuple[NormalizedToolCall, ...]: + return tuple( + NormalizedToolCall(id=call_id, name=name, arguments=argument) + for call_id, argument in zip(call_ids, arguments, strict=True) + ) + + +def _normalized_tool_calls_from_parses(parses: Sequence[_ParsedToolCall]) -> tuple[NormalizedToolCall, ...]: + id_groups: Final = _split_tool_call_ids(tuple((call_id, len(arguments)) for call_id, _, arguments in parses)) + grouped: Final = tuple( + _normalized_tool_calls_for_parse(name, call_ids, arguments) + for (_, name, arguments), call_ids in zip(parses, id_groups, strict=True) + ) + return tuple(itertools.chain.from_iterable(grouped)) def _tool_calls_from_chat_completion_response( response: object, include_all_choices: bool = False -) -> list[NormalizedToolCall]: - choices: Final = get_attribute_or_key(response, "choices", None) - if not (isinstance(choices, list) and choices): - return [] - tool_calls: Final[list[object]] = [] - for choice in choices if include_all_choices else choices[:1]: - message = get_attribute_or_key(choice, "message", None) - choice_tool_calls = get_attribute_or_key(message, "tool_calls", None) if message else None - if isinstance(choice_tool_calls, list): - tool_calls.extend(choice_tool_calls) - result: Final[list[NormalizedToolCall]] = [] - for tc in tool_calls: - fn = get_attribute_or_key(tc, "function", None) - if fn is None: - continue - name = get_attribute_or_key(fn, "name") - result.append( - NormalizedToolCall( - id=get_attribute_or_key(tc, "id"), - name=name, - arguments=_parse_tool_call_arguments( - get_attribute_or_key(fn, "arguments", "{}"), - tool_name=name, - context="chat completions", - ), - ) - ) - return result +) -> tuple[NormalizedToolCall, ...]: + return _normalized_tool_calls_from_parses(_parsed_chat_tool_calls(response, include_all_choices)) -def _tool_calls_from_responses_api_response(response: object) -> list[NormalizedToolCall]: +def _response_function_calls(response: object) -> tuple[object, ...]: output: Final = get_attribute_or_key(response, "output", None) if not isinstance(output, list): - return [] - result: Final[list[NormalizedToolCall]] = [] - for item in output: - if get_attribute_or_key(item, "type") != "function_call": - continue - name = get_attribute_or_key(item, "name") - result.append( - NormalizedToolCall( - id=get_attribute_or_key(item, "call_id") or get_attribute_or_key(item, "id"), - name=name, - arguments=_parse_tool_call_arguments( - get_attribute_or_key(item, "arguments", "{}"), - tool_name=name, - context="responses API", - ), - ) - ) - return result + return () + return tuple(item for item in output if get_attribute_or_key(item, "type") == "function_call") + + +def _parsed_response_tool_call(item: object) -> _ParsedToolCall: + name: Final = _optional_tool_name(get_attribute_or_key(item, "name")) + raw_id: Final = get_attribute_or_key(item, "call_id") or get_attribute_or_key(item, "id") + return ( + _optional_call_id(raw_id), + name, + _parse_tool_call_arguments( + get_attribute_or_key(item, "arguments", "{}"), + tool_name=name, + context="responses API", + ), + ) + + +def _tool_calls_from_responses_api_response(response: object) -> tuple[NormalizedToolCall, ...]: + parses: Final = tuple(_parsed_response_tool_call(item) for item in _response_function_calls(response)) + return _normalized_tool_calls_from_parses(parses) def _tool_calls_from_anthropic_messages_response(response: object) -> list[NormalizedToolCall]: @@ -5494,16 +5584,18 @@ def get_tool_calls_from_response(response: object, include_all_choices: bool = F Callers that only care about a specific tool should filter the result by ``name`` themselves -- this returns every tool call found. """ - chat_tool_calls = _tool_calls_from_chat_completion_response(response, include_all_choices=include_all_choices) + chat_tool_calls: Final = _tool_calls_from_chat_completion_response( + response, include_all_choices=include_all_choices + ) if chat_tool_calls: - return chat_tool_calls + return list(chat_tool_calls) for extractor in ( _tool_calls_from_responses_api_response, _tool_calls_from_anthropic_messages_response, ): tool_calls = extractor(response) if tool_calls: - return tool_calls + return list(tool_calls) return [] diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index 79c50bf2369..45fc93f04c1 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -20,7 +20,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( hoist_images_from_tool_messages, is_encrypted_reasoning_block, merge_consecutive_system_messages, + parse_tool_call_arguments, responses_reasoning_items_from_thinking_blocks, + salvage_concatenated_tool_arguments, split_concatenated_json_objects, strip_encrypted_reasoning_from_messages, system_messages_first, @@ -269,6 +271,40 @@ def test_split_concatenated_json_salvages_prefix_before_truncated_tail(): assert result == [{"a": 1}, {"b": 2}] +def test_parse_tool_call_arguments_rejects_concatenated_json() -> None: + with pytest.raises(ValueError, match="Failed to parse tool call arguments"): + parse_tool_call_arguments('{"a":1}{"b":2}') + + +def _distinct_json_objects(count: int) -> str: + return "".join(json.dumps({"n": index}, separators=(",", ":")) for index in range(count)) + + +@pytest.mark.parametrize( + ("raw", "expected"), + ( + ('{"a":1}{"b":2}', ({"a": 1}, {"b": 2})), + ('{"a":1}{"a":1}{"a":1}', ({"a": 1},)), + ('{"a":1}{"a":1}{"b":2}', ({"a": 1}, {"a": 1}, {"b": 2})), + (_distinct_json_objects(8), tuple({"n": index} for index in range(8))), + (_distinct_json_objects(9), ()), + (_distinct_json_objects(9) + " junk", ()), + ('{"a":1}' * 7 + '{"b":2}', tuple({"a": 1} for _ in range(7)) + ({"b": 2},)), + ('{"a":1}' * 8 + '{"b":2}', ()), + ('{"a":1}' * 5000, ({"a": 1},)), + ('{"a":1}' * 20, ({"a": 1},)), + ('{"a":1}{"b":', ()), + ('0{"x":1}', ()), + ('{"x":1}0', ()), + ('[1]{"x":1}', ()), + ('{"a":1}{"b":2}}', ()), + ('{"a":1} junk', ()), + ), +) +def test_salvage_concatenated_tool_arguments(raw: str, expected: tuple[dict[str, object], ...]) -> None: + assert salvage_concatenated_tool_arguments(raw) == expected + + # --------------------------------------------------------------------------- # Regression tests for non-OpenAI file content blocks. # @@ -1949,6 +1985,8 @@ class TestMergeConsecutiveSystemMessages: assert merged == [{"role": "system", "content": expected_content}, {"role": "user", "content": "Hello"}] def test_keeps_the_first_message_when_no_system_message_in_the_run_has_content(self): - merged = merge_consecutive_system_messages([{"role": "system"}, {"role": "system"}, {"role": "user", "content": "Hi"}]) + merged = merge_consecutive_system_messages( + [{"role": "system"}, {"role": "system"}, {"role": "user", "content": "Hi"}] + ) assert merged == [{"role": "system"}, {"role": "user", "content": "Hi"}] diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 0c08c5dfd85..1e12a973cdb 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -1,4 +1,5 @@ import base64 +import json import logging import os import re @@ -16,10 +17,12 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( _bedrock_tools_pt, _rename_duplicate_bedrock_document_names, _convert_to_bedrock_tool_call_invoke, + _sanitize_anthropic_tool_use_id, _convert_to_bedrock_tool_call_result, anthropic_messages_pt, convert_to_anthropic_tool_result, convert_to_gemini_tool_call_result, + get_tool_calls_from_response, make_valid_bedrock_tool_name, ollama_pt, sanitize_messages_for_tool_calling, @@ -31,9 +34,7 @@ def _get_gemini_function_response_inline_data_parts(result): assert isinstance(result, list), "expected Gemini parts list" assert len(result) == 1, "multimodal function responses should stay in one part" function_response_part = result[0] - assert ( - "inline_data" not in function_response_part - ), "inline_data should be nested under function_response.parts" + assert "inline_data" not in function_response_part, "inline_data should be nested under function_response.parts" function_response = function_response_part["function_response"] nested_parts = function_response["parts"] return [part["inline_data"] for part in nested_parts if "inline_data" in part] @@ -49,7 +50,9 @@ def test_ollama_pt_simple_messages(): result = ollama_pt(model="llama2", messages=messages) - expected_prompt = "### System:\nYou are a helpful assistant\n\n### Assistant:\nHow can I help you?\n\n### User:\nHello\n\n" + expected_prompt = ( + "### System:\nYou are a helpful assistant\n\n### Assistant:\nHow can I help you?\n\n### User:\nHello\n\n" + ) assert isinstance(result, dict) assert result["prompt"] == expected_prompt assert result["images"] == [] @@ -104,10 +107,7 @@ async def test_anthropic_bedrock_thinking_blocks_with_none_content(): # verify the result assert len(result) == 2 - assert ( - result[1]["content"][0]["reasoningContent"]["reasoningText"]["text"] - == "This is a test thinking block" - ) + assert result[1]["content"][0]["reasoningContent"]["reasoningText"]["text"] == "This is a test thinking block" def test_bedrock_converse_assistant_with_empty_thinking_block_and_tool_calls(): @@ -175,11 +175,7 @@ def test_bedrock_converse_assistant_with_empty_thinking_block_and_tool_calls(): assert len(assistant_blocks) == 1 for block in assistant_blocks[0]["content"]: if "text" in block: - assert block[ - "text" - ].strip(), ( - f"Bedrock Converse rejects blank-text ContentBlocks; got {block!r}" - ) + assert block["text"].strip(), f"Bedrock Converse rejects blank-text ContentBlocks; got {block!r}" # toolUse blocks must still be present tool_use_blocks = [b for b in assistant_blocks[0]["content"] if "toolUse" in b] assert len(tool_use_blocks) == 2 @@ -220,19 +216,16 @@ def test_anthropic_messages_pt_drops_unsignable_thinking_block(thinking_block): {"role": "user", "content": "Now what is 3+3?"}, ] - result = anthropic_messages_pt( - messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic") assistant = next(m for m in result if m["role"] == "assistant") content = assistant["content"] - assert all( - block.get("type") not in ("thinking", "redacted_thinking") for block in content - ), f"unsignable thinking block must be dropped, got {content!r}" - assert any( - block.get("type") == "text" and block.get("text") == "2+2 equals 4." - for block in content - ), f"assistant answer text must be preserved, got {content!r}" + assert all(block.get("type") not in ("thinking", "redacted_thinking") for block in content), ( + f"unsignable thinking block must be dropped, got {content!r}" + ) + assert any(block.get("type") == "text" and block.get("text") == "2+2 equals 4." for block in content), ( + f"assistant answer text must be preserved, got {content!r}" + ) def test_anthropic_messages_pt_keeps_signed_thinking_block(): @@ -255,9 +248,7 @@ def test_anthropic_messages_pt_keeps_signed_thinking_block(): {"role": "user", "content": "Now what is 3+3?"}, ] - result = anthropic_messages_pt( - messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic") assistant = next(m for m in result if m["role"] == "assistant") thinking_blocks = [b for b in assistant["content"] if b.get("type") == "thinking"] @@ -373,9 +364,7 @@ def test_bedrock_get_document_format_fallback_mimes(): """ # Test DOCX fallback - docx_mime = ( - "application/vnd.openxmlformats-officedocument.wordprocessingml.document" - ) + docx_mime = "application/vnd.openxmlformats-officedocument.wordprocessingml.document" supported_formats = ["pdf", "docx", "xlsx", "csv"] # Mock mimetypes.guess_all_extensions to return empty list (simulating Docker container scenario) @@ -399,15 +388,11 @@ def test_bedrock_get_document_format_mimetypes_success(): """ Test the _get_document_format method when mimetypes.guess_all_extensions works normally. """ - docx_mime = ( - "application/vnd.openxmlformats-officedocument.wordprocessingml.document" - ) + docx_mime = "application/vnd.openxmlformats-officedocument.wordprocessingml.document" supported_formats = ["pdf", "docx", "xlsx", "csv"] # Test normal mimetypes behavior (should not hit fallback) - result = BedrockImageProcessor._get_document_format( - mime_type=docx_mime, supported_doc_formats=supported_formats - ) + result = BedrockImageProcessor._get_document_format(mime_type=docx_mime, supported_doc_formats=supported_formats) assert result == "docx", f"Expected 'docx', got '{result}'" @@ -623,9 +608,7 @@ async def test_bedrock_process_image_async_factory(): image_url = "data:application/pdf; qs=0.001;base64,JVBERi0xLjQKJcOkw7zDtsOfCjIgMCBvYmoKPDwvTGVuZ3RoIDMgMCBSL0ZpbHRlci9GbGF0ZURlY29kZT4" - content_block = await BedrockImageProcessor.process_image_async( - image_url=image_url, format=None - ) + content_block = await BedrockImageProcessor.process_image_async(image_url=image_url, format=None) print(f"content_block: {content_block}") @@ -668,9 +651,7 @@ def test_unpack_defs_resolves_nested_ref_inside_anyof_items(): items_schema = schema["properties"]["vatAmounts"]["anyOf"][0]["items"] # Assertions: items_schema should now be the resolved object, not an empty dict - assert isinstance( - items_schema, dict - ), "Items schema should be a dict after unpacking" + assert isinstance(items_schema, dict), "Items schema should be a dict after unpacking" assert items_schema.get("type") == "object" # Ensure essential properties are present assert set(items_schema.get("properties", {}).keys()) == {"vatRate", "vatAmount"} @@ -861,9 +842,7 @@ def test_convert_gemini_tool_call_result_with_multiple_anthropic_image_blocks(): last_message_with_tool_calls=last_message_with_tool_calls, ) inline_parts = _get_gemini_function_response_inline_data_parts(result) - assert ( - len(inline_parts) == 2 - ), f"expected 2 inline_data parts, got {len(inline_parts)}" + assert len(inline_parts) == 2, f"expected 2 inline_data parts, got {len(inline_parts)}" mime_types = {p["mime_type"] for p in inline_parts} assert mime_types == {"image/png", "image/jpeg"} @@ -899,9 +878,7 @@ def test_convert_gemini_tool_call_result_with_data_url_string(): last_message_with_tool_calls=last_message_with_tool_calls, ) inline_parts = _get_gemini_function_response_inline_data_parts(result) - assert ( - len(inline_parts) == 1 - ), "data-URL image string was not converted to inline_data" + assert len(inline_parts) == 1, "data-URL image string was not converted to inline_data" assert inline_parts[0]["mime_type"] == "image/png" assert inline_parts[0]["data"] == tiny_png_b64 @@ -937,9 +914,9 @@ def test_convert_gemini_tool_call_result_with_data_url_extra_params(): ) inline_parts = _get_gemini_function_response_inline_data_parts(result) assert len(inline_parts) == 1 - assert ( - inline_parts[0]["mime_type"] == "image/png" - ), f"expected clean 'image/png', got '{inline_parts[0]['mime_type']}'" + assert inline_parts[0]["mime_type"] == "image/png", ( + f"expected clean 'image/png', got '{inline_parts[0]['mime_type']}'" + ) def test_bedrock_tools_unpack_defs(): @@ -1036,9 +1013,7 @@ def test_bedrock_tools_pt_strict_parameter(): }, } ] - result = _bedrock_tools_pt( - tools_with_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools_with_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0") assert result[0]["toolSpec"]["strict"] is True assert result[0]["toolSpec"]["inputSchema"]["json"]["additionalProperties"] is False @@ -1060,9 +1035,7 @@ def test_bedrock_tools_pt_strict_parameter(): }, } ] - result = _bedrock_tools_pt( - tools_without_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools_without_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0") assert "strict" not in result[0]["toolSpec"] assert "additionalProperties" not in result[0]["toolSpec"]["inputSchema"]["json"] @@ -1085,9 +1058,7 @@ def test_bedrock_image_processor_content_type_fallback_url_extension(): # Test with .png URL image_url = "https://example.com/test-image.png" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/png" assert base64_bytes == base64.b64encode(png_content).decode("utf-8") @@ -1111,9 +1082,7 @@ def test_bedrock_image_processor_content_type_fallback_binary_detection(): # Test with URL without extension image_url = "https://example.com/test-image-without-extension" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/jpeg" assert base64_bytes == base64.b64encode(jpeg_content).decode("utf-8") @@ -1136,9 +1105,7 @@ def test_bedrock_image_processor_content_type_fallback_application_octet_stream( # Test with .gif URL image_url = "https://s3.amazonaws.com/bucket/image.gif" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/gif" assert base64_bytes == base64.b64encode(gif_content).decode("utf-8") @@ -1161,9 +1128,7 @@ def test_bedrock_image_processor_content_type_with_query_params(): # Test with URL containing query parameters (common in S3 signed URLs) image_url = "https://s3.amazonaws.com/bucket/image.webp?AWSAccessKeyId=123&Expires=456&Signature=789" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/webp" assert base64_bytes == base64.b64encode(webp_content).decode("utf-8") @@ -1185,9 +1150,7 @@ def test_bedrock_image_processor_content_type_normal_header(): mock_response.content = png_content image_url = "https://example.com/test-image.png" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/png" assert base64_bytes == base64.b64encode(png_content).decode("utf-8") @@ -1207,7 +1170,7 @@ def test_bedrock_image_processor_content_type_fallback_failure(): # Test with URL without recognizable extension image_url = "https://example.com/unknown-file" - with pytest.raises(ValueError, match='Unable to determine content type from URL: https') as excinfo: + with pytest.raises(ValueError, match="Unable to determine content type from URL: https") as excinfo: BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert "Unable to determine content type" in str(excinfo.value) @@ -1227,16 +1190,12 @@ def test_bedrock_image_processor_content_type_jpeg_variants(): # Test with .jpg extension image_url_jpg = "https://example.com/photo.jpg" - _, content_type_jpg = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url_jpg - ) + _, content_type_jpg = BedrockImageProcessor._post_call_image_processing(mock_response, image_url_jpg) assert content_type_jpg == "image/jpeg" # Test with .jpeg extension image_url_jpeg = "https://example.com/photo.jpeg" - _, content_type_jpeg = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url_jpeg - ) + _, content_type_jpeg = BedrockImageProcessor._post_call_image_processing(mock_response, image_url_jpeg) assert content_type_jpeg == "image/jpeg" @@ -1258,9 +1217,7 @@ def test_bedrock_image_processor_content_type_pdf_document(): # Test with .pdf URL pdf_url = "https://s3.amazonaws.com/bucket/document.pdf" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, pdf_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, pdf_url) assert content_type == "application/pdf" assert base64_bytes == base64.b64encode(pdf_content).decode("utf-8") @@ -1293,12 +1250,8 @@ def test_bedrock_image_processor_content_type_document_formats(): ] for url, expected_mime in test_cases: - _, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, url - ) - assert ( - content_type == expected_mime - ), f"Expected {expected_mime} for {url}, got {content_type}" + _, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, url) + assert content_type == expected_mime, f"Expected {expected_mime} for {url}, got {content_type}" def test_bedrock_image_processor_content_type_s3_pdf_with_query(): @@ -1317,9 +1270,7 @@ def test_bedrock_image_processor_content_type_s3_pdf_with_query(): # S3 signed URL with query parameters s3_url = "https://my-bucket.s3.us-east-1.amazonaws.com/documents/report.pdf?AWSAccessKeyId=AKIAIOSFODNN7EXAMPLE&Expires=1234567890&Signature=abcdef123456" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, s3_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, s3_url) assert content_type == "application/pdf" assert base64_bytes == base64.b64encode(pdf_content).decode("utf-8") @@ -1428,12 +1379,8 @@ def test_bedrock_create_bedrock_block_normalized_base64(): base64_content = base64.b64encode(pdf_content).decode("utf-8") # Create versions with different whitespace - base64_with_newlines = "\n".join( - [base64_content[i : i + 64] for i in range(0, len(base64_content), 64)] - ) - base64_with_spaces = " ".join( - [base64_content[i : i + 32] for i in range(0, len(base64_content), 32)] - ) + base64_with_newlines = "\n".join([base64_content[i : i + 64] for i in range(0, len(base64_content), 64)]) + base64_with_spaces = " ".join([base64_content[i : i + 32] for i in range(0, len(base64_content), 32)]) # Create blocks block1 = BedrockImageProcessor._create_bedrock_block( @@ -1565,9 +1512,7 @@ def test_bedrock_create_bedrock_block_document_name_format(): # Check format: DocumentPDFmessages_{16_hex_chars}_{format} pattern = r"^DocumentPDFmessages_[0-9a-f]{16}_pdf$" - assert re.match( - pattern, document_name - ), f"Document name format mismatch: {document_name}" + assert re.match(pattern, document_name), f"Document name format mismatch: {document_name}" def test_bedrock_create_bedrock_block_different_document_formats(): @@ -1620,9 +1565,7 @@ def test_bedrock_nova_web_search_options_mapping(): assert system_tool["name"] == "nova_grounding" # Test with search_context_size (should be ignored for Nova) - result2 = config._map_web_search_options( - {"search_context_size": "high"}, "us.amazon.nova-premier-v1:0" - ) + result2 = config._map_web_search_options({"search_context_size": "high"}, "us.amazon.nova-premier-v1:0") assert result2 is not None system_tool2 = result2.get("systemTool") @@ -1688,9 +1631,7 @@ def test_bedrock_tools_pt_drops_unmappable_responses_builtin_tools(): {"type": "custom", "name": "free_form"}, ] - result = _bedrock_tools_pt( - tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0") names = [block["toolSpec"]["name"] for block in result if "toolSpec" in block] assert names == ["noop"] @@ -1720,9 +1661,7 @@ def test_bedrock_tools_pt_keeps_anthropic_input_schema_tools(): }, ] - result = _bedrock_tools_pt( - tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0") names = [block["toolSpec"]["name"] for block in result if "toolSpec" in block] assert names == ["lookup"] @@ -1924,9 +1863,7 @@ def test_anthropic_messages_pt_server_tool_use_passthrough(): "tool_use_id": "srvtoolu_01ABC123", "content": { "type": "tool_search_tool_search_result", - "tool_references": [ - {"type": "tool_reference", "tool_name": "get_time"} - ], + "tool_references": [{"type": "tool_reference", "tool_name": "get_time"}], }, }, {"type": "text", "text": "I found the time tool. How can I help you?"}, @@ -1954,20 +1891,14 @@ def test_anthropic_messages_pt_server_tool_use_passthrough(): # Verify server_tool_use block is preserved assert "server_tool_use" in content_types - server_tool_use_block = next( - b for b in assistant_msg["content"] if b.get("type") == "server_tool_use" - ) + server_tool_use_block = next(b for b in assistant_msg["content"] if b.get("type") == "server_tool_use") assert server_tool_use_block["id"] == "srvtoolu_01ABC123" assert server_tool_use_block["name"] == "tool_search_tool_regex" assert server_tool_use_block["input"] == {"query": ".*time.*"} # Verify tool_search_tool_result block is preserved assert "tool_search_tool_result" in content_types - tool_result_block = next( - b - for b in assistant_msg["content"] - if b.get("type") == "tool_search_tool_result" - ) + tool_result_block = next(b for b in assistant_msg["content"] if b.get("type") == "tool_search_tool_result") assert tool_result_block["tool_use_id"] == "srvtoolu_01ABC123" assert tool_result_block["content"]["type"] == "tool_search_tool_search_result" assert tool_result_block["content"]["tool_references"][0]["tool_name"] == "get_time" @@ -2019,9 +1950,7 @@ def test_bedrock_tools_unpack_defs_no_oom_with_nested_refs(): "anyOf": [ {"$ref": "#/$defs/Literal"}, {"$ref": "#/$defs/FieldRef"}, - { - "$ref": "#/$defs/Expression" - }, # Circular: Operand -> Expression -> Operand + {"$ref": "#/$defs/Expression"}, # Circular: Operand -> Expression -> Operand ], }, "Literal": { @@ -2155,9 +2084,7 @@ def test_anthropic_messages_pt_file_block_cache_control_with_explicit_provider() file_block = content_blocks[0] assert file_block["type"] == "document" - assert ( - "cache_control" in file_block - ), "cache_control should be preserved on file/document content blocks" + assert "cache_control" in file_block, "cache_control should be preserved on file/document content blocks" assert file_block["cache_control"]["type"] == "ephemeral" text_block = content_blocks[1] @@ -2365,22 +2292,16 @@ def test_bedrock_tool_call_invoke_concatenated_json(): # First block keeps original tool id assert result[0]["toolUse"]["toolUseId"] == "tooluse_L7I3TewYAUhoheJZQEuwVN" assert result[0]["toolUse"]["name"] == "shell" - assert result[0]["toolUse"]["input"] == { - "command": ["curl", "-i", "http://localhost:9009", "-m", "10"] - } + assert result[0]["toolUse"]["input"] == {"command": ["curl", "-i", "http://localhost:9009", "-m", "10"]} # Subsequent blocks get suffixed ids assert result[1]["toolUse"]["toolUseId"] == "tooluse_L7I3TewYAUhoheJZQEuwVN_1" assert result[1]["toolUse"]["name"] == "shell" - assert result[1]["toolUse"]["input"] == { - "command": ["curl", "-i", "http://localhost:9009/robots.txt", "-m", "5"] - } + assert result[1]["toolUse"]["input"] == {"command": ["curl", "-i", "http://localhost:9009/robots.txt", "-m", "5"]} assert result[2]["toolUse"]["toolUseId"] == "tooluse_L7I3TewYAUhoheJZQEuwVN_2" assert result[2]["toolUse"]["name"] == "shell" - assert result[2]["toolUse"]["input"] == { - "command": ["curl", "-i", "http://localhost:9009/sitemap.xml", "-m", "5"] - } + assert result[2]["toolUse"]["input"] == {"command": ["curl", "-i", "http://localhost:9009/sitemap.xml", "-m", "5"]} def test_bedrock_tool_call_invoke_concatenated_json_with_cache_control(): @@ -2535,9 +2456,7 @@ def test_bedrock_tool_call_invoke_unconvertible_raises_non_retryable_bad_request def test_make_valid_bedrock_tool_name_preserves_hyphens(): assert make_valid_bedrock_tool_name("my-tool") == "my-tool" assert ( - make_valid_bedrock_tool_name( - "CreateCaseKnowledgeArticle_foTWsqR6yDt-OnSsvR5e6Q" - ) + make_valid_bedrock_tool_name("CreateCaseKnowledgeArticle_foTWsqR6yDt-OnSsvR5e6Q") == "CreateCaseKnowledgeArticle_foTWsqR6yDt-OnSsvR5e6Q" ) @@ -2564,9 +2483,7 @@ def test_bedrock_tool_name_sanitized_consistently_in_tools_and_tool_use(): "function": {"name": raw_name, "arguments": "{}"}, } ] - tool_use_name = _convert_to_bedrock_tool_call_invoke(tool_calls)[0]["toolUse"][ - "name" - ] + tool_use_name = _convert_to_bedrock_tool_call_invoke(tool_calls)[0]["toolUse"]["name"] assert tool_spec_name == "foo_bar" assert tool_use_name == tool_spec_name @@ -2589,15 +2506,8 @@ def test_bedrock_converse_messages_pt_tool_use_matches_tool_spec_hyphen_name(): ], }, ] - translated = _bedrock_converse_messages_pt( - messages=messages, model="", llm_provider="" - ) - tool_use_blocks = [ - block - for msg in translated - for block in msg.get("content", []) - if "toolUse" in block - ] + translated = _bedrock_converse_messages_pt(messages=messages, model="", llm_provider="") + tool_use_blocks = [block for msg in translated for block in msg.get("content", []) if "toolUse" in block] assert len(tool_use_blocks) == 1 assert tool_use_blocks[0]["toolUse"]["name"] == tool_name @@ -2694,11 +2604,7 @@ def test_sanitize_messages_deduplicates_tool_results(): result = sanitize_messages_for_tool_calling(messages) # Count tool messages with this ID — should be exactly 1 - tool_results = [ - m - for m in result - if m.get("role") == "tool" and m.get("tool_call_id") == "call_abc123" - ] + tool_results = [m for m in result if m.get("role") == "tool" and m.get("tool_call_id") == "call_abc123"] assert len(tool_results) == 1 # Should keep the LAST occurrence (most complete) assert tool_results[0]["content"] == '{"temperature": 72, "condition": "sunny"}' @@ -2833,11 +2739,7 @@ def test_sanitize_messages_dedup_scoped_per_turn_preserves_cross_turn(): result = sanitize_messages_for_tool_calling(messages) # Both tool results must survive — one per turn - tool_results = [ - m - for m in result - if m.get("role") == "tool" and m.get("tool_call_id") == "call_X" - ] + tool_results = [m for m in result if m.get("role") == "tool" and m.get("tool_call_id") == "call_X"] assert len(tool_results) == 2, ( f"Expected 2 tool results (one per turn), got {len(tool_results)}. " "Dedup may be global instead of per-turn scoped." @@ -2891,32 +2793,26 @@ def test_sanitize_messages_combined_case_a_and_case_d(): tool_results = [m for m in result if m.get("role") in ("tool", "function")] # Case A: call_missing should have a dummy result injected - missing_results = [ - m for m in tool_results if m.get("tool_call_id") == "call_missing" - ] - assert ( - len(missing_results) == 1 - ), f"Expected 1 dummy result for call_missing (Case A), got {len(missing_results)}" + missing_results = [m for m in tool_results if m.get("tool_call_id") == "call_missing"] + assert len(missing_results) == 1, ( + f"Expected 1 dummy result for call_missing (Case A), got {len(missing_results)}" + ) # Case D: call_duped should have exactly 1 result (the fresh one) - duped_results = [ - m for m in tool_results if m.get("tool_call_id") == "call_duped" - ] - assert ( - len(duped_results) == 1 - ), f"Expected 1 result for call_duped after dedup (Case D), got {len(duped_results)}" - assert ( - duped_results[0]["content"] == "fresh_result" - ), f"Expected last-wins 'fresh_result', got '{duped_results[0]['content']}'" + duped_results = [m for m in tool_results if m.get("tool_call_id") == "call_duped"] + assert len(duped_results) == 1, ( + f"Expected 1 result for call_duped after dedup (Case D), got {len(duped_results)}" + ) + assert duped_results[0]["content"] == "fresh_result", ( + f"Expected last-wins 'fresh_result', got '{duped_results[0]['content']}'" + ) # Verify tool results immediately follow the assistant message asst_idx = next(i for i, m in enumerate(result) if m.get("role") == "assistant") - tool_msgs_after_asst = [ - m for m in result[asst_idx + 1 :] if m.get("role") in ("tool", "function") - ] - assert ( - len(tool_msgs_after_asst) == 2 - ), f"Expected 2 tool results after assistant, got {len(tool_msgs_after_asst)}" + tool_msgs_after_asst = [m for m in result[asst_idx + 1 :] if m.get("role") in ("tool", "function")] + assert len(tool_msgs_after_asst) == 2, ( + f"Expected 2 tool results after assistant, got {len(tool_msgs_after_asst)}" + ) # Both tool_call_ids should be present (order may vary) tool_ids = {m["tool_call_id"] for m in tool_msgs_after_asst} assert tool_ids == { @@ -2958,9 +2854,7 @@ def test_anthropic_messages_pt_file_block_preserves_cache_control(): } ] - result = anthropic_messages_pt( - messages, model="claude-sonnet-4-20250514", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages, model="claude-sonnet-4-20250514", llm_provider="anthropic") content_blocks = result[0]["content"] assert len(content_blocks) == 2 @@ -2968,9 +2862,7 @@ def test_anthropic_messages_pt_file_block_preserves_cache_control(): # Document block (from file) should preserve cache_control doc_block = content_blocks[0] assert doc_block["type"] == "document" - assert ( - "cache_control" in doc_block - ), "cache_control was dropped from file/document block" + assert "cache_control" in doc_block, "cache_control was dropped from file/document block" assert doc_block["cache_control"]["type"] == "ephemeral" # Text block should also preserve cache_control @@ -3013,9 +2905,7 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch): } # Claude 4.5 model: ttl should be preserved - result = add_cache_point_tool_block( - tool_with_1h, model="jp.anthropic.claude-opus-4-7" - ) + result = add_cache_point_tool_block(tool_with_1h, model="jp.anthropic.claude-opus-4-7") assert result is not None assert result["cachePoint"]["type"] == "default" assert result["cachePoint"]["ttl"] == "1h" @@ -3024,16 +2914,12 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch): tool_with_5m = { "cache_control": {"type": "ephemeral", "ttl": "5m"}, } - result_5m = add_cache_point_tool_block( - tool_with_5m, model="jp.anthropic.claude-opus-4-7" - ) + result_5m = add_cache_point_tool_block(tool_with_5m, model="jp.anthropic.claude-opus-4-7") assert result_5m is not None assert result_5m["cachePoint"]["ttl"] == "5m" # Older model: ttl should be stripped - result_old = add_cache_point_tool_block( - tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0" - ) + result_old = add_cache_point_tool_block(tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0") assert result_old is not None assert result_old["cachePoint"]["type"] == "default" assert "ttl" not in result_old["cachePoint"] @@ -3052,9 +2938,7 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch): # cache_control without ttl: returns default cachePoint (unchanged behavior) tool_no_ttl = {"cache_control": {"type": "ephemeral"}} - result_no_ttl = add_cache_point_tool_block( - tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result_no_ttl = add_cache_point_tool_block(tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0") assert result_no_ttl is not None assert result_no_ttl["cachePoint"]["type"] == "default" assert "ttl" not in result_no_ttl["cachePoint"] @@ -3127,9 +3011,7 @@ def test_bedrock_tools_pt_passes_ttl_for_claude_4_5(monkeypatch): assert cache_blocks[0]["cachePoint"]["ttl"] == "1h" # Older model: cachePoint should not have ttl - result_old = _bedrock_tools_pt( - tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0" - ) + result_old = _bedrock_tools_pt(tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0") cache_blocks_old = [b for b in result_old if "cachePoint" in b] assert len(cache_blocks_old) == 1 assert "ttl" not in cache_blocks_old[0]["cachePoint"] @@ -3204,9 +3086,7 @@ def test_bedrock_converse_messages_pt_document_various_formats(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") doc_block = result[0]["content"][0] assert doc_block["document"]["format"] == expected_format, ( @@ -3233,12 +3113,8 @@ def test_bedrock_converse_messages_pt_document_deterministic_name(): } ] - result1 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) - result2 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + result1 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") + result2 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") name1 = result1[0]["content"][0]["document"]["name"] name2 = result2[0]["content"][0]["document"]["name"] @@ -3272,34 +3148,18 @@ def test_bedrock_converse_messages_pt_renames_duplicate_document_names(): }, ] - result1 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) - result2 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + result1 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") + result2 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") - names1 = [ - block["document"]["name"] - for message in result1 - for block in message["content"] - if "document" in block - ] - names2 = [ - block["document"]["name"] - for message in result2 - for block in message["content"] - if "document" in block - ] + names1 = [block["document"]["name"] for message in result1 for block in message["content"] if "document" in block] + names2 = [block["document"]["name"] for message in result2 for block in message["content"] if "document" in block] assert len(names1) == 2 assert len(set(names1)) == 2 assert names1[1] == f"{names1[0]}_2" assert names1 == names2 - single_turn = _bedrock_converse_messages_pt( - [messages[0]], "anthropic.claude-sonnet-4-6", "bedrock" - ) + single_turn = _bedrock_converse_messages_pt([messages[0]], "anthropic.claude-sonnet-4-6", "bedrock") assert names1[0] == single_turn[0]["content"][0]["document"]["name"] @@ -3321,14 +3181,10 @@ def test_rename_duplicate_bedrock_document_names_skips_organic_suffixes(): def _names(contents): return [block["document"]["name"] for block in contents[0]["content"]] - organic_first = _rename_duplicate_bedrock_document_names( - _contents(["report", "report_2", "report"]) - ) + organic_first = _rename_duplicate_bedrock_document_names(_contents(["report", "report_2", "report"])) assert _names(organic_first) == ["report", "report_2", "report_3"] - organic_last = _rename_duplicate_bedrock_document_names( - _contents(["report", "report", "report_2"]) - ) + organic_last = _rename_duplicate_bedrock_document_names(_contents(["report", "report", "report_2"])) assert _names(organic_last) == ["report", "report_3", "report_2"] @@ -3350,18 +3206,11 @@ def test_bedrock_converse_messages_pt_document_rejects_url_source(): ] with pytest.raises(ValueError, match="only supports base64-encoded"): - _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") def _collect_cache_points(blocks): - return [ - block["cachePoint"] - for message in blocks - for block in message["content"] - if "cachePoint" in block - ] + return [block["cachePoint"] for message in blocks for block in message["content"] if "cachePoint" in block] @pytest.mark.parametrize( @@ -3527,6 +3376,189 @@ def test_get_tool_calls_from_response_warns_for_malformed_arguments(caplog): assert "Failed to parse tool call arguments" in caplog.text +def _concatenated_json(*payloads: dict[str, object]) -> str: + return "".join(json.dumps(payload, separators=(",", ":")) for payload in payloads) + + +def _function_tool_call(call_id: str | None, name: str, arguments: str) -> dict[str, object]: + return {"id": call_id, "function": {"name": name, "arguments": arguments}} + + +def _chat_tool_response(*tool_calls: dict[str, object]) -> dict[str, object]: + return {"choices": [{"message": {"tool_calls": list(tool_calls)}}]} + + +def test_get_tool_calls_from_response_expands_distinct_concatenated_arguments(caplog): + raw = '{"flag":true}{"box":"A","limit":50}' + response: Final = _chat_tool_response(_function_tool_call("call_move", "move", raw)) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + tool_calls: Final = get_tool_calls_from_response(response) + + assert tool_calls == [ + {"id": "call_move", "name": "move", "arguments": {"flag": True}}, + {"id": "call_move__concat_1", "name": "move", "arguments": {"box": "A", "limit": 50}}, + ] + assert "Recovered 2 tool call(s)" in caplog.text + assert "move" in caplog.text + assert "flag" not in caplog.text + + +def test_get_tool_calls_from_response_expands_responses_api_concatenated_arguments(): + response: Final = { + "output": [ + { + "type": "function_call", + "call_id": "call_move", + "name": "move", + "arguments": '{"flag":true}{"box":"A","limit":50}', + } + ] + } + + assert get_tool_calls_from_response(response) == [ + {"id": "call_move", "name": "move", "arguments": {"flag": True}}, + {"id": "call_move__concat_1", "name": "move", "arguments": {"box": "A", "limit": 50}}, + ] + + +def test_get_tool_calls_from_response_collapses_identical_concatenated_arguments(): + response: Final = _chat_tool_response(_function_tool_call("call_move", "move", '{"flag":true}' * 3)) + + assert get_tool_calls_from_response(response) == [ + {"id": "call_move", "name": "move", "arguments": {"flag": True}}, + ] + + +def test_get_tool_calls_from_response_does_not_expand_a_valid_json_array(): + response: Final = _chat_tool_response(_function_tool_call("call_batch", "batch", '[{"a":1},{"b":2}]')) + + assert get_tool_calls_from_response(response) == [ + {"id": "call_batch", "name": "batch", "arguments": {}}, + ] + + +@pytest.mark.parametrize("arguments", ('{"a":1}{"b":', '0{"x":1}')) +def test_get_tool_calls_from_response_drops_partial_concatenated_arguments(arguments: str, caplog): + response: Final = _chat_tool_response(_function_tool_call("call_move", "move", arguments)) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + tool_calls: Final = get_tool_calls_from_response(response) + + assert tool_calls == [{"id": "call_move", "name": "move", "arguments": {}}] + assert "Failed to parse tool call arguments" in caplog.text + + +@pytest.mark.parametrize(("count", "expands"), ((8, True), (9, False))) +def test_get_tool_calls_from_response_caps_distinct_concatenated_arguments(count: int, expands: bool): + raw = _concatenated_json(*({"n": index} for index in range(count))) + response: Final = _chat_tool_response(_function_tool_call("call", "move", raw)) + + tool_calls: Final = get_tool_calls_from_response(response) + + if expands: + assert [call["id"] for call in tool_calls] == ["call", *(f"call__concat_{index}" for index in range(1, count))] + assert [call["arguments"] for call in tool_calls] == [{"n": index} for index in range(count)] + return + assert tool_calls == [{"id": "call", "name": "move", "arguments": {}}] + + +def test_get_tool_calls_from_response_skips_concat_ids_taken_by_a_sibling(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("call", "move", raw), + _function_tool_call("call__concat_1", "look", '{"x":1}'), + ) + + assert [call["id"] for call in get_tool_calls_from_response(response)] == [ + "call", + "call__concat_2", + "call__concat_1", + ] + + +def test_get_tool_calls_from_response_keeps_sanitized_concat_ids_distinct(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("a:b", "move", raw), + _function_tool_call("a_b__concat_1", "look", '{"x":1}'), + ) + + ids: Final = [call["id"] for call in get_tool_calls_from_response(response)] + sanitized: Final = [_sanitize_anthropic_tool_use_id(call_id) for call_id in ids if isinstance(call_id, str)] + + assert len(sanitized) == len(set(sanitized)) + assert ids == ["a:b", "a:b__concat_2", "a_b__concat_1"] + + +def test_get_tool_calls_from_response_bumps_suffix_when_sibling_sanitizes_onto_it(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("a_b", "move", raw), + _function_tool_call("a:b__concat_1", "look", '{"x":1}'), + ) + + ids: Final = [call["id"] for call in get_tool_calls_from_response(response)] + sanitized: Final = [_sanitize_anthropic_tool_use_id(call_id) for call_id in ids if isinstance(call_id, str)] + + assert len(ids) == len(sanitized) + assert len(sanitized) == len(set(sanitized)) + assert ids == ["a_b", "a_b__concat_2", "a:b__concat_1"] + + +def test_get_tool_calls_from_response_continues_concat_suffixes_per_sanitized_base(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("x", "move", raw), + _function_tool_call("x", "move", raw), + ) + + assert [call["id"] for call in get_tool_calls_from_response(response)] == [ + "x", + "x__concat_1", + "x", + "x__concat_2", + ] + + +def test_get_tool_calls_from_response_skips_a_run_of_reserved_concat_ids(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + siblings: Final = tuple(_function_tool_call(f"call__concat_{index}", "look", '{"x":1}') for index in range(1, 51)) + response: Final = _chat_tool_response(_function_tool_call("call", "move", raw), *siblings) + + ids: Final = [call["id"] for call in get_tool_calls_from_response(response)] + + assert ids[0] == "call" + assert ids[1] == "call__concat_51" + + +def test_get_tool_calls_from_response_reserves_concat_ids_across_choices(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = { + "choices": [ + {"message": {"tool_calls": [_function_tool_call("call", "move", raw)]}}, + {"message": {"tool_calls": [_function_tool_call("call__concat_1", "look", '{"x":1}')]}}, + ] + } + + assert [call["id"] for call in get_tool_calls_from_response(response, include_all_choices=True)] == [ + "call", + "call__concat_2", + "call__concat_1", + ] + + +def test_get_tool_calls_from_response_does_not_invent_ids_for_a_missing_call_id(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response(_function_tool_call(None, "move", raw)) + + tool_calls: Final = get_tool_calls_from_response(response) + + assert len(tool_calls) == 2 + assert all(call["id"] is None for call in tool_calls) + assert [call["arguments"] for call in tool_calls] == [{"a": 1}, {"b": 2}] + + def test_group_tool_exchanges_pairs_assistant_with_its_tool_rows(): from litellm.litellm_core_utils.prompt_templates.factory import group_tool_exchanges @@ -3625,9 +3657,7 @@ def test_bedrock_converse_pdf_only_user_message_gets_text_block(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert len(result) == 1 assert any("document" in block for block in result[0]["content"]) @@ -3645,9 +3675,7 @@ def test_bedrock_converse_document_with_text_gets_no_extra_text_block(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert _text_blocks(result[0]) == ["summarize this"] @@ -3660,9 +3688,7 @@ def test_bedrock_converse_image_only_user_message_gets_no_text_block(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert any("image" in block for block in result[0]["content"]) assert _text_blocks(result[0]) == [] @@ -3705,9 +3731,7 @@ def test_bedrock_converse_tool_round_trip_document_injects_text_before_cache_poi }, ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert _text_blocks(result[0]) == ["read the pdf"] document_message = result[-1] From 829cba1bf18c22593ddf65737e30d1b905651259 Mon Sep 17 00:00:00 2001 From: Thippaluri Yaseen Basha Date: Sun, 27 Sep 2026 09:48:33 +0530 Subject: [PATCH 15/65] fix(gemini): forward seed to the Gemini API instead of rejecting it (#43197) * fix(gemini): forward seed to the Gemini API instead of rejecting it The gemini/ provider left seed out of its supported params, so requests with seed failed with UnsupportedParamsError, or lost the seed silently when drop_params was on. The Gemini API accepts generationConfig.seed and the inherited mapping already translates it, so adding it to the allowlist is enough * test(gemini): assert the forwarded seed without mutating shared state --- litellm/llms/gemini/chat/transformation.py | 1 + ...test_vertex_and_google_ai_studio_gemini.py | 28 +++++++++++++++++++ 2 files changed, 29 insertions(+) diff --git a/litellm/llms/gemini/chat/transformation.py b/litellm/llms/gemini/chat/transformation.py index 285350aecba..cae0ba49c3f 100644 --- a/litellm/llms/gemini/chat/transformation.py +++ b/litellm/llms/gemini/chat/transformation.py @@ -96,6 +96,7 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig): "logprobs", "frequency_penalty", "presence_penalty", + "seed", "modalities", "parallel_tool_calls", "web_search_options", diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 7548f3c2daa..fd735afb16e 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -3413,6 +3413,34 @@ def test_google_ai_studio_presence_penalty_supported(): assert "presence_penalty" in supported_params +@pytest.mark.asyncio +@pytest.mark.parametrize("drop_params", [False, True]) +async def test_google_ai_studio_forwards_seed_to_generation_config(drop_params: bool): + def echo_seed_sent_upstream(request: httpx.Request) -> httpx.Response: + seed_sent: Final = json.loads(request.content).get("generationConfig", {}).get("seed") + return httpx.Response( + 200, + json={ + "candidates": [ + {"content": {"parts": [{"text": f"seed={seed_sent}"}], "role": "model"}, "finishReason": "STOP"} + ], + "usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }, + request=request, + ) + + response: Final = await litellm.acompletion( + model="gemini/gemini-3.8-flash", + messages=[{"role": "user", "content": "hi"}], + seed=42, + drop_params=drop_params, + api_key="fake-gemini-key", + client=AsyncHTTPHandler(transport=httpx.MockTransport(echo_seed_sent_upstream)), + ) + + assert response.choices[0].message.content == "seed=42" + + # ==================== Tool Type Separation Tests ==================== # These tests verify that each Tool object contains exactly one type per Vertex AI API spec # Ref: https://cloud.google.com/vertex-ai/generative-ai/docs/reference/rest/v1beta1/Tool From 2101c860c25546250fa55c5a67653093a9d28873 Mon Sep 17 00:00:00 2001 From: Stewart Park <388348+stewartpark@users.noreply.github.com> Date: Sat, 26 Sep 2026 21:26:38 -0700 Subject: [PATCH 16/65] fix(vertex_ai): make Gemma fake streams work with traced Responses (#43147) * test(vertex_ai): reproduce traced Gemma Responses stream failure * fix(vertex_ai): wrap Gemma fake streams for Responses tracing * test(vertex_ai): cover Gemma traced streams and usage options * test(vertex_ai): inject gemma test deps and assert hidden usage accounting Replace class-level patches in the Vertex AI shard test with the provider's documented dependency-injection seams (httpx.MockTransport client + credential cache), and pin the default/omit-usage trace behavior: LiteLLM still accounts all tokens; ddtrace's metric is absent by design, asserted rather than silent. Mutation-checked: commenting out CustomStreamWrapper chunk accumulation turns the new assertions red; restoring them turns green. * test(vertex_ai): drop explanatory comment from usage-option assertions --- .../vertex_gemma_models/transformation.py | 30 ++++- .../test_vertex_gemma_transformation.py | 93 ++++++++++++++ .../test_vertex_gemma_transformation.py | 117 +++++++++++++++++- 3 files changed, 229 insertions(+), 11 deletions(-) create mode 100644 tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index ea97f0a0a9a..33922e38674 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -28,8 +28,8 @@ from litellm.types.utils import ModelResponse if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer - from litellm.llms.base_llm.base_model_iterator import MockResponseIterator def parse_vertex_gemma_container_error(predictions: object) -> VertexGemmaContainerError | None: @@ -73,7 +73,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): self, model_response: ModelResponse, stream: bool, - ) -> "ModelResponse | MockResponseIterator": + model: str, + logging_obj: "LiteLLMLoggingObj", + ) -> "ModelResponse | CustomStreamWrapper": """ Helper method to return fake stream iterator if streaming is requested. @@ -82,12 +84,18 @@ class VertexGemmaConfig(OpenAIGPTConfig): stream: Whether streaming was requested Returns: - MockResponseIterator if stream=True, otherwise the model_response + CustomStreamWrapper if stream=True, otherwise the model_response """ if stream: + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.base_llm.base_model_iterator import MockResponseIterator - return MockResponseIterator(model_response=model_response) + return CustomStreamWrapper( + completion_stream=MockResponseIterator(model_response=model_response), + model=model, + custom_llm_provider="vertex_ai", + logging_obj=logging_obj, + ) return model_response def transform_request( @@ -373,7 +381,12 @@ class VertexGemmaConfig(OpenAIGPTConfig): ) # Return fake stream iterator if streaming was requested - return self._handle_fake_stream_response(model_response=model_response, stream=stream) + return self._handle_fake_stream_response( + model_response=model_response, + stream=stream, + model=model, + logging_obj=logging_obj, + ) async def _async_completion( self, @@ -463,4 +476,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): ) # Return fake stream iterator if streaming was requested - return self._handle_fake_stream_response(model_response=model_response, stream=stream) + return self._handle_fake_stream_response( + model_response=model_response, + stream=stream, + model=model, + logging_obj=logging_obj, + ) diff --git a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py new file mode 100644 index 00000000000..294d26b2e58 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -0,0 +1,93 @@ +import json +from collections.abc import AsyncIterator +from types import SimpleNamespace +from typing import Any, cast + +import httpx +import pytest + +import litellm +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.main import vertex_gemma_chat_completion +from litellm.types.llms.openai import OutputTextDeltaEvent, ResponseCompletedEvent, ResponsesAPIStreamingResponse + +_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict" +_MESSAGES = [{"role": "user", "content": "Reply exactly READY"}] +_FAKE_CREDENTIALS = "gemma-test-credentials" + + +def _vertex_response(): + return { + "predictions": { + "id": "chatcmpl-stream-test", + "created": 1759863903, + "model": "google/gemma-3-12b-it", + "object": "chat.completion", + "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "READY"}}], + "usage": {"prompt_tokens": 14, "completion_tokens": 1, "total_tokens": 15}, + } + } + + +@pytest.fixture(autouse=True) +def _cached_access_token(): + """Serve a fake token from the handler's credential cache so no auth round-trip runs.""" + cache = vertex_gemma_chat_completion._credentials_project_mapping + key = (_FAKE_CREDENTIALS, "test") + cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test") + yield + cache.pop(key, None) + + +def test_sync_gemma_stream(): + captured: dict[str, Any] = {} + + def handle(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_vertex_response()) + + stream = litellm.completion( + model="vertex_ai/gemma/test-model", + messages=_MESSAGES, + stream=True, + api_base=_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_CREDENTIALS, + client=httpx.Client(transport=httpx.MockTransport(handle)), + ) + + assert isinstance(stream, CustomStreamWrapper) + chunks = list(stream) + + assert "stream" not in captured["body"]["instances"][0] + assert len(chunks) == 2 + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_async_gemma_responses_stream(): + captured: dict[str, Any] = {} + + def handle(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_vertex_response()) + + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base=_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_CREDENTIALS, + client=httpx.AsyncClient(transport=httpx.MockTransport(handle)), + ) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + + assert "stream" not in captured["body"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage is not None + assert events[-1].response.usage.total_tokens == 15 diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index e5ca31833ce..97f4f290958 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -5,11 +5,18 @@ Maps to: litellm/llms/vertex_ai/vertex_gemma_models/transformation.py """ import json +from collections.abc import AsyncIterator +from typing import cast from unittest.mock import AsyncMock, Mock, patch import pytest import litellm +from litellm.types.llms.openai import ( + OutputTextDeltaEvent, + ResponseCompletedEvent, + ResponsesAPIStreamingResponse, +) @pytest.fixture(autouse=True) @@ -439,8 +446,9 @@ class TestVertexGemmaCompletion: Verifies: 1. Request body does NOT include 'stream' parameter (model doesn't support it) - 2. Response returns a MockResponseIterator that yields chunks + 2. Response wraps a MockResponseIterator and yields chunks """ + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.base_llm.base_model_iterator import MockResponseIterator # Mock Vertex response @@ -502,8 +510,8 @@ class TestVertexGemmaCompletion: vertex_location="us-central1", ) - # Verify the response is a MockResponseIterator - assert isinstance(response, MockResponseIterator), f"Expected MockResponseIterator, got {type(response)}" + assert isinstance(response, CustomStreamWrapper) + assert isinstance(response.completion_stream, MockResponseIterator) # Verify the request sent to Vertex does NOT include 'stream' call_args = mock_client.post.call_args @@ -520,8 +528,9 @@ class TestVertexGemmaCompletion: async for chunk in response: chunks.append(chunk) - # Should get exactly one chunk (fake streaming) - assert len(chunks) == 1, f"Expected 1 chunk from fake stream, got {len(chunks)}" + assert len(chunks) == 2 + assert chunks[1].choices[0].finish_reason == "stop" + assert all(getattr(chunk, "usage", None) is None for chunk in chunks) # Verify the chunk has the expected content chunk = chunks[0] @@ -529,6 +538,104 @@ class TestVertexGemmaCompletion: assert len(chunk.choices) > 0 assert chunk.choices[0].delta.content == "Streaming test response" + @pytest.mark.asyncio + async def test_aresponses_streams_vertex_gemma_with_llm_tracing(self): + pytest.importorskip("ddtrace") + from ddtrace.contrib.internal.litellm.patch import patch as patch_litellm + from ddtrace.contrib.internal.litellm.patch import unpatch as unpatch_litellm + from ddtrace.llmobs._integrations.base_stream_handler import TracedAsyncStream + + from litellm.responses.litellm_completion_transformation.streaming_iterator import ( + LiteLLMCompletionStreamingIterator, + ) + + reply = Mock(status_code=200) + reply.json.return_value = _make_gemma_vertex_response(content="READY") + client = Mock() + client.post = AsyncMock(return_value=reply) + + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client), + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + patch_litellm() + try: + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base="https://example.invalid/v1/projects/test-project/locations/us-central1/endpoints/test:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + bridge = cast(LiteLLMCompletionStreamingIterator, response) + traced_stream = bridge.litellm_custom_stream_wrapper + assert isinstance(traced_stream, TracedAsyncStream) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + span = traced_stream.handler.primary_span + assert span.finished + assert span.get_tag("_dd.llmobs.span_kind") == "llm" + assert span.get_metric("_dd.llmobs.total_tokens") == 114 + finally: + unpatch_litellm() + + assert "stream" not in client.post.call_args.kwargs["json"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage.total_tokens == 114 + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream_options", [None, {"include_usage": False}, {"include_usage": True}]) + async def test_acompletion_stream_respects_usage_option_with_llm_tracing(self, stream_options): + pytest.importorskip("ddtrace") + from ddtrace.contrib.internal.litellm.patch import patch as patch_litellm + from ddtrace.contrib.internal.litellm.patch import unpatch as unpatch_litellm + + reply = Mock(status_code=200) + reply.json.return_value = _make_gemma_vertex_response(content="READY") + client = Mock(post=AsyncMock(return_value=reply)) + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client), + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + patch_litellm() + try: + stream = await litellm.acompletion( + model="vertex_ai/gemma/test-model", + messages=[{"role": "user", "content": "Reply exactly READY"}], + stream=True, + **({"stream_options": stream_options} if stream_options is not None else {}), + api_base="https://example.invalid/v1/projects/test-project/locations/us-central1/endpoints/test:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + chunks = [chunk async for chunk in stream] + span = stream.handler.primary_span + assert span.finished + assert span.get_tag("_dd.llmobs.span_kind") == "llm" + finally: + unpatch_litellm() + + assert len(chunks) == (3 if stream_options and stream_options["include_usage"] else 2) + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + if stream_options and stream_options["include_usage"]: + assert chunks[-1].choices[0].delta.content is None + assert chunks[-1].usage.total_tokens == 114 + assert span.get_metric("_dd.llmobs.total_tokens") == 114 + else: + from litellm.litellm_core_utils.streaming_handler import calculate_total_usage + + assert all(getattr(chunk, "usage", None) is None for chunk in chunks) + assert calculate_total_usage(chunks=stream.chunks).total_tokens == 114 + assert span.get_metric("_dd.llmobs.total_tokens") is None + @pytest.mark.asyncio async def test_acompletion_filters_stream_and_stream_options(self): """ From 491d454826342aa8b53aa69edd0242ba3b6f8b4d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 04:51:56 +0000 Subject: [PATCH 17/65] fix(responses): emit the reasoning item on streaming /v1/responses for signature-only thinking (#43414) * fix(responses): emit the reasoning item on streaming /v1/responses for signature-only thinking Anthropic models return thinking blocks with empty text and the reasoning carried in the signature: Claude Fable 5.1 and Claude Opus 5.5 by default, and Bedrock adaptive thinking with or without an effort. On streaming /v1/responses the chat->Responses bridge opened a reasoning output item only on reasoning_content text (LiteLLMCompletionStreamingIterator._ensure_output_item_for_chunk), and ChunkProcessor.get_combined_thinking_content kept an assembled thinking block only when it had thinking text. Such a response emitted no reasoning item mid-stream and none in response.completed, so a streaming Responses client could not replay the reasoning even though the reasoning tokens were billed. Non-streaming /v1/responses was unaffected. Open the reasoning item when the delta carries a signed or redacted thinking block, and keep a signed block through stream assembly even when its thinking text is empty. Unsigned text-only fragments are still dropped. The reasoning-text path is unchanged. (cherry picked from commit bc9b6f8a5c3ac9a2b46e3f9f01f7c2c5f9b688e7) * test(vertex_ai): move orphaned gemma streaming tests into the llm-vertex-ai shard PR #43147 left a copy of the Gemma streaming tests under tests/test_litellm/llms, a tree no CI shard claims, which broke assert-ci-coverage and assert-shard-coverage on main. Fold the two streaming tests into the existing tests/unit/llms/vertex_ai file so the llm-vertex-ai shard runs them Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Chloe Lu Co-authored-by: Krrish Dholakia Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../streaming_chunk_builder_utils.py | 2 +- .../streaming_iterator.py | 7 +- .../test_vertex_gemma_transformation.py | 93 ------------------- .../test_streaming_chunk_builder_utils.py | 25 +++++ .../test_vertex_gemma_transformation.py | 78 ++++++++++++++++ .../test_streaming_iterator_transformation.py | 40 ++++++++ 6 files changed, 150 insertions(+), 95 deletions(-) delete mode 100644 tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index d975c3551f3..67684a230e3 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -685,7 +685,7 @@ class ChunkProcessor: def _flush_thinking_block() -> None: nonlocal current_thinking_text_parts, current_signature - if len(current_thinking_text_parts) > 0 and current_signature: + if current_signature: thinking_blocks.append( ChatCompletionThinkingBlock( type="thinking", diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index 5173cd04a89..21a33c17ab8 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -73,6 +73,11 @@ def _output_items_with_id(items: tuple[Any, ...], item_type: str, item_id: str | ) +def _delta_has_signed_thinking_block(delta: object) -> bool: + blocks: Final = getattr(delta, "thinking_blocks", None) or () + return any(isinstance(b, dict) and (b.get("signature") or b.get("data")) for b in blocks) + + class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): """ Async iterator for processing streaming responses from the Responses API. @@ -936,7 +941,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self.sent_output_item_added_event = True # Reasoning-first - if hasattr(delta, "reasoning_content") and delta.reasoning_content: + if (hasattr(delta, "reasoning_content") and delta.reasoning_content) or _delta_has_signed_thinking_block(delta): self._reasoning_active = True if self._cached_reasoning_item_id is None: self._cached_reasoning_item_id = f"rs_{uuid.uuid4()}" diff --git a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py deleted file mode 100644 index 294d26b2e58..00000000000 --- a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ /dev/null @@ -1,93 +0,0 @@ -import json -from collections.abc import AsyncIterator -from types import SimpleNamespace -from typing import Any, cast - -import httpx -import pytest - -import litellm -from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper -from litellm.main import vertex_gemma_chat_completion -from litellm.types.llms.openai import OutputTextDeltaEvent, ResponseCompletedEvent, ResponsesAPIStreamingResponse - -_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict" -_MESSAGES = [{"role": "user", "content": "Reply exactly READY"}] -_FAKE_CREDENTIALS = "gemma-test-credentials" - - -def _vertex_response(): - return { - "predictions": { - "id": "chatcmpl-stream-test", - "created": 1759863903, - "model": "google/gemma-3-12b-it", - "object": "chat.completion", - "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "READY"}}], - "usage": {"prompt_tokens": 14, "completion_tokens": 1, "total_tokens": 15}, - } - } - - -@pytest.fixture(autouse=True) -def _cached_access_token(): - """Serve a fake token from the handler's credential cache so no auth round-trip runs.""" - cache = vertex_gemma_chat_completion._credentials_project_mapping - key = (_FAKE_CREDENTIALS, "test") - cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test") - yield - cache.pop(key, None) - - -def test_sync_gemma_stream(): - captured: dict[str, Any] = {} - - def handle(request: httpx.Request) -> httpx.Response: - captured["body"] = json.loads(request.content) - return httpx.Response(200, json=_vertex_response()) - - stream = litellm.completion( - model="vertex_ai/gemma/test-model", - messages=_MESSAGES, - stream=True, - api_base=_VERTEX_URL, - vertex_project="test", - vertex_location="us-central1", - vertex_credentials=_FAKE_CREDENTIALS, - client=httpx.Client(transport=httpx.MockTransport(handle)), - ) - - assert isinstance(stream, CustomStreamWrapper) - chunks = list(stream) - - assert "stream" not in captured["body"]["instances"][0] - assert len(chunks) == 2 - assert chunks[0].choices[0].delta.content == "READY" - assert chunks[1].choices[0].finish_reason == "stop" - - -@pytest.mark.asyncio -async def test_async_gemma_responses_stream(): - captured: dict[str, Any] = {} - - def handle(request: httpx.Request) -> httpx.Response: - captured["body"] = json.loads(request.content) - return httpx.Response(200, json=_vertex_response()) - - response = await litellm.aresponses( - model="vertex_ai/gemma/test-model", - input="Reply exactly READY", - stream=True, - api_base=_VERTEX_URL, - vertex_project="test", - vertex_location="us-central1", - vertex_credentials=_FAKE_CREDENTIALS, - client=httpx.AsyncClient(transport=httpx.MockTransport(handle)), - ) - events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] - - assert "stream" not in captured["body"]["instances"][0] - assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) - assert isinstance(events[-1], ResponseCompletedEvent) - assert events[-1].response.usage is not None - assert events[-1].response.usage.total_tokens == 15 diff --git a/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py index af763da2d87..aaf877df364 100644 --- a/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -236,6 +236,31 @@ def test_get_combined_thinking_content_preserves_interleaved_blocks(): assert result[2]["signature"] == "sig_block2" +def test_get_combined_thinking_content_keeps_signed_block_without_thinking_text(): + chunks: Final = [ + ModelResponseStream( + id="chatcmpl-123", + object="chat.completion.chunk", + created=1234567890, + model="claude-sonnet-4-20250514", + choices=[ + StreamingChoices( + index=0, + delta=Delta(thinking_blocks=[{"type": "thinking", "thinking": "", "signature": "sig_only"}]), + finish_reason=None, + ) + ], + ) + ] + + result: Final = ChunkProcessor(chunks=chunks).get_combined_thinking_content(chunks) + + assert result is not None + assert [(block["type"], block["thinking"], block["signature"]) for block in result] == [ + ("thinking", "", "sig_only") + ] + + def test_cache_read_input_tokens_retained(): chunk1 = ModelResponseStream( id="chatcmpl-95aabb85-c39f-443d-ae96-0370c404d70c", diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index 97f4f290958..92e684e42c8 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -1303,3 +1303,81 @@ class TestVertexGemmaCompletion: mock_async_post.assert_awaited_once() assert mock_async_post.call_args.kwargs["client"] is None assert response.choices[0].message.content == "default async handler fallback" + + +_GEMMA_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict" +_FAKE_GEMMA_CREDENTIALS = "gemma-test-credentials" + + +@pytest.fixture +def _gemma_cached_access_token(): + """Serve a fake token from the handler's credential cache so no auth round-trip runs.""" + from types import SimpleNamespace + + from litellm.main import vertex_gemma_chat_completion + + cache = vertex_gemma_chat_completion._credentials_project_mapping + key = (_FAKE_GEMMA_CREDENTIALS, "test") + cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test") + yield + cache.pop(key, None) + + +def test_sync_gemma_stream(_gemma_cached_access_token): + import httpx + + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper + + captured = {} + + def handle(request): + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15)) + + stream = litellm.completion( + model="vertex_ai/gemma/test-model", + messages=[{"role": "user", "content": "Reply exactly READY"}], + stream=True, + api_base=_GEMMA_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_GEMMA_CREDENTIALS, + client=httpx.Client(transport=httpx.MockTransport(handle)), + ) + + assert isinstance(stream, CustomStreamWrapper) + chunks = list(stream) + + assert "stream" not in captured["body"]["instances"][0] + assert len(chunks) == 2 + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_async_gemma_responses_stream(_gemma_cached_access_token): + import httpx + + captured = {} + + def handle(request): + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15)) + + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base=_GEMMA_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_GEMMA_CREDENTIALS, + client=httpx.AsyncClient(transport=httpx.MockTransport(handle)), + ) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + + assert "stream" not in captured["body"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage is not None + assert events[-1].response.usage.total_tokens == 15 diff --git a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py index 8fbba0dbf87..041bcf1b6d7 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py +++ b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py @@ -978,6 +978,25 @@ def _reasoning_chunk(reasoning: str, finish_reason: str | None = None) -> ModelR ) +def _signature_only_thinking_chunk(signature: str) -> ModelResponseStream: + return ModelResponseStream( + id=CHAT_COMPLETION_ID, + created=1748575031, + model="claude-haiku-4-5", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta( + role="assistant", + thinking_blocks=[{"type": "thinking", "thinking": "", "signature": signature}], + ), + finish_reason=None, + ) + ], + ) + + async def _collect_events( iterator: LiteLLMCompletionStreamingIterator, sync_mode: bool ) -> list[BaseLiteLLMOpenAIResponseObject]: @@ -1015,6 +1034,27 @@ async def test_tool_only_stream_emits_no_message_item_events(sync_mode: bool): assert any(getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED for event in events) +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_signature_only_thinking_streams_a_replayable_reasoning_item(sync_mode: bool): + iterator: Final = _build_iterator([_signature_only_thinking_chunk("sig_only"), _chunk("4", finish_reason="stop")]) + + events: Final = await _collect_events(iterator, sync_mode) + + added_item_types: Final = [ + event.item.type + for event in events + if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + ] + completed: Final = next( + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ) + reasoning_items: Final = [item for item in completed.response.output if getattr(item, "type", None) == "reasoning"] + assert added_item_types[0] == "reasoning" + assert len(reasoning_items) == 1 + assert json.loads(reasoning_items[0].encrypted_content)[0]["signature"] == "sig_only" + + @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_reasoning_then_text_announces_message_item_before_text_events(sync_mode: bool): From 9a0ff249d5935ca73216d19597603083e6a0845c Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 05:02:35 +0000 Subject: [PATCH 18/65] fix(anthropic): forward the per-turn-control beta to Azure AI Foundry (#43415) Co-authored-by: Krrish Dholakia Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/anthropic_beta_headers_config.json | 2 +- .../messages/test_anthropic_messages_per_turn_control.py | 8 +++++++- 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/litellm/anthropic_beta_headers_config.json b/litellm/anthropic_beta_headers_config.json index 3a28d65e47c..71e7081b440 100644 --- a/litellm/anthropic_beta_headers_config.json +++ b/litellm/anthropic_beta_headers_config.json @@ -57,7 +57,7 @@ "mcp-servers-2025-12-04": null, "output-128k-2025-02-19": null, "structured-output-2024-03-01": null, - "per-turn-control-2026-07-01": null, + "per-turn-control-2026-07-01": "per-turn-control-2026-07-01", "prompt-caching-scope-2026-01-05": "prompt-caching-scope-2026-01-05", "skills-2025-10-02": "skills-2025-10-02", "structured-outputs-2025-11-13": "structured-outputs-2025-11-13", diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py index 557305a945c..4197192e4af 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py @@ -95,13 +95,19 @@ def test_added_per_turn_control_beta_survives_the_anthropic_allowlist(): assert PER_TURN_CONTROL in _betas(filtered) -@pytest.mark.parametrize("provider", ["bedrock", "bedrock_converse", "vertex_ai", "azure_ai", "databricks"]) +@pytest.mark.parametrize("provider", ["bedrock", "bedrock_converse", "vertex_ai", "databricks"]) def test_per_turn_control_beta_is_dropped_for_providers_without_it(provider): filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider=provider) assert "anthropic-beta" not in filtered +def test_per_turn_control_beta_is_forwarded_for_azure_ai(): + filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider="azure_ai") + + assert _betas(filtered) == {PER_TURN_CONTROL} + + def test_json_provider_passthrough_adds_per_turn_control_beta(): config = JSONProviderAnthropicMessagesConfig( SimpleProviderConfig( From c1f761eba50bb344259bb7a3ff2ef538ff94600f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 05:02:44 +0000 Subject: [PATCH 19/65] test(vertex_ai): move stray Gemma streaming tests to tests/unit so CI coverage passes (#43422) Co-authored-by: Krrish Dholakia Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../vertex_gemma_models/test_vertex_gemma_transformation.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index 92e684e42c8..efe97ce33a8 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -1332,7 +1332,7 @@ def test_sync_gemma_stream(_gemma_cached_access_token): def handle(request): captured["body"] = json.loads(request.content) - return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15)) + return httpx.Response(200, json=_make_gemma_vertex_response(content="READY")) stream = litellm.completion( model="vertex_ai/gemma/test-model", @@ -1362,7 +1362,7 @@ async def test_async_gemma_responses_stream(_gemma_cached_access_token): def handle(request): captured["body"] = json.loads(request.content) - return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15)) + return httpx.Response(200, json=_make_gemma_vertex_response(content="READY")) response = await litellm.aresponses( model="vertex_ai/gemma/test-model", @@ -1380,4 +1380,4 @@ async def test_async_gemma_responses_stream(_gemma_cached_access_token): assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) assert isinstance(events[-1], ResponseCompletedEvent) assert events[-1].response.usage is not None - assert events[-1].response.usage.total_tokens == 15 + assert events[-1].response.usage.total_tokens == 114 From b831e9b4ac8a1f221704663acb9cb542ba40fd11 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 05:10:24 +0000 Subject: [PATCH 20/65] fix(bedrock): keep the provider status code on unprocessable image errors (#43416) Co-authored-by: Krrish Dholakia Co-authored-by: dbalintx Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../exception_mapping_utils.py | 2 +- .../test_exception_mapping_utils.py | 37 ++++++++++++++++++- 2 files changed, 37 insertions(+), 2 deletions(-) diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index f09dd9fe75a..0fdfb301291 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -954,7 +954,7 @@ def _map_bedrock_exception( llm_provider="bedrock", response=getattr(original_exception, "response", None), ) - elif "Could not process image" in error_str: + elif "Could not process image" in error_str and getattr(original_exception, "status_code", 500) == 500: raise litellm.InternalServerError( message=f"BedrockException - {error_str}", model=model, diff --git a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py index 9fce0441a58..9de768ea47b 100644 --- a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py +++ b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py @@ -1280,7 +1280,7 @@ def test_bedrock_500_preserves_provider_response_headers(): "bedrock", 400, '{"message":"Could not process image"}', - litellm.InternalServerError, + litellm.BadRequestError, ), ], ) @@ -1313,6 +1313,41 @@ def test_bedrock_classified_errors_preserve_provider_response_headers( assert exc_info.value.response.headers["x-amzn-requestid"] == "req-classified" +@pytest.mark.parametrize( + "status_code, expected_exception", + [ + (400, litellm.BadRequestError), + (503, litellm.ServiceUnavailableError), + (500, litellm.InternalServerError), + ], +) +def test_bedrock_unprocessable_image_keeps_provider_status_code(status_code, expected_exception): + """An unprocessable image maps to the status Bedrock sent, so the 400 it returns stays a client error.""" + provider_message = '{"message":"The model returned the following errors: Could not process image"}' + provider_response = httpx.Response( + status_code=status_code, + text=provider_message, + request=httpx.Request("POST", "https://bedrock-runtime.us-east-1.amazonaws.com/"), + ) + original_exception = BedrockError( + status_code=status_code, + message=provider_message, + headers=provider_response.headers, + response=provider_response, + ) + + with pytest.raises(expected_exception) as exc_info: + exception_type( + model="anthropic.claude-haiku-4-5-20251001-v1:0", + original_exception=original_exception, + custom_llm_provider="bedrock", + completion_kwargs={}, + extra_kwargs={}, + ) + + assert exc_info.value.status_code == status_code + + @pytest.mark.parametrize( "status_code, provider_message", [ From 4274bdda441527c8ab9601c44e4dbbc79a63e707 Mon Sep 17 00:00:00 2001 From: Anmol Jaiswal <68013660+anmolg1997@users.noreply.github.com> Date: Sun, 27 Sep 2026 10:44:53 +0530 Subject: [PATCH 21/65] fix(vertex_ai): stop importing the vertexai SDK in partner-model completion (#42274) completion() imported vertexai only to check that the package exists. Partner models are reached with an authenticated httpx client and never use that SDK, the same reasoning count_tokens in this file already follows (#28084). The import loads all of google-cloud-aiplatform on the first request of every process and made a google-auth-only install fail with a 400 --- .../vertex_ai_partner_models/main.py | 9 +---- .../test_partner_models_credential_reuse.py | 38 +++++++++++++++++++ 2 files changed, 39 insertions(+), 8 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py index 2a36e5cc785..40503edbb9e 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py @@ -109,8 +109,6 @@ class VertexAIPartnerModels(VertexBase): client=None, ): try: - import vertexai - from litellm.llms.anthropic.chat import AnthropicChatCompletion from litellm.llms.codestral.completion.handler import ( CodestralTextCompletion, @@ -119,14 +117,9 @@ class VertexAIPartnerModels(VertexBase): except Exception as e: raise VertexAIError( status_code=400, - message=f"""vertexai import failed please run `pip install -U "google-cloud-aiplatform>=1.38"`. Got error: {e}""", + message=f"Failed to import a partner model handler. Got error: {e}", ) - if not (hasattr(vertexai, "preview") or hasattr(vertexai.preview, "language_models")): - raise VertexAIError( - status_code=400, - message="""Upgrade vertex ai. Run `pip install "google-cloud-aiplatform>=1.38"`""", - ) try: access_token, project_id = self._ensure_access_token( credentials=vertex_credentials, diff --git a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py index b20442a032e..8e6270e41a1 100644 --- a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py +++ b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py @@ -127,6 +127,44 @@ class TestPartnerModelsCredentialReuse: assert mock_load.call_count == 1 + def test_completion_works_without_the_vertexai_sdk(self): + """completion() reaches the HTTP handler when `import vertexai` raises ImportError.""" + partner = VertexAIPartnerModels() + + with ( + patch.dict(sys.modules, {"vertexai": None}), + patch.object( + partner, + "_ensure_access_token", + return_value=("cached-token", "test-project"), + ), + patch( + "litellm.llms.vertex_ai.vertex_ai_partner_models.main.base_llm_http_handler" + ) as mock_handler, + ): + mock_handler.completion.return_value = "response" + + result = partner.completion( + model="meta/llama-3.1-405b-instruct-maas", + messages=[{"role": "user", "content": "hello"}], + model_response=MagicMock(), + print_verbose=lambda *a, **kw: None, + encoding=MagicMock(), + logging_obj=MagicMock(), + api_base=None, + optional_params={}, + custom_prompt_dict={}, + headers=None, + timeout=30.0, + litellm_params={}, + vertex_project="test-project", + vertex_location="us-central1", + vertex_credentials=None, + ) + + assert result == "response" + mock_handler.completion.assert_called_once() + class TestGemmaModelsCredentialReuse: def test_completion_uses_self_ensure_access_token(self): From 8e6d99d74a63c61e39628baeabea0c07dbeda5f4 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 22:28:31 -0700 Subject: [PATCH 22/65] fix(token_counter): count Gemini function_declarations tools (#43417) * fix(token_counter): count Gemini function_declarations tools Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(token_counter): skip non-dict tools when formatting definitions 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/litellm_core_utils/token_counter.py | 78 +++++++++++-------- .../litellm_core_utils/test_token_counter.py | 72 +++++++++++++++++ .../test_vertex_ai_context_caching.py | 9 ++- 3 files changed, 126 insertions(+), 33 deletions(-) diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 5d7956059e4..cdd2d0654be 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -951,7 +951,7 @@ def _count_content_list( ) -def _format_function_definitions(tools): +def _format_function_definitions(tools: Sequence[object]) -> str: """Formats tool definitions in the format that OpenAI appears to use. Based on https://github.com/forestwanglin/openai-java/blob/main/jtokkit/src/main/java/xyz/felh/openai/jtokkit/utils/TikTokenUtils.java """ @@ -959,41 +959,57 @@ def _format_function_definitions(tools): lines.append("namespace functions {") lines.append("") for tool in tools: - if not isinstance(tool, dict): + if not isinstance(tool, Mapping): continue - function = tool.get("function") - if not isinstance(function, dict): - # Anthropic tool shape → OpenAI function dict for token counting. - params = tool.get("input_schema") or tool.get("parameters") or {} - if not isinstance(params, dict): - params = {} - function = { - "name": tool.get("name"), - "description": tool.get("description"), - "parameters": params, - } - function_name = function.get("name") - if not function_name: - # Skip malformed tools missing a name to avoid emitting - # ``type None = ...`` which would produce inaccurate token counts. - continue - if function_description := function.get("description"): - lines.append(f"// {function_description}") - parameters = function.get("parameters") or {} - if not isinstance(parameters, dict): - parameters = {} - properties = parameters.get("properties") - if properties and properties.keys(): - lines.append(f"type {function_name} = (_: {{") - lines.append(_format_object_parameters(parameters, 0)) - lines.append("}) => any;") - else: - lines.append(f"type {function_name} = () => any;") - lines.append("") + for function in _function_definitions_for_tool(cast(Mapping[str, object], tool)): + lines.extend(_format_single_function_definition(function)) lines.append("} // namespace functions") return "\n".join(lines) +def _function_definitions_for_tool(tool: Mapping[str, object]) -> Iterable[Mapping[str, object]]: + function: Final = tool.get("function") + if isinstance(function, Mapping): + yield function + return + declarations: Final = tool.get("function_declarations") or tool.get("functionDeclarations") + if isinstance(declarations, list): + for declaration in declarations: + if isinstance(declaration, Mapping): + yield declaration + return + parameters: Final = tool.get("input_schema") or tool.get("parameters") or {} + normalized_parameters: Final = parameters if isinstance(parameters, Mapping) else {} + yield { + "name": tool.get("name"), + "description": tool.get("description"), + "parameters": normalized_parameters, + } + + +def _format_single_function_definition(function: Mapping[str, object]) -> tuple[str, ...]: + function_name: Final = function.get("name") + if not function_name: + return () + function_description: Final = function.get("description") + parameters_value: Final = function.get("parameters") or {} + parameters: Final = parameters_value if isinstance(parameters_value, Mapping) else {} + properties: Final = parameters.get("properties") + if isinstance(properties, Mapping) and properties: + return ( + *((f"// {function_description}",) if function_description else ()), + f"type {function_name} = (_: {{", + _format_object_parameters(parameters, 0), + "}) => any;", + "", + ) + return ( + *((f"// {function_description}",) if function_description else ()), + f"type {function_name} = () => any;", + "", + ) + + def _format_object_parameters(parameters, indent): properties: Final = parameters.get("properties") if not properties: diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index f7ded4f3fa8..c71b1496bdd 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -442,6 +442,78 @@ def test_token_counter_with_tools(message_count_pair): ), f"Expected {expected_tokens} tokens, got {counted_tokens}." +def test_token_counter_counts_gemini_function_declarations(): + openai_tools: Final = [ + { + "type": "function", + "function": { + "name": "lookup_weather", + "description": "Find current weather conditions for a location", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string", "description": "City and region"}, + "units": {"type": "string", "enum": ["celsius", "fahrenheit"]}, + }, + "required": ["location"], + }, + }, + } + ] + gemini_tools: Final = litellm.utils.get_optional_params( + model="gemini-2.5-pro", + custom_llm_provider="gemini", + tools=openai_tools, + )["tools"] + camel_case_tools: Final = [{"functionDeclarations": gemini_tools[0]["function_declarations"]}] + + openai_tokens: Final = token_counter_new( + model="gemini-2.5-pro", + messages=[{"role": "user", "content": "What's the weather?"}], + tools=openai_tools, + ) + gemini_tokens: Final = token_counter_new( + model="gemini-2.5-pro", + messages=[{"role": "user", "content": "What's the weather?"}], + tools=gemini_tools, + ) + camel_case_tokens: Final = token_counter_new( + model="gemini-2.5-pro", + messages=[{"role": "user", "content": "What's the weather?"}], + tools=camel_case_tools, + ) + + assert openai_tokens == gemini_tokens == camel_case_tokens + + +def test_token_counter_skips_non_mapping_tools(): + openai_tool: Final = { + "type": "function", + "function": { + "name": "lookup_weather", + "description": "Find current weather conditions for a location", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string", "description": "City and region"}}, + "required": ["location"], + }, + }, + } + messages: Final = [{"role": "user", "content": "What's the weather?"}] + valid_tokens: Final = token_counter_new( + model="gemini-2.5-pro", + messages=messages, + tools=[openai_tool], + ) + mixed_tokens: Final = token_counter_new( + model="gemini-2.5-pro", + messages=messages, + tools=["bad", None, openai_tool], + ) + + assert mixed_tokens == valid_tokens + + class NeedsToleranceUpdateError(Exception): """Custom exception to mark tests that have improved""" diff --git a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 283ed3710d0..67d78d6030e 100644 --- a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1,4 +1,4 @@ -from typing import List +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -1530,7 +1530,7 @@ class TestContextCachingEndpoints: ] all_messages = short_cached_messages + non_cached_messages - large_tools = [ + openai_large_tools: Final = [ { "type": "function", "function": { @@ -1548,6 +1548,11 @@ class TestContextCachingEndpoints: } for i in range(12) ] + large_tools: Final = litellm.utils.get_optional_params( + model="gemini-1.5-pro", + custom_llm_provider="gemini", + tools=openai_large_tools, + )["tools"] optional_params = { **self.sample_optional_params, From f4308bc124eebc783dfc51790ce8db27ed21ae00 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 01:28:02 -0700 Subject: [PATCH 23/65] refactor(types): replace Any with proven types in 5 files (#43304) * refactor(types): replace Any with proven types in 6 files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(types): keep enterprise email import inside try-except for unsafe-import check Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(types): keep email_logging_instance annotation as Any pending a guarded alias Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(types): revert iterator override typing in proxy utils 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/assistants/main.py | 14 +++++++------- litellm/litellm_core_utils/litellm_logging.py | 10 +++++----- litellm/llms/custom_httpx/llm_http_handler.py | 16 +++++++++------- litellm/proxy/common_request_processing.py | 8 +++++--- litellm/utils.py | 2 +- 5 files changed, 27 insertions(+), 23 deletions(-) diff --git a/litellm/assistants/main.py b/litellm/assistants/main.py index 1ce40e94320..c14c4aec093 100644 --- a/litellm/assistants/main.py +++ b/litellm/assistants/main.py @@ -3,7 +3,7 @@ import asyncio import contextvars import os -from collections.abc import Coroutine, Iterable +from collections.abc import Coroutine, Iterable, Mapping, Sequence from functools import partial from typing import Any, Final, Literal @@ -233,8 +233,8 @@ def create_assistants( name: str | None = None, description: str | None = None, instructions: str | None = None, - tools: list[dict[str, Any]] | None = None, - tool_resources: dict[str, Any] | None = None, + tools: Sequence[Mapping[str, object]] | None = None, + tool_resources: Mapping[str, object] | None = None, metadata: dict[str, str] | None = None, temperature: float | None = None, top_p: float | None = None, @@ -244,7 +244,7 @@ def create_assistants( api_base: str | None = None, api_version: str | None = None, **kwargs, -) -> Assistant | Coroutine[Any, Any, Assistant]: +) -> Assistant | Coroutine[None, None, Assistant]: async_create_assistants: Final[bool | None] = kwargs.pop("async_create_assistants", None) if async_create_assistants is not None and not isinstance(async_create_assistants, bool): raise ValueError("Invalid value passed in for async_create_assistants. Only bool or None allowed") @@ -283,7 +283,7 @@ def create_assistants( # only send params that are not None create_assistant_data = {k: v for k, v in create_assistant_data.items() if v is not None} - response: Coroutine[Any, Any, Assistant] | Assistant | None = None + response: Coroutine[None, None, Assistant] | Assistant | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -415,7 +415,7 @@ def delete_assistant( api_base: str | None = None, api_version: str | None = None, **kwargs, -) -> AssistantDeleted | Coroutine[Any, Any, AssistantDeleted]: +) -> AssistantDeleted | Coroutine[None, None, AssistantDeleted]: optional_params: Final = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs) litellm_params_dict: Final = get_litellm_params(**kwargs) @@ -440,7 +440,7 @@ def delete_assistant( elif timeout is None: timeout = 600.0 - response: AssistantDeleted | Coroutine[Any, Any, AssistantDeleted] | None = None + response: AssistantDeleted | Coroutine[None, None, AssistantDeleted] | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 152fd54e55d..5b4187846ff 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -640,8 +640,8 @@ class Logging(LiteLLMLoggingBaseClass): self._own_session_id: str = session_id_var.get() self.function_id = function_id - self.streaming_chunks: list[Any] = [] # for generating complete stream response - self.sync_streaming_chunks: list[Any] = [] # for generating complete stream response + self.streaming_chunks: list[object] = [] # for generating complete stream response + self.sync_streaming_chunks: list[object] = [] # for generating complete stream response self.log_raw_request_response = log_raw_request_response self.raw_request_only = raw_request_only @@ -693,7 +693,7 @@ class Logging(LiteLLMLoggingBaseClass): self.response_timing_metrics: Mapping[str, float] = {} # mutable-ok: kept deep-copyable # Passthrough endpoint guardrails config for field targeting - self.passthrough_guardrails_config: dict[str, Any] | None = None + self.passthrough_guardrails_config: dict[str, object] | None = None self.model_call_details: dict[str, Any] = { "litellm_trace_id": self.litellm_trace_id, @@ -4479,7 +4479,7 @@ def set_callbacks(callback_list, function_id=None): def _init_custom_logger_compatible_class( logging_integration: _custom_logger_compatible_callbacks_literal, internal_usage_cache: DualCache | None, - llm_router: Any | None, # expect litellm.Router, but typing errors due to circular import + llm_router: object, # expect litellm.Router, but typing errors due to circular import custom_logger_init_args: dict | None = {}, ) -> CustomLogger | None: """ @@ -6439,7 +6439,7 @@ def _autorouter_savings_for_payload( def get_standard_logging_object_payload( kwargs: dict | None, - init_response_obj: Any | BaseModel | dict, + init_response_obj: object, start_time: dt_object, end_time: dt_object, logging_obj: Logging, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 8aa38ff3341..ce9f7a2ea54 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6375,7 +6375,7 @@ class BaseLLMHTTPHandler: custom_llm_provider: str | None = None, first_message: str | None = None, request_defaults: ResponsesWebSocketRequestDefaults | None = None, - **kwargs: Any, + **kwargs: object, ) -> Exception | None: """ Handles Responses API WebSocket mode. @@ -10378,13 +10378,14 @@ class BaseLLMHTTPHandler: encoded_vector_store_id: Final = encode_url_path_segment(vector_store_id, field_name="vector_store_id") url: Final = f"{api_base}/{encoded_vector_store_id}" - request_body: Final[dict[str, Any]] = dict(vector_store_update_optional_params) + request_body: Final[dict[str, object]] = dict(vector_store_update_optional_params) + metadata: Final = vector_store_update_optional_params.get("metadata") # Clean metadata to only include string values (OpenAI requirement) - if "metadata" in request_body and request_body["metadata"] is not None: + if metadata is not None: from litellm.utils import add_openai_metadata - request_body["metadata"] = add_openai_metadata(request_body["metadata"]) + request_body["metadata"] = add_openai_metadata(metadata) if extra_body: request_body.update(extra_body) @@ -10456,13 +10457,14 @@ class BaseLLMHTTPHandler: encoded_vector_store_id: Final = encode_url_path_segment(vector_store_id, field_name="vector_store_id") url: Final = f"{api_base}/{encoded_vector_store_id}" - request_body: Final[dict[str, Any]] = dict(vector_store_update_optional_params) + request_body: Final[dict[str, object]] = dict(vector_store_update_optional_params) + metadata: Final = vector_store_update_optional_params.get("metadata") # Clean metadata to only include string values (OpenAI requirement) - if "metadata" in request_body and request_body["metadata"] is not None: + if metadata is not None: from litellm.utils import add_openai_metadata - request_body["metadata"] = add_openai_metadata(request_body["metadata"]) + request_body["metadata"] = add_openai_metadata(metadata) if extra_body: request_body.update(extra_body) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 64b0c6c1967..15610da9aec 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -739,7 +739,7 @@ async def _parse_event_data_for_error(event_line: str | bytes) -> int | None: if not json_str or json_str == "[DONE]": # handle empty data or [DONE] message return None try: - data: Final = orjson.loads(json_str) + data: Final[object] = orjson.loads(json_str) if isinstance(data, dict) and "error" in data and isinstance(data["error"], dict): error_code_raw: Final = data["error"].get("code") error_code: int | None = None @@ -792,7 +792,7 @@ def _extract_error_from_sse_chunk(event_line: str | bytes) -> dict: return default_error try: - data: Final = orjson.loads(json_str) + data: Final[object] = orjson.loads(json_str) if isinstance(data, dict) and "error" in data: error_obj: Final = data["error"] if isinstance(error_obj, dict): @@ -4131,7 +4131,9 @@ class ProxyBaseLLMRequestProcessing: if stripped_ln.startswith("data:"): json_part = stripped_ln.split("data:", 1)[1].strip() if json_part and json_part != "[DONE]": - obj = json.loads(json_part) + obj: object = json.loads(json_part) + if not isinstance(obj, dict): + return None maybe_modified = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict( obj, model_name, litellm_logging_obj ) diff --git a/litellm/utils.py b/litellm/utils.py index 7ce412e818c..d45b29c0f16 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1469,7 +1469,7 @@ async def async_pre_call_deployment_hook(kwargs: dict[str, Any], call_type: str) async def async_post_call_success_deployment_hook( request_data: dict, response: object, call_type: CallTypes | None -) -> Any | None: +) -> object: """ Allow modifying / reviewing the response just after it's received from the deployment. """ From ff462f7a77a5af4da86129265692530dbf04fe69 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 08:56:50 -0700 Subject: [PATCH 24/65] chore(cost-map): update azure_ai/grok-4.6 input price from Azure pricing page (#43440) --- litellm/model_prices_and_context_window_backup.json | 4 ++-- model_prices_and_context_window.json | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 09fc442e5a7..5362b1b042c 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -12257,7 +12257,7 @@ "azure_ai/grok-4.6": { "cache_read_input_token_cost": 5e-07, "cache_read_input_token_cost_above_200k_tokens": 1e-06, - "input_cost_per_token": 2e-06, + "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 4e-06, "litellm_provider": "azure_ai", "max_input_tokens": 200000, @@ -12266,7 +12266,7 @@ "mode": "chat", "output_cost_per_token": 6e-06, "output_cost_per_token_above_200k_tokens": 1.2e-05, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/grok/", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 09fc442e5a7..5362b1b042c 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -12257,7 +12257,7 @@ "azure_ai/grok-4.6": { "cache_read_input_token_cost": 5e-07, "cache_read_input_token_cost_above_200k_tokens": 1e-06, - "input_cost_per_token": 2e-06, + "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 4e-06, "litellm_provider": "azure_ai", "max_input_tokens": 200000, @@ -12266,7 +12266,7 @@ "mode": "chat", "output_cost_per_token": 6e-06, "output_cost_per_token_above_200k_tokens": 1.2e-05, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/grok/", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, From 22b36cbcf6583e2d6b552cc0e87ae6ab82c46341 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 09:14:33 -0700 Subject: [PATCH 25/65] chore(cost-map): update azure_ai/grok-4.6 input price and add azure_ai/MAI-Cyber-1-Flash (#43446) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- .../model_prices_and_context_window_backup.json | 15 +++++++++++++++ model_prices_and_context_window.json | 15 +++++++++++++++ 2 files changed, 30 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 5362b1b042c..b36d84ea027 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -77895,5 +77895,20 @@ "source": "https://aws.amazon.com/bedrock/pricing/", "supports_function_calling": true, "supports_tool_choice": false + }, + "azure_ai/MAI-Cyber-1-Flash": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 6e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 256000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 3.5e-06, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/microsoft/", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true } } diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 5362b1b042c..b36d84ea027 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -77895,5 +77895,20 @@ "source": "https://aws.amazon.com/bedrock/pricing/", "supports_function_calling": true, "supports_tool_choice": false + }, + "azure_ai/MAI-Cyber-1-Flash": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 6e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 256000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 3.5e-06, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/microsoft/", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true } } From 268e8bb735b6871bfed8e593be1b0b53e277d949 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 14:53:12 -0700 Subject: [PATCH 26/65] refactor(rust): share anthropic types, request helpers, and streaming contracts across crates (#43426) * refactor(rust): standardize Azure Messages module path * docs(rust): define shared types crate boundaries * refactor(rust): share request helpers and type Anthropic blocks * docs(rust): format shared type invariants as bullets * test(rust): parameterize repeated cases with rstest * refactor(rust): move Responses transform result into llms * fix(anthropic): validate chat and batch responses * docs(rust): clarify API format ownership boundaries * docs: clarify Rust error message construction * refactor(auth): keep shared Rust errors provider-neutral * refactor(rust): separate format contracts from provider policy * fix(rust): type Anthropic chat response text collection * fix(rust): pass audio secret sources through hosts * fix(rust): unblock batch lint and OCR error assertions * test(rust): assert response failures at the adapter boundary * refactor(rust): declare error messages with typed context * wip * fix(rust): adapt Bedrock error details * style(rust): cargo fmt bedrock audio transcription Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): adapt tests and dead code to typed error details Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): keep converse error contracts and read env secrets without litellm Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci(rust): raise the native wheel size gate to 45 MB Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): tolerate missing usage in converse responses on the transcription route Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/scripts/verify_linux_native_wheel.py | 2 +- litellm-rust/AGENTS.md | 2 + litellm-rust/Cargo.lock | 10 +- litellm-rust/crates/auth-azure/src/native.rs | 74 ++-- litellm-rust/crates/auth-azure/src/resolve.rs | 70 +++- litellm-rust/crates/auth-azure/src/types.rs | 7 +- litellm-rust/crates/auth-gcp/Cargo.toml | 3 + litellm-rust/crates/auth-gcp/src/lib.rs | 51 +-- litellm-rust/crates/auth-types/Cargo.toml | 1 + .../crates/auth-types/src/credential.rs | 16 +- litellm-rust/crates/auth-types/src/error.rs | 173 +++++----- litellm-rust/crates/auth-types/src/http.rs | 8 +- litellm-rust/crates/auth-types/src/lib.rs | 2 +- litellm-rust/crates/auth-types/src/policy.rs | 10 +- litellm-rust/crates/auth-types/tests/error.rs | 74 ++++ .../crates/core-utils/src/call_arguments.rs | 34 +- .../crates/core-utils/src/core_helpers.rs | 34 +- .../crates/core-utils/src/serde_compat.rs | 29 +- .../crates/core-utils/src/settings.rs | 16 + .../crates/core-utils/src/url_utils.rs | 24 +- .../crates/core-utils/tests/settings.rs | 33 ++ litellm-rust/crates/core/AGENTS.md | 2 + litellm-rust/crates/core/Cargo.toml | 2 - .../core/src/audio_transcription/handler.rs | 10 +- .../core/src/audio_transcription/mod.rs | 4 +- .../core/src/audio_transcription/prepare.rs | 13 +- .../core/src/audio_transcription/types.rs | 2 + .../core/src/chat_completions/handler.rs | 13 +- .../crates/core/src/chat_completions/mod.rs | 6 +- .../core/src/chat_completions/prepare.rs | 39 ++- .../crates/core/src/chat_completions/types.rs | 2 + litellm-rust/crates/core/src/error.rs | 35 +- .../crates/core/src/messages/AGENTS.md | 7 + .../crates/core/src/messages/common_utils.rs | 4 +- .../crates/core/src/messages/handler.rs | 19 +- .../crates/core/src/messages/prepare.rs | 17 +- .../crates/core/src/messages/route.rs | 2 +- .../crates/core/src/messages/types.rs | 17 +- litellm-rust/crates/core/src/ocr/document.rs | 34 +- litellm-rust/crates/core/src/ocr/route.rs | 2 +- .../crates/core/src/responses/websocket.rs | 82 +---- .../crates/core/tests/audio_transcription.rs | 71 +++- .../crates/core/tests/chat_completions.rs | 169 +++++++++- .../crates/core/tests/messages/host.rs | 6 +- .../crates/core/tests/messages/request.rs | 126 +++++-- .../crates/core/tests/messages/response.rs | 2 +- .../crates/core/tests/messages/secrets.rs | 8 +- .../crates/core/tests/ocr/azure_ai.rs | 4 +- litellm-rust/crates/cost/Cargo.toml | 1 + litellm-rust/crates/cost/tests/calculation.rs | 28 +- .../src/audio_transcription.rs | 1 + .../gateway-inference/src/chat_completions.rs | 1 + litellm-rust/crates/host/src/machine/auth.rs | 2 +- litellm-rust/crates/http/AGENTS.md | 6 + litellm-rust/crates/http/Cargo.toml | 3 + litellm-rust/crates/http/src/lib.rs | 1 + litellm-rust/crates/http/src/media.rs | 48 +-- litellm-rust/crates/http/src/request.rs | 34 +- litellm-rust/crates/http/src/websocket.rs | 61 ++++ litellm-rust/crates/http/tests/request.rs | 69 ++++ litellm-rust/crates/http/tests/websocket.rs | 59 ++++ litellm-rust/crates/llms/AGENTS.md | 22 +- .../crates/llms/src/anthropic/AGENTS.md | 8 + .../src/anthropic/batches/transformation.rs | 64 +++- .../crates/llms/src/anthropic/chat/handler.rs | 23 +- .../llms/src/anthropic/chat/transformation.rs | 97 ++++-- .../crates/llms/src/anthropic/common_utils.rs | 319 +++++++----------- .../llms/src/anthropic/messages/AGENTS.md | 10 +- .../llms/src/anthropic/messages/handler.rs | 27 +- .../crates/llms/src/anthropic/messages/mod.rs | 1 - .../llms/src/anthropic/messages/thinking.rs | 205 +++++------ .../src/anthropic/messages/transformation.rs | 160 ++++----- .../crates/llms/src/azure_ai/anthropic/mod.rs | 1 - .../crates/llms/src/azure_ai/common_utils.rs | 25 ++ .../llms/src/azure_ai/messages/AGENTS.md | 3 + .../messages}/mod.rs | 1 - .../transformation.rs} | 191 ++++------- litellm-rust/crates/llms/src/azure_ai/mod.rs | 3 +- .../llms/src/azure_ai/ocr/common_utils.rs | 7 +- .../llms/src/azure_ai/ocr/transformation.rs | 6 +- .../audio_transcription/transformation.rs | 18 +- litellm-rust/crates/llms/src/base_llm/auth.rs | 45 +-- .../llms/src/base_llm/chat/transformation.rs | 2 + .../llms/src/base_llm/messages/AGENTS.md | 5 + .../llms/src/base_llm/messages/context.rs | 134 ++++++++ .../crates/llms/src/base_llm/messages/mod.rs | 4 + .../src/base_llm/messages/normalization.rs | 50 +++ .../streaming.rs | 47 ++- .../transformation.rs | 50 ++- litellm-rust/crates/llms/src/base_llm/mod.rs | 2 +- .../src/base_llm/responses/transformation.rs | 101 +----- .../src/bedrock/audio_transcription/mod.rs | 105 ++++-- .../bedrock/chat/converse_transformation.rs | 170 ++++++++-- .../llms/src/bedrock/chat/invoke_handler.rs | 66 ++-- .../llms/src/bedrock/messages/AGENTS.md | 3 + .../anthropic_claude3_transformation.rs | 83 ++--- litellm-rust/crates/llms/src/error.rs | 126 ++++++- litellm-rust/crates/llms/src/lib.rs | 2 +- .../src/openai/responses/transformation.rs | 94 +++++- .../src/openai_like/chat/transformation.rs | 4 + .../llms/src/openai_like/common_utils.rs | 2 +- .../llms/src/vertex_ai/ocr/common_utils.rs | 5 +- .../tests/anthropic_chat_transformation.rs | 61 ++-- .../tests/bedrock_converse_transformation.rs | 78 +++-- .../llms/tests/messages_normalization.rs | 48 +++ .../tests/openai_like_chat_transformation.rs | 2 +- .../crates/python-bridge/src/coercion.rs | 49 ++- .../crates/python-bridge/src/credentials.rs | 45 +-- .../crates/python-bridge/src/errors.rs | 4 +- .../src/routes/audio_transcription.rs | 20 +- .../src/routes/chat_completions.rs | 6 + .../python-bridge/src/routes/messages/host.rs | 4 +- .../python-bridge/src/secrets/config.rs | 10 +- .../crates/python-bridge/src/secrets/mod.rs | 8 +- .../src/secret_manager/client.rs | 25 +- .../crates/token-counter-fast/src/error.rs | 16 +- .../crates/token-counter-fast/src/lib.rs | 2 +- .../crates/token-counter-fast/src/tiktoken.rs | 31 +- .../token-counter-huggingface/Cargo.toml | 3 + .../token-counter-huggingface/src/lib.rs | 18 +- .../crates/token-counter-tiktoken/Cargo.toml | 3 + .../crates/token-counter-tiktoken/src/lib.rs | 49 ++- .../token-counter-tiktoken/src/ranks.rs | 19 +- .../crates/token-counter/src/error.rs | 2 +- litellm-rust/crates/token-counter/src/fast.rs | 2 +- .../crates/token-counter/src/tiktoken.rs | 2 +- litellm-rust/crates/types/AGENTS.md | 60 ++++ .../crates/types/src/audio_transcription.rs | 15 + litellm-rust/crates/types/src/lib.rs | 2 + .../anthropic_messages/anthropic_request.rs | 53 ++- .../crates/types/src/messages/AGENTS.md | 5 + litellm-rust/crates/types/src/messages/mod.rs | 1 + .../src/messages/streaming.rs} | 30 +- .../src/responses/streaming_websocket.rs | 79 ++--- .../crates/types/tests/anthropic_request.rs | 49 +++ .../crates/types/tests/messages_streaming.rs | 37 ++ 136 files changed, 3199 insertions(+), 1615 deletions(-) create mode 100644 litellm-rust/crates/auth-types/tests/error.rs create mode 100644 litellm-rust/crates/core-utils/tests/settings.rs create mode 100644 litellm-rust/crates/core/src/messages/AGENTS.md create mode 100644 litellm-rust/crates/http/src/websocket.rs create mode 100644 litellm-rust/crates/http/tests/request.rs create mode 100644 litellm-rust/crates/http/tests/websocket.rs create mode 100644 litellm-rust/crates/llms/src/anthropic/AGENTS.md delete mode 100644 litellm-rust/crates/llms/src/azure_ai/anthropic/mod.rs create mode 100644 litellm-rust/crates/llms/src/azure_ai/common_utils.rs create mode 100644 litellm-rust/crates/llms/src/azure_ai/messages/AGENTS.md rename litellm-rust/crates/llms/src/{base_llm/anthropic_messages => azure_ai/messages}/mod.rs (55%) rename litellm-rust/crates/llms/src/azure_ai/{anthropic/messages_transformation.rs => messages/transformation.rs} (80%) create mode 100644 litellm-rust/crates/llms/src/base_llm/messages/AGENTS.md create mode 100644 litellm-rust/crates/llms/src/base_llm/messages/context.rs create mode 100644 litellm-rust/crates/llms/src/base_llm/messages/mod.rs create mode 100644 litellm-rust/crates/llms/src/base_llm/messages/normalization.rs rename litellm-rust/crates/llms/src/base_llm/{anthropic_messages => messages}/streaming.rs (67%) rename litellm-rust/crates/llms/src/base_llm/{anthropic_messages => messages}/transformation.rs (75%) create mode 100644 litellm-rust/crates/llms/src/bedrock/messages/AGENTS.md create mode 100644 litellm-rust/crates/llms/tests/messages_normalization.rs create mode 100644 litellm-rust/crates/types/AGENTS.md create mode 100644 litellm-rust/crates/types/src/audio_transcription.rs create mode 100644 litellm-rust/crates/types/src/messages/AGENTS.md create mode 100644 litellm-rust/crates/types/src/messages/mod.rs rename litellm-rust/crates/{llms/src/anthropic/messages/streaming_iterator.rs => types/src/messages/streaming.rs} (87%) create mode 100644 litellm-rust/crates/types/tests/anthropic_request.rs create mode 100644 litellm-rust/crates/types/tests/messages_streaming.rs diff --git a/.github/scripts/verify_linux_native_wheel.py b/.github/scripts/verify_linux_native_wheel.py index 6b7fcd57bbc..465918f5a81 100644 --- a/.github/scripts/verify_linux_native_wheel.py +++ b/.github/scripts/verify_linux_native_wheel.py @@ -214,7 +214,7 @@ def main( native_module: Final = load_native_module(native_path) native_module_loads: Final = native_module is not None panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test") - native_size_limit: Final = 40_000_000 + native_size_limit: Final = 45_000_000 native_size_within_limit: Final = native_member.file_size <= native_size_limit validations: Final = ( (f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG), diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index bc6a2552e4c..b1dc35d3698 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -16,7 +16,9 @@ Use [`#[rstest]`](https://docs.rs/rstest/latest/rstest/attr.rstest.html) for new ## Error definitions - A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs` +- Put message templates in the variant's `#[error(...)]` declaration. Callers pass only the small typed arguments needed to fill them, never `Error::Variant(format!(...))` or a preformatted message. Keep the smallest set of neutral variants that callers need to distinguish; different wording or providers do not justify new variants - Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each. A failure mode is something a caller handles differently (phase, status code, retry, a message Python parity pins exactly); failures no caller tells apart share one variant and differ only in its message +- Keep shared error enums minimal and provider-neutral. Provider names, credential types, configuration fields, and setup guidance belong in caller-supplied data, not dedicated variants or hardcoded shared messages. Reuse a variant for the same failure mode across providers, such as `MissingApiBase { provider: "Azure", guidance: "..." }`. An exact parity message does not justify a provider-specific variant when caller-supplied context can preserve it - Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string - Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return - Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index d67623feffd..bee4421f1e9 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2891,6 +2891,7 @@ dependencies = [ "http 1.4.2", "litellm-auth-types", "moka", + "rstest", "serde_json", "sha2 0.10.9", "tokio", @@ -2900,6 +2901,7 @@ dependencies = [ name = "litellm-auth-types" version = "0.1.0" dependencies = [ + "rstest", "serde", "subtle", "thiserror 2.0.19", @@ -3146,8 +3148,6 @@ dependencies = [ "reqwest 0.12.28", "rstest", "rstest_reuse", - "rustls 0.23.42", - "rustls-native-certs", "serde", "serde_json", "sha2 0.10.9", @@ -3194,6 +3194,7 @@ version = "0.1.0" dependencies = [ "criterion", "proptest", + "rstest", ] [[package]] @@ -3305,6 +3306,7 @@ dependencies = [ name = "litellm-http" version = "0.1.0" dependencies = [ + "futures-util", "http 1.4.2", "hyper-util", "litellm-core-utils", @@ -3312,11 +3314,13 @@ dependencies = [ "reqwest 0.12.28", "rstest", "rustls 0.23.42", + "rustls-native-certs", "serde", "serde_json", "tempfile", "thiserror 2.0.19", "tokio", + "tokio-tungstenite", "veil", "webpki-roots", ] @@ -3661,6 +3665,7 @@ dependencies = [ name = "litellm-token-counter-huggingface" version = "0.1.0" dependencies = [ + "rstest", "serde_json", "thiserror 2.0.19", "tokenizers", @@ -3672,6 +3677,7 @@ version = "0.1.0" dependencies = [ "base64 0.22.1", "once_cell", + "rstest", "rustc-hash", "thiserror 2.0.19", "tiktoken-rs", diff --git a/litellm-rust/crates/auth-azure/src/native.rs b/litellm-rust/crates/auth-azure/src/native.rs index d635e559641..64752162384 100644 --- a/litellm-rust/crates/auth-azure/src/native.rs +++ b/litellm-rust/crates/auth-azure/src/native.rs @@ -133,7 +133,7 @@ impl NativeAzureTokenAcquirer { let token = credential .get_token(&[scope.as_str()], None) .await - .map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?; + .map_err(|error| Error::CredentialAcquisition(error.to_string().into()))?; let expires_on = u64::try_from(token.expires_on.unix_timestamp()) .ok() .map(|seconds| UNIX_EPOCH + Duration::from_secs(seconds)); @@ -250,7 +250,12 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> { let Some(authority) = authority else { return Ok(()); }; - let url = url::Url::parse(authority.value()).map_err(|_| Error::InvalidAzureAuthority)?; + let url = url::Url::parse(authority.value()).map_err(|_| { + Error::InvalidConfiguration( + "Azure authority must be an HTTPS origin without credentials, query, or fragment" + .into(), + ) + })?; if url.scheme() != "https" || url.host_str().is_none() || !url.username().is_empty() @@ -259,7 +264,10 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> { || url.fragment().is_some() || !matches!(url.path(), "" | "/") { - return Err(Error::InvalidAzureAuthority); + return Err(Error::InvalidConfiguration( + "Azure authority must be an HTTPS origin without credentials, query, or fragment" + .into(), + )); } Ok(()) } @@ -368,7 +376,9 @@ fn trusted_source(sources: &[InputSource]) -> InputSource { } fn mixed_sources() -> Result { - Err(Error::MixedAzureCredentialSources) + Err(Error::InvalidConfiguration( + "request-controlled Azure auth inputs cannot be combined with host credentials".into(), + )) } fn build_credential( @@ -433,7 +443,12 @@ fn build_credential( NativeAzureRequest::DeveloperTools { .. } => DeveloperToolsCredential::new(None) .map(|credential| credential as Arc), } - .map_err(|error| Error::AzureCredentialInitialization(error.to_string())) + .map_err(|error| { + Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed( + "Azure credential initialization", + error, + )) + }) } fn client_options( @@ -638,7 +653,7 @@ mod tests { assert_eq!(transport.requests.lock().unwrap().len(), 6); } - #[test] + #[rstest::rstest] fn request_authority_requires_request_owned_client_secret_identity() { let error = ValidatedAzureRequest::new(sourced_client_secret( InputSource::Deployment, @@ -647,10 +662,13 @@ mod tests { )) .unwrap_err(); - assert!(matches!( + assert_eq!( error, - litellm_auth_types::Error::MixedAzureCredentialSources - )); + litellm_auth_types::Error::InvalidConfiguration( + "request-controlled Azure auth inputs cannot be combined with host credentials" + .into() + ) + ); } #[test] @@ -665,24 +683,24 @@ mod tests { assert_eq!(request.credential_source(), InputSource::Request); } - #[test] - fn authority_is_restricted_to_an_https_origin() { - for authority in [ - "http://login.example", - "https://user@login.example", - "https://login.example/tenant", - "https://login.example?target=other", - ] { - let error = ValidatedAzureRequest::new(sourced_client_secret( - InputSource::Deployment, - InputSource::Deployment, - authority, - )) - .unwrap_err(); - assert!(matches!( - error, - litellm_auth_types::Error::InvalidAzureAuthority - )); - } + #[rstest::rstest] + #[case::http("http://login.example")] + #[case::userinfo("https://user@login.example")] + #[case::path("https://login.example/tenant")] + #[case::query("https://login.example?target=other")] + fn authority_is_restricted_to_an_https_origin(#[case] authority: &str) { + let error = ValidatedAzureRequest::new(sourced_client_secret( + InputSource::Deployment, + InputSource::Deployment, + authority, + )) + .unwrap_err(); + assert_eq!( + error, + litellm_auth_types::Error::InvalidConfiguration( + "Azure authority must be an HTTPS origin without credentials, query, or fragment" + .into() + ) + ); } } diff --git a/litellm-rust/crates/auth-azure/src/resolve.rs b/litellm-rust/crates/auth-azure/src/resolve.rs index 9a7afe645db..2142e22db50 100644 --- a/litellm-rust/crates/auth-azure/src/resolve.rs +++ b/litellm-rust/crates/auth-azure/src/resolve.rs @@ -91,7 +91,9 @@ impl AzureAuthService { AzureCredentialPlan::Caller(caller) => { let credential = caller.acquire().await?; if credential.secret().expose().is_empty() { - return Err(Error::EmptyAzureToken); + return Err(Error::EmptyCallerCredential( + "Azure AD token provider returned an empty token", + )); } Ok(Some(Sourced::new(credential, InputSource::Deployment))) } @@ -104,7 +106,11 @@ impl AzureAuthService { } => { let assertion = resolve_reference(inputs, env_lookup, reference.value()) .await? - .ok_or(Error::UnresolvedOidcReference)?; + .ok_or_else(|| { + Error::CredentialAcquisition( + "Azure OIDC reference did not resolve to a value".into(), + ) + })?; let request = ValidatedAzureRequest::new(NativeAzureRequest::ClientAssertion { tenant_id, client_id, @@ -167,7 +173,7 @@ pub(crate) fn select_auth_plan( .map(|selector| Sourced::new(selector, value.source())) }) .transpose() - .map_err(|_| Error::InvalidAzureSelector)?; + .map_err(|_| Error::InvalidConfiguration("invalid Azure credential selector".into()))?; let federated_token_file = configured_string( &inputs.federated_token_file, AZURE_FEDERATED_TOKEN_FILE_ENV, @@ -257,7 +263,9 @@ fn select_native_plan( let selection_source = selected.source(); match selected.into_value() { - AzureCredentialType::ClientSecretCredential => Err(Error::MissingClientSecretFields), + AzureCredentialType::ClientSecretCredential => Err(Error::InvalidConfiguration( + "ClientSecretCredential requires tenant_id, client_id, and client_secret".into(), + )), AzureCredentialType::WorkloadIdentityCredential => { Ok(AzureCredentialPlan::Native(ValidatedAzureRequest::new( workload_request(tenant_id, client_id, federated_token_file, scope, authority)?, @@ -341,9 +349,17 @@ fn workload_request( authority: Option>, ) -> Result { Ok(NativeAzureRequest::WorkloadIdentity { - tenant_id: tenant_id.ok_or(Error::MissingWorkloadTenant)?, - client_id: client_id.ok_or(Error::MissingWorkloadClient)?, - token_file_path: token_file_path.ok_or(Error::MissingWorkloadTokenFile)?, + tenant_id: tenant_id.ok_or_else(|| { + Error::InvalidConfiguration("WorkloadIdentityCredential requires tenant_id".into()) + })?, + client_id: client_id.ok_or_else(|| { + Error::InvalidConfiguration("WorkloadIdentityCredential requires client_id".into()) + })?, + token_file_path: token_file_path.ok_or_else(|| { + Error::InvalidConfiguration( + "WorkloadIdentityCredential requires azure_federated_token_file".into(), + ) + })?, scope, authority, }) @@ -394,10 +410,11 @@ async fn resolve_reference( .map_or(CredentialLookup::Missing, CredentialLookup::Found), CredentialRef::None => return Ok(None), CredentialRef::File(_) | CredentialRef::Request(_) | CredentialRef::Host(_) => { - let resolver = inputs - .credential_resolver - .as_ref() - .ok_or(Error::MissingHostResolver)?; + let resolver = inputs.credential_resolver.as_ref().ok_or_else(|| { + Error::InvalidConfiguration( + "credential reference requires a host credential resolver".into(), + ) + })?; resolver.resolve(reference).await? } }; @@ -415,7 +432,9 @@ fn oidc_reference( }; let value = token.value().expose(); if token.source() == InputSource::Request && value.starts_with("oidc/") { - return Err(Error::RequestAzureCredentialReference); + return Err(Error::InvalidConfiguration( + "request-controlled Azure credential references are not allowed".into(), + )); } if let Some(name) = value.strip_prefix("oidc/env/") { return non_empty_reference(name, "OIDC environment reference") @@ -437,14 +456,20 @@ fn oidc_reference( ))); } if value.starts_with("oidc/") { - return Err(Error::UnsupportedOidcReference); + return Err(Error::InvalidConfiguration( + "unsupported OIDC reference".into(), + )); } Ok(None) } fn non_empty_reference(value: &str, kind: &str) -> Result { if value.is_empty() { - return Err(Error::EmptyReference(kind.to_string())); + return Err(Error::InvalidConfiguration( + litellm_auth_types::ErrorDetail::Empty { + subject: kind.into(), + }, + )); } Ok(value.to_string()) } @@ -493,7 +518,7 @@ mod tests { expires_on: None, }) } else { - Err(Error::AzureTokenAcquisition(format!("{kind} failed"))) + Err(Error::CredentialAcquisition(kind.into())) } }) } @@ -602,7 +627,7 @@ mod tests { assert!(error.to_string().contains("unsupported OIDC reference")); } - #[test] + #[rstest::rstest] fn request_oidc_reference_is_rejected_before_lookup() { let params = json!({ "azure_ad_token": "oidc/env/ASSERTION", @@ -624,7 +649,12 @@ mod tests { }) .unwrap_err(); - assert!(matches!(error, Error::RequestAzureCredentialReference)); + assert_eq!( + error, + Error::InvalidConfiguration( + "request-controlled Azure credential references are not allowed".into() + ) + ); } #[tokio::test] @@ -723,6 +753,7 @@ mod tests { assert_eq!(credential.value().secret().expose(), "caller-token"); } + #[rstest::rstest] #[tokio::test] async fn empty_caller_token_is_rejected() { let error = AzureAuthService::default() @@ -730,6 +761,9 @@ mod tests { .await .unwrap_err(); - assert!(matches!(error, Error::EmptyAzureToken)); + assert_eq!( + error, + Error::EmptyCallerCredential("Azure AD token provider returned an empty token") + ); } } diff --git a/litellm-rust/crates/auth-azure/src/types.rs b/litellm-rust/crates/auth-azure/src/types.rs index a3a898f000f..a042937a047 100644 --- a/litellm-rust/crates/auth-azure/src/types.rs +++ b/litellm-rust/crates/auth-azure/src/types.rs @@ -117,7 +117,12 @@ fn string_config( None => Ok(ConfigValue::Absent), Some(Value::Null) => Ok(ConfigValue::ExplicitNone(source)), Some(Value::String(value)) => Ok(ConfigValue::Value(Sourced::new(value.clone(), source))), - Some(_) => Err(Error::InvalidFieldType(name.to_string())), + Some(_) => Err(Error::InvalidConfiguration( + litellm_auth_types::ErrorDetail::InvalidType { + field: name.into(), + expected: "a string or null", + }, + )), } } diff --git a/litellm-rust/crates/auth-gcp/Cargo.toml b/litellm-rust/crates/auth-gcp/Cargo.toml index 0c6258a193c..8a3598234e1 100644 --- a/litellm-rust/crates/auth-gcp/Cargo.toml +++ b/litellm-rust/crates/auth-gcp/Cargo.toml @@ -19,3 +19,6 @@ tokio.workspace = true gcp_auth = "0.12.7" google-cloud-auth = { workspace = true, optional = true } http = { workspace = true, optional = true } + +[dev-dependencies] +rstest.workspace = true diff --git a/litellm-rust/crates/auth-gcp/src/lib.rs b/litellm-rust/crates/auth-gcp/src/lib.rs index 4374dff95aa..97bc2c482c3 100644 --- a/litellm-rust/crates/auth-gcp/src/lib.rs +++ b/litellm-rust/crates/auth-gcp/src/lib.rs @@ -299,7 +299,7 @@ fn validate_request_credentials(configured: &str) -> Result<&str, Error> { .map(str::to_string) }); if token_uri.as_deref() != Some(GOOGLE_OAUTH_TOKEN_ENDPOINT) { - return Err(Error::RequestVertexTokenEndpoint); + return Err(Error::InvalidConfiguration("request-controlled Vertex credentials must use the canonical Google OAuth token endpoint".into())); } Ok(configured) } @@ -376,10 +376,20 @@ fn optional_credentials( .map(SecretValue::new) .map(|value| Sourced::new(value, source)) .map(Some) - .map_err(|error| Error::InvalidFieldType(format!("{}: {error}", names[0]))); + .map_err(|error| { + Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed( + "credential serialization", + error, + )) + }); } Some(_) => { - return Err(Error::InvalidFieldType(names[0].to_string())); + return Err(Error::InvalidConfiguration( + litellm_auth_types::ErrorDetail::InvalidType { + field: names[0].into(), + expected: "a string or null", + }, + )); } } } @@ -397,7 +407,12 @@ fn optional_string(params: &Map, names: &[&str]) -> Result