mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
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:
parent
bf5a9245ff
commit
0509c8649e
6 changed files with 119 additions and 10 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue