From 0509c8649e7ca98ebe7f48e93aea53aeeb30f952 Mon Sep 17 00:00:00 2001 From: Aryan Gorde Date: Sun, 13 Sep 2026 08:10:17 +0000 Subject: [PATCH] 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 --- litellm/a2a_protocol/card_resolver.py | 14 ++++-- litellm/models/organization.py | 27 +++++++++++ litellm/proxy/auth/auth_object_prefetch.py | 19 ++++++-- .../auth/test_auth_object_prefetch.py | 4 ++ .../a2a_protocol/test_card_resolver.py | 46 +++++++++++++++++++ .../proxy/auth/test_auth_object_prefetch.py | 19 +++++++- 6 files changed, 119 insertions(+), 10 deletions(-) diff --git a/litellm/a2a_protocol/card_resolver.py b/litellm/a2a_protocol/card_resolver.py index 12123618372..0ee51b6a322 100644 --- a/litellm/a2a_protocol/card_resolver.py +++ b/litellm/a2a_protocol/card_resolver.py @@ -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" diff --git a/litellm/models/organization.py b/litellm/models/organization.py index 894c178af0d..331df435957 100644 --- a/litellm/models/organization.py +++ b/litellm/models/organization.py @@ -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() + } diff --git a/litellm/proxy/auth/auth_object_prefetch.py b/litellm/proxy/auth/auth_object_prefetch.py index 52e26e885c9..1e7b7ade9f0 100644 --- a/litellm/proxy/auth/auth_object_prefetch.py +++ b/litellm/proxy/auth/auth_object_prefetch.py @@ -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) diff --git a/tests/proxy_behavior/auth/test_auth_object_prefetch.py b/tests/proxy_behavior/auth/test_auth_object_prefetch.py index e2d947f4284..a7695f526f3 100644 --- a/tests/proxy_behavior/auth/test_auth_object_prefetch.py +++ b/tests/proxy_behavior/auth/test_auth_object_prefetch.py @@ -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 diff --git a/tests/test_litellm/a2a_protocol/test_card_resolver.py b/tests/test_litellm/a2a_protocol/test_card_resolver.py index 0c691db2341..872cf95686e 100644 --- a/tests/test_litellm/a2a_protocol/test_card_resolver.py +++ b/tests/test_litellm/a2a_protocol/test_card_resolver.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py b/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py index 0fd0dda3017..4ac05af720a 100644 --- a/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py +++ b/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py @@ -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()