mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(anthropic): stop model discovery crashing on the second page
The paginated /v1/models fetch handed the shared http client a read-only mapping for its query params. That client merges the URL's own query string in by mutating the mapping it is given, and a mappingproxy has no update, so the call raised AttributeError as soon as after_id was set. Any org with more models than fit on one page could not discover them at all. after_id now rides the URL instead, which is the same mechanism the client uses to carry query params, so nothing needs to mutate. The existing pagination test missed this because its stub client only recorded params and never merged them the way the real one does. The new test drives the real HTTPHandler over a mock transport, so the second page goes through the code that actually mutates. It fails with the original AttributeError against the previous implementation. Also clears the basedpyright reportArgumentType errors this branch added: make_spec now takes typed keyword arguments rather than splatting an untyped dict, both token_response helpers build their body in one shot instead of assigning an int into a str-inferred dict, the duplicated credential decrypt block is one helper, and HTTPHandler.get accepts a Mapping for headers, which it only forwards.
This commit is contained in:
parent
743b918349
commit
0a923aa10b
6 changed files with 102 additions and 48 deletions
|
|
@ -8,6 +8,7 @@ from collections.abc import Mapping, Sequence
|
|||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import Any, ClassVar, Final, Literal
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
|
|
@ -205,10 +206,12 @@ def _sanitized_anthropic_error(response: httpx.Response, detail: str | None = No
|
|||
def _fetch_anthropic_models_page(
|
||||
api_base: str, headers: Mapping[str, str], after_id: str | None
|
||||
) -> _AnthropicModelsPage:
|
||||
# after_id rides the URL because the client mutates the params mapping it is handed,
|
||||
# which a read-only one cannot support
|
||||
query: Final = f"?after_id={quote(after_id)}" if after_id else ""
|
||||
response: Final = litellm.module_level_client.get(
|
||||
url=f"{api_base}/v1/models",
|
||||
url=f"{api_base}/v1/models{query}",
|
||||
headers=headers,
|
||||
params=MappingProxyType({"after_id": after_id}) if after_id else MappingProxyType({}),
|
||||
follow_redirects=False,
|
||||
)
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -1257,7 +1257,7 @@ class HTTPHandler:
|
|||
self,
|
||||
url: str,
|
||||
params: dict | None = None,
|
||||
headers: dict | None = None,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
follow_redirects: bool | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -23,6 +23,21 @@ from litellm.types.router import (
|
|||
from litellm.types.utils import CredentialItem
|
||||
|
||||
|
||||
def _decrypted(db_credential: CredentialItem) -> CredentialItem:
|
||||
"""The stored credential with every value decrypted, leaving already-plaintext values alone."""
|
||||
decrypted_values: Final = MappingProxyType(
|
||||
{
|
||||
key: decrypt_value_helper(value=value, key=key) or value
|
||||
for key, value in db_credential.credential_values.items()
|
||||
}
|
||||
)
|
||||
return CredentialItem(
|
||||
credential_name=db_credential.credential_name,
|
||||
credential_values=decrypted_values, # pyright: ignore[reportArgumentType] # declared dict[str, str], and pydantic copies this mapping into one on validation; LIT002 rules out building that dict here
|
||||
credential_info=db_credential.credential_info,
|
||||
)
|
||||
|
||||
|
||||
async def hydrate_named_credential_authoritative(
|
||||
credential_name: str,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -39,17 +54,7 @@ async def hydrate_named_credential_authoritative(
|
|||
db_credential: Final = await CredentialsRepository(prisma_client).find_by_name(credential_name)
|
||||
if db_credential is None:
|
||||
return await hydrate_named_credential(credential_name, prisma_client)
|
||||
decrypted_values: Final = MappingProxyType(
|
||||
{
|
||||
key: decrypt_value_helper(value=value, key=key) or value
|
||||
for key, value in db_credential.credential_values.items()
|
||||
}
|
||||
)
|
||||
return CredentialItem(
|
||||
credential_name=db_credential.credential_name,
|
||||
credential_values=decrypted_values,
|
||||
credential_info=db_credential.credential_info,
|
||||
)
|
||||
return _decrypted(db_credential)
|
||||
|
||||
|
||||
async def hydrate_named_credential(
|
||||
|
|
@ -64,17 +69,7 @@ async def hydrate_named_credential(
|
|||
db_credential: Final = await CredentialsRepository(prisma_client).find_by_name(credential_name)
|
||||
if db_credential is None:
|
||||
return None
|
||||
decrypted_values: Final = MappingProxyType(
|
||||
{
|
||||
key: decrypt_value_helper(value=value, key=key) or value
|
||||
for key, value in db_credential.credential_values.items()
|
||||
}
|
||||
)
|
||||
return CredentialItem(
|
||||
credential_name=db_credential.credential_name,
|
||||
credential_values=decrypted_values,
|
||||
credential_info=db_credential.credential_info,
|
||||
)
|
||||
return _decrypted(db_credential)
|
||||
|
||||
|
||||
async def named_credential_wif_fields(
|
||||
|
|
|
|||
|
|
@ -15,8 +15,10 @@ import os
|
|||
import sys
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../..")))
|
||||
|
|
@ -3181,8 +3183,42 @@ class TestModelDiscovery:
|
|||
|
||||
assert models == ["anthropic/claude-a", "anthropic/claude-b", "anthropic/claude-c"]
|
||||
assert len(client.calls) == 2
|
||||
assert client.calls[0].params == {}
|
||||
assert client.calls[1].params == {"after_id": "claude-b"}
|
||||
assert client.calls[0].url == "https://api.anthropic.com/v1/models"
|
||||
assert client.calls[1].url == "https://api.anthropic.com/v1/models?after_id=claude-b"
|
||||
|
||||
def test_paginated_fetch_survives_the_real_http_client(self, monkeypatch, clean_anthropic_env):
|
||||
"""Regression: the second page is fetched through the real HTTPHandler, which merges the
|
||||
URL's query string into the params mapping by mutating it. Handing that client a
|
||||
read-only mapping raised AttributeError, so discovery blew up for any org holding more
|
||||
models than one page, while the stubbed client here never exercised the mutation."""
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY)
|
||||
requested: Final = [] # mutable-ok: a test spy recording the URLs the client was asked for
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
requested.append(str(request.url))
|
||||
first: Final = "after_id" not in request.url.params
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"data": [{"id": "claude-a"}] if first else [{"id": "claude-b"}],
|
||||
"has_more": first,
|
||||
"last_id": "claude-a" if first else "claude-b",
|
||||
},
|
||||
)
|
||||
|
||||
handler: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond)))
|
||||
monkeypatch.setattr("litellm.module_level_client", handler)
|
||||
|
||||
models = AnthropicModelInfo().get_models(api_base="https://api.anthropic.com")
|
||||
|
||||
assert models == ["anthropic/claude-a", "anthropic/claude-b"]
|
||||
assert requested == [
|
||||
"https://api.anthropic.com/v1/models",
|
||||
"https://api.anthropic.com/v1/models?after_id=claude-a",
|
||||
]
|
||||
|
||||
def test_get_models_refuses_to_follow_redirects(self, monkeypatch, clean_anthropic_env):
|
||||
"""Only the configured api_base is validated, so a redirected /v1/models must not be
|
||||
|
|
|
|||
|
|
@ -102,9 +102,11 @@ class ManualExecutor(concurrent.futures.Executor):
|
|||
|
||||
|
||||
def token_response(token: str = "sk-ant-oat01-minted", expires_in: int | None = 3600) -> httpx.Response:
|
||||
body = {"access_token": token, "token_type": "Bearer"}
|
||||
if expires_in is not None:
|
||||
body["expires_in"] = expires_in
|
||||
body: Final[dict[str, str | int]] = {
|
||||
"access_token": token,
|
||||
"token_type": "Bearer",
|
||||
**({} if expires_in is None else {"expires_in": expires_in}),
|
||||
}
|
||||
return httpx.Response(200, json=body)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ import logging
|
|||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from urllib.parse import parse_qsl
|
||||
|
||||
|
|
@ -33,7 +34,9 @@ from litellm.llms.base_llm.auth.token_exchange import (
|
|||
redact_oauth_error_body,
|
||||
)
|
||||
from litellm.llms.base_llm.auth.types import (
|
||||
AssertionSource,
|
||||
AssertionSourceError,
|
||||
BodyEncoding,
|
||||
ExchangeError,
|
||||
ExchangeResult,
|
||||
InsecureTokenUrl,
|
||||
|
|
@ -143,29 +146,45 @@ class NeverRunsExecutor(concurrent.futures.Executor):
|
|||
|
||||
|
||||
def token_response(token: str = "sk-ant-oat01-minted", expires_in: int | None = 3600) -> httpx.Response:
|
||||
body = {"access_token": token, "token_type": "Bearer"}
|
||||
if expires_in is not None:
|
||||
body["expires_in"] = expires_in
|
||||
body: Final[dict[str, str | int]] = {
|
||||
"access_token": token,
|
||||
"token_type": "Bearer",
|
||||
**({} if expires_in is None else {"expires_in": expires_in}),
|
||||
}
|
||||
return httpx.Response(200, json=body)
|
||||
|
||||
|
||||
def make_spec(**overrides) -> TokenExchangeSpec:
|
||||
base = {
|
||||
"token_url": "https://token.example/v1/oauth/token",
|
||||
"assertion_ref": DEFAULT_REF,
|
||||
"assertion_field": "assertion",
|
||||
"static_body": {
|
||||
def make_spec(
|
||||
*,
|
||||
token_url: str = "https://token.example/v1/oauth/token",
|
||||
assertion_ref: str = DEFAULT_REF,
|
||||
assertion_field: str = "assertion",
|
||||
static_body: Mapping[str, str] = MappingProxyType(
|
||||
{
|
||||
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
|
||||
"federation_rule_id": "fdrl_1",
|
||||
"organization_id": "org-1",
|
||||
},
|
||||
"body_encoding": "json",
|
||||
"request_headers": {"anthropic-beta": "oauth-2025-04-20,oidc-federation-2026-04-01"},
|
||||
"cache_key_identity": ("fdrl_1", "org-1", "", ""),
|
||||
"timeout_seconds": 2.0,
|
||||
}
|
||||
base.update(overrides)
|
||||
return TokenExchangeSpec(**base)
|
||||
}
|
||||
),
|
||||
body_encoding: BodyEncoding = "json",
|
||||
request_headers: Mapping[str, str] = MappingProxyType(
|
||||
{"anthropic-beta": "oauth-2025-04-20,oidc-federation-2026-04-01"}
|
||||
),
|
||||
cache_key_identity: tuple[str, ...] = ("fdrl_1", "org-1", "", ""),
|
||||
timeout_seconds: float = 2.0,
|
||||
assertion_source: AssertionSource | None = None,
|
||||
) -> TokenExchangeSpec:
|
||||
return TokenExchangeSpec(
|
||||
token_url=token_url,
|
||||
assertion_ref=assertion_ref,
|
||||
assertion_field=assertion_field,
|
||||
static_body=static_body,
|
||||
body_encoding=body_encoding,
|
||||
request_headers=request_headers,
|
||||
cache_key_identity=cache_key_identity,
|
||||
timeout_seconds=timeout_seconds,
|
||||
assertion_source=assertion_source,
|
||||
)
|
||||
|
||||
|
||||
class RecordingMetricsSink:
|
||||
|
|
@ -862,7 +881,6 @@ class TestAssertionGuards:
|
|||
|
||||
assert isinstance(result, AssertionSourceError)
|
||||
assert result.detail == overlong_message[:_REDACTION_CAP]
|
||||
assert len(result.detail) == _REDACTION_CAP
|
||||
|
||||
|
||||
class TestAssertionSourceOverridesEngineReader:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue