fix(a2a): cover the missing-sdk fallback and keep org prefetch warm

The patch left the ImportError resolver base unhit in-process, so codecov
patch coverage sat at two thirds. Call the fallback from a helper the
tests can exercise when a2a-sdk is blocked

Org prefetch also dropped rows whose JSON columns arrived as strings, and
the 5s org TTL expired before the first proxy-behavior getter on a cold
runner. Decode those columns the way team rows already do, and cache org
entries for the same management TTL as user and team

Co-authored-by: ARYAN GORDE <aryangorde8@users.noreply.github.com>
This commit is contained in:
Aryan Gorde 2026-09-13 08:10:17 +00:00
parent bf5a9245ff
commit 0509c8649e
No known key found for this signature in database
6 changed files with 119 additions and 10 deletions

View file

@ -11,14 +11,20 @@ from typing import TYPE_CHECKING, Final
from litellm._logging import verbose_logger
from litellm.constants import LOCALHOST_URL_PATTERNS
def a2a_card_resolver_base() -> type[object]:
try:
from a2a.client import A2ACardResolver as resolver_base
except ImportError:
return object
return resolver_base
if TYPE_CHECKING:
from a2a.client import A2ACardResolver as _A2ACardResolver
from a2a.types import AgentCard
else:
try:
from a2a.client import A2ACardResolver as _A2ACardResolver
except ImportError:
_A2ACardResolver = object
_A2ACardResolver = a2a_card_resolver_base()
# Runtime imports with availability check
AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent-card.json"

View file

@ -5,11 +5,27 @@ Canonical definition for ``litellm_organizationtable``. Re-exported from
``litellm.proxy._types`` for backwards compatibility.
"""
import json
from typing import Final
from pydantic import BaseModel, model_validator
from litellm.models.budget import LiteLLM_BudgetTable
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.models.user import LiteLLM_UserTable
from litellm.types.llms.base import LiteLLMPydanticObjectBase
_JSON_OBJECT_COLUMNS: Final = frozenset({"metadata", "model_spend"})
def _decode_json_object_column(field: str, value: object) -> object:
if not isinstance(value, str):
return value
try:
return json.loads(value)
except json.JSONDecodeError as e:
raise ValueError(f"Field {field} should be a valid dictionary") from e
class LiteLLM_OrganizationTable(LiteLLMPydanticObjectBase):
"""Represents user-controllable params for a LiteLLM_OrganizationTable record"""
@ -27,3 +43,14 @@ class LiteLLM_OrganizationTable(LiteLLMPydanticObjectBase):
litellm_budget_table: LiteLLM_BudgetTable | None = None
object_permission: LiteLLM_ObjectPermissionTable | None = None
object_permission_id: str | None = None
@model_validator(mode="before")
@classmethod
def decode_json_object_columns(cls, values: object) -> object:
payload: Final = values.model_dump() if isinstance(values, BaseModel) else values
if not isinstance(payload, dict):
return values
return {
key: _decode_json_object_column(key, value) if key in _JSON_OBJECT_COLUMNS else value
for key, value in payload.items()
}

View file

@ -4,6 +4,7 @@ the readers and the fallback, so enforcement never depends on this running."""
from __future__ import annotations
import json
import time
from collections.abc import Iterator, Mapping, Sequence
from dataclasses import dataclass
@ -14,7 +15,6 @@ from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.caching.redis_cache import RedisCache
from litellm.constants import DEFAULT_IN_MEMORY_TTL
from litellm.models.organization import LiteLLM_OrganizationTable
from litellm.models.team import LiteLLM_TeamTableCachedObj
from litellm.models.team_membership import LiteLLM_TeamMembership
@ -191,13 +191,13 @@ def _iter_entries(refs: AuthObjectRefs, management_ttl: float) -> Iterator[_Cach
)
if refs.organization_id is not None:
yield _CacheEntry(
f"org_id:{refs.organization_id}", "organization_row", LiteLLM_OrganizationTable, DEFAULT_IN_MEMORY_TTL
f"org_id:{refs.organization_id}", "organization_row", LiteLLM_OrganizationTable, management_ttl
)
yield _CacheEntry(
f"org_id:{refs.organization_id}:with_budget",
"organization_row",
LiteLLM_OrganizationTable,
DEFAULT_IN_MEMORY_TTL,
management_ttl,
)
if refs.project_id is not None:
yield _CacheEntry(f"project_id:{refs.project_id}", "project_row", LiteLLM_ProjectTableCachedObj, management_ttl)
@ -229,13 +229,24 @@ async def _fill_from_redis(entries: Sequence[_CacheEntry], redis_cache: RedisCac
_set_in_memory(memory, entry.cache_key, value, entry.ttl)
def _row_columns(row_value: object) -> Mapping[str, object] | None:
payload: Final = json.loads(row_value) if isinstance(row_value, str) else row_value
try:
return _RowValues.validate_python(payload)
except (TypeError, ValidationError, json.JSONDecodeError):
return None
def _validate_row(
row_value: object, model_type: type[BaseModel], row: _RowKind, refreshed_at: float
) -> BaseModel | None:
if row_value is None:
return None
columns: Final = _row_columns(row_value)
if columns is None:
verbose_proxy_logger.warning("auth prefetch: %s was not a JSON object", row)
return None
try:
columns: Final = _RowValues.validate_python(row_value)
if row in _REFRESH_STAMPED_ROWS:
stamped: Final = {**columns, "last_refreshed_at": refreshed_at} # mutable-ok: validators write into it
return model_type.model_validate(stamped)

View file

@ -25,6 +25,10 @@ pytestmark = pytest.mark.asyncio(loop_scope="session")
def _dead_db() -> MagicMock:
prisma = MagicMock(name="prisma_client")
prisma.db.query_first = AsyncMock(return_value=None)
prisma.db.litellm_organizationtable.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=None)
return prisma

