diff --git a/litellm/identity/cache.py b/litellm/identity/cache.py index 83882bbe6c0..b238d11dc56 100644 --- a/litellm/identity/cache.py +++ b/litellm/identity/cache.py @@ -43,10 +43,6 @@ def org_generation_key(org_id: str) -> str: return f"{IDENTITY_GENERATION_PREFIX}:org:{org_id}" -def _generation_attr_key(scope: str) -> str: - return f"identity_cache_generation_{scope}" - - def _attach_generations(uak: "UserAPIKeyAuth", generations: dict) -> None: """Stash the generation counters this entry was minted under. diff --git a/tests/test_litellm/identity/extractors/test_client.py b/tests/test_litellm/identity/extractors/test_client.py index 4cf55077062..039bbbd4154 100644 --- a/tests/test_litellm/identity/extractors/test_client.py +++ b/tests/test_litellm/identity/extractors/test_client.py @@ -4,7 +4,11 @@ from types import SimpleNamespace sys.path.insert(0, os.path.abspath("../../..")) -from litellm.identity.extractors.client import extract_client_info +from litellm.identity.extractors.client import ( + _direct_client_host, + _split_forwarded_chain, + extract_client_info, +) def _fake_request(headers, client_host=None): @@ -57,3 +61,20 @@ def test_extract_client_info_uses_starlette_headers_case_insensitively(): info = extract_client_info(req, general_settings={}) assert info.forwarded_chain == ["1.2.3.4"] assert info.user_agent == "curl/8" + + +def test_split_forwarded_chain_returns_empty_for_missing_or_non_string(): + assert _split_forwarded_chain(None) == [] + assert _split_forwarded_chain(12345) == [] # type: ignore[arg-type] + assert _split_forwarded_chain("1.2.3.4, 10.0.0.1") == ["1.2.3.4", "10.0.0.1"] + + +def test_direct_client_host_returns_none_without_a_peer(): + assert _direct_client_host(SimpleNamespace(client=None)) is None + assert ( + _direct_client_host(SimpleNamespace(client=SimpleNamespace(host=None))) is None + ) + assert ( + _direct_client_host(SimpleNamespace(client=SimpleNamespace(host="5.6.7.8"))) + == "5.6.7.8" + ) diff --git a/tests/test_litellm/identity/test_cache.py b/tests/test_litellm/identity/test_cache.py index 433e14e3639..c570a9a4914 100644 --- a/tests/test_litellm/identity/test_cache.py +++ b/tests/test_litellm/identity/test_cache.py @@ -8,6 +8,8 @@ import pytest from litellm.caching.dual_cache import DualCache from litellm.identity.cache import ( IdentityCache, + _attach_generations, + _read_generations, identity_cache_key, team_generation_key, user_generation_key, @@ -132,3 +134,53 @@ async def test_snapshot_generations_maps_returned_values_by_scope(): snapshot = await cache._snapshot_generations_for(uak) assert snapshot == {"team": 7, "user": 3} + + +def test_attach_generations_initializes_non_dict_metadata(): + uak = UserAPIKeyAuth(token="hash-x", user_id="u1") + uak.metadata = None # type: ignore[assignment] + + _attach_generations(uak, {"team": 4}) + + assert isinstance(uak.metadata, dict) + assert _read_generations(uak) == {"team": 4} + + +@pytest.mark.asyncio +async def test_is_stale_is_false_when_entry_carries_no_generations(): + fake = _CountingCache() + cache = IdentityCache(dual_cache=fake) # type: ignore[arg-type] + uak = UserAPIKeyAuth(token="hash-x", user_id="u1", team_id="t1") + assert _read_generations(uak) == {} + + assert await cache._is_stale(uak) is False + # No stored generations means there is nothing to compare, so the cache + # layer must not issue a generation read. + assert fake.batch_get_calls == 0 + + +@pytest.mark.asyncio +async def test_snapshot_generations_returns_empty_for_unscoped_entry(): + fake = _CountingCache() + cache = IdentityCache(dual_cache=fake) # type: ignore[arg-type] + uak = UserAPIKeyAuth(token="hash-x") + + assert await cache._snapshot_generations_for(uak) == {} + assert fake.batch_get_calls == 0 + + +def test_get_identity_cache_falls_back_to_proxy_cache_when_arg_omitted(): + import litellm.identity.cache as cache_module + import litellm.proxy.proxy_server as proxy_server + + saved = cache_module._identity_cache + saved_proxy_cache = proxy_server.user_api_key_cache + sentinel = _user_api_key_cache() + cache_module._identity_cache = None + proxy_server.user_api_key_cache = sentinel + try: + built = cache_module.get_identity_cache() + assert built._cache is sentinel + finally: + cache_module._identity_cache = saved + proxy_server.user_api_key_cache = saved_proxy_cache diff --git a/tests/test_litellm/identity/test_observability.py b/tests/test_litellm/identity/test_observability.py index fc873115973..ddfc6bc48e8 100644 --- a/tests/test_litellm/identity/test_observability.py +++ b/tests/test_litellm/identity/test_observability.py @@ -6,7 +6,11 @@ sys.path.insert(0, os.path.abspath("../..")) import pytest from litellm.integrations.otel.model.spans import SpanRole -from litellm.integrations.otel.runtime import traced +from litellm.integrations.otel.runtime import ( + _apply_span_attrs, + _resolve_attrs, + traced, +) @pytest.mark.asyncio @@ -73,3 +77,47 @@ def test_decorator_rejects_sync_functions(): @traced("identity.sync-not-allowed", role=SpanRole.SERVICE) def sync_fn(): return None + + +def test_resolve_attrs_passes_result_positionally_when_builder_lacks_result_param(): + seen = {} + + def builder(value): + seen["value"] = value + return {"identity.kind": "api_key"} + + out = _resolve_attrs(builder, args=("a",), kwargs={"k": "v"}, result="RESULT") + + assert seen["value"] == "RESULT" + assert out == {"identity.kind": "api_key"} + + +def test_resolve_attrs_returns_empty_without_a_builder(): + assert _resolve_attrs(None, args=(), kwargs={}, result="x") == {} + + +class _RecordingSpan: + def __init__(self): + self.attrs = {} + + def set_attribute(self, key, value): + self.attrs[key] = value + + +def test_apply_span_attrs_sets_non_null_values_and_skips_none(): + span = _RecordingSpan() + _apply_span_attrs(span, {"a": 1, "b": None, "c": "x"}) + assert span.attrs == {"a": 1, "c": "x"} + + +def test_apply_span_attrs_is_noop_without_span_or_attrs(): + _apply_span_attrs(None, {"a": 1}) + _apply_span_attrs(_RecordingSpan(), None) + + +def test_apply_span_attrs_swallows_set_attribute_errors(): + class _Boom: + def set_attribute(self, key, value): + raise RuntimeError("span closed") + + _apply_span_attrs(_Boom(), {"a": 1}) diff --git a/tests/test_litellm/identity/test_resolver.py b/tests/test_litellm/identity/test_resolver.py index 4f9cf8b693f..389333ef5f9 100644 --- a/tests/test_litellm/identity/test_resolver.py +++ b/tests/test_litellm/identity/test_resolver.py @@ -74,3 +74,28 @@ async def test_client_info_from_request(): ctx = await resolve_identity(request=req) assert ctx.client.ip == "127.0.0.1" assert ctx.client.user_agent == "curl/8" + + +@pytest.mark.asyncio +async def test_resolve_user_api_key_auth_hashes_token_then_loads_identity(): + from unittest.mock import AsyncMock, patch + + from litellm.identity.extractors.api_key import hash_principal_token + from litellm.identity.resolver import resolve_user_api_key_auth + + sentinel = object() + with patch( + "litellm.identity.store.load_identity", + new=AsyncMock(return_value=sentinel), + ) as mock_load: + result = await resolve_user_api_key_auth( + api_key="sk-resolve-me", + prisma_client="prisma", # type: ignore[arg-type] + identity_cache="idcache", # type: ignore[arg-type] + user_api_key_cache="uakcache", # type: ignore[arg-type] + ) + + assert result is sentinel + kwargs = mock_load.await_args.kwargs + assert kwargs["hashed_token"] == hash_principal_token("sk-resolve-me") + assert kwargs["cache"] == "idcache" diff --git a/tests/test_litellm/identity/test_store.py b/tests/test_litellm/identity/test_store.py index b32f35f91ad..aece8e7713a 100644 --- a/tests/test_litellm/identity/test_store.py +++ b/tests/test_litellm/identity/test_store.py @@ -218,3 +218,134 @@ async def test_hydrated_copy_is_request_scoped(): assert b.parent_otel_span is None assert b.request_route is None + + +@pytest.mark.asyncio +async def test_cache_miss_without_prisma_raises_no_db_connected(): + cache_backend = UserApiKeyCache() + identity_cache = IdentityCache(dual_cache=cache_backend) + + with pytest.raises(Exception) as excinfo: + await load_identity( + hashed_token="hash-no-db", + prisma_client=None, + cache=identity_cache, + user_api_key_cache=cache_backend, + ) + + assert "No DB Connected" in str(excinfo.value) + + +@pytest.mark.asyncio +async def test_cold_path_fetches_object_permission_when_only_id_present(): + cache_backend = UserApiKeyCache() + identity_cache = IdentityCache(dual_cache=cache_backend) + prisma = _stub_prisma_client() + prisma.get_data.return_value = _verification_token_view( + token="hash-op", + user_id=None, + team_id=None, + object_permission_id="op-99", + ) + + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + fetched = LiteLLM_ObjectPermissionTable(object_permission_id="op-99") + with ( + patch("litellm.identity.store._populate_legacy_cache", new=AsyncMock()), + patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + new=AsyncMock(return_value=fetched), + ) as mock_get_perm, + ): + result = await load_identity( + hashed_token="hash-op", + prisma_client=prisma, + cache=identity_cache, + user_api_key_cache=cache_backend, + ) + + assert mock_get_perm.await_count == 1 + assert result.object_permission is not None + assert result.object_permission.object_permission_id == "op-99" + + +@pytest.mark.asyncio +async def test_cold_path_swallows_object_permission_lookup_failure(): + cache_backend = UserApiKeyCache() + identity_cache = IdentityCache(dual_cache=cache_backend) + prisma = _stub_prisma_client() + prisma.get_data.return_value = _verification_token_view( + token="hash-op-err", + user_id=None, + team_id=None, + object_permission_id="op-err", + ) + + with ( + patch("litellm.identity.store._populate_legacy_cache", new=AsyncMock()), + patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + new=AsyncMock(side_effect=RuntimeError("db blip")), + ), + ): + result = await load_identity( + hashed_token="hash-op-err", + prisma_client=prisma, + cache=identity_cache, + user_api_key_cache=cache_backend, + ) + + assert result.object_permission is None + assert result.object_permission_id == "op-err" + + +@pytest.mark.asyncio +async def test_populate_legacy_cache_delegates_to_cache_key_object(): + from litellm.identity.store import _populate_legacy_cache + from litellm.proxy._types import UserAPIKeyAuth + + uak = UserAPIKeyAuth(token="hash-legacy", user_id="u1") + with patch( + "litellm.proxy.auth.auth_checks._cache_key_object", new=AsyncMock() + ) as mock_cache_key: + await _populate_legacy_cache( + hashed_token="hash-legacy", + uak=uak, + user_api_key_cache=UserApiKeyCache(), + proxy_logging_obj=None, + ) + + assert mock_cache_key.await_count == 1 + assert mock_cache_key.await_args.kwargs["hashed_token"] == "hash-legacy" + assert mock_cache_key.await_args.kwargs["user_api_key_obj"] is uak + + +@pytest.mark.asyncio +async def test_populate_legacy_cache_swallows_write_failures(): + from litellm.identity.store import _populate_legacy_cache + from litellm.proxy._types import UserAPIKeyAuth + + uak = UserAPIKeyAuth(token="hash-legacy", user_id="u1") + with patch( + "litellm.proxy.auth.auth_checks._cache_key_object", + new=AsyncMock(side_effect=RuntimeError("redis down")), + ): + await _populate_legacy_cache( + hashed_token="hash-legacy", + uak=uak, + user_api_key_cache=UserApiKeyCache(), + proxy_logging_obj=None, + ) + + +def test_rehydrate_bundled_user_nulls_out_uncoercible_dict(): + from litellm.identity.store import _rehydrate_bundled_user + from litellm.proxy._types import UserAPIKeyAuth + + uak = UserAPIKeyAuth(token="hash-bad-user") + uak.user = {"user_id": "u1", "tpm_limit": "not-an-int"} + + _rehydrate_bundled_user(uak) + + assert uak.user is None