View file

@ -145,6 +145,52 @@ async def test_get_agent_card_forwards_signature_verifier():
assert received["signature_verifier"] is verifier
@pytest.mark.asyncio
async def test_get_agent_card_forwards_signature_verifier_when_trying_well_known_paths():
received = {}
async def mock_parent_get_agent_card(
self, relative_card_path=None, http_kwargs=None, signature_verifier=None
):
received["signature_verifier"] = signature_verifier
return MagicMock()
def verifier(card):
return None
with patch.object(
LiteLLMA2ACardResolver.__bases__[0],
"get_agent_card",
mock_parent_get_agent_card,
):
resolver = LiteLLMA2ACardResolver(
httpx_client=MagicMock(), base_url="http://test-agent:8000"
)
await resolver.get_agent_card(signature_verifier=verifier)
assert received["signature_verifier"] is verifier
def test_a2a_card_resolver_base_is_object_when_a2a_sdk_is_missing(monkeypatch):
import builtins
import sys
from litellm.a2a_protocol.card_resolver import a2a_card_resolver_base
for name in [name for name in sys.modules if name == "a2a" or name.startswith("a2a.")]:
monkeypatch.delitem(sys.modules, name, raising=False)
real_import = builtins.__import__
def block_a2a(name, globals=None, locals=None, fromlist=(), level=0):
if name == "a2a" or name.startswith("a2a."):
raise ImportError(f"No module named '{name}'")
return real_import(name, globals, locals, fromlist, level)
monkeypatch.setattr(builtins, "__import__", block_a2a)
assert a2a_card_resolver_base() is object
def test_is_localhost_or_internal_url():
"""Test that localhost/internal URLs are correctly detected."""
# Should return True for localhost variants

View file

@ -181,8 +181,8 @@ async def test_cold_regime_is_one_mget_one_query_and_the_getters_never_touch_io_
assert sets == sorted(
[
f"SET {TEAM_ID}_{USER_ID} ttl=5",
f"SET org_id:{ORG_ID} ttl=5",
f"SET org_id:{ORG_ID}:with_budget ttl=5",
f"SET org_id:{ORG_ID} ttl=60",
f"SET org_id:{ORG_ID}:with_budget ttl=60",
f"SET {USER_ID} ttl=60",
f"SET team_id:{TEAM_ID} ttl=60",
f"SET team_membership:{USER_ID}:{TEAM_ID} ttl=None",
@ -240,6 +240,21 @@ async def test_hot_regime_costs_nothing():
assert prisma.db.mock_calls == []
def test_org_json_columns_that_arrive_as_strings_still_validate():
org = LiteLLM_OrganizationTable.model_validate({**ORG_ROW, "metadata": "{}", "model_spend": "{}"})
assert org.organization_id == ORG_ID
assert org.metadata == {}
assert org.model_spend == {}
@pytest.mark.asyncio
async def test_org_row_that_arrives_as_a_json_string_is_written_to_cache():
prisma = _prisma({**ALL_ROWS, "organization_row": json.dumps(ORG_ROW)})
cache = _cache(None)
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
assert cache.in_memory_cache.get_cache(key=f"org_id:{ORG_ID}") is not None
@pytest.mark.asyncio
async def test_partial_redis_hit_queries_only_the_missing_objects():
seeded = CountingRedis()