From 53d2ab6b0527ad8c376d236c0a8e462a0162e5d0 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 09:20:27 -0700 Subject: [PATCH] fix(proxy): emit postgres service spans only on real DB reads in auth cache helpers (#44148) log_db_metrics wrapped whole cache-first auth helpers and always emitted a ServiceTypes.DB success event, so in-memory cache hits showed up as postgres spans and DB service metrics. The decorator now installs a ContextVar witness that _TrackedPrismaEngine marks on every Prisma query and transaction call, and the DB event is emitted only when the witness was marked. Real reads keep their existing call_type names, the failure path and the PROXY batch-write branch are unchanged, and Redis instrumentation is untouched. A decorated helper that reaches Prisma only through another decorated helper (get_key_object -> get_object_permission, get_team_object_by_alias -> get_object_permission, get_tag_object -> get_tag_objects_batch) used to emit two events for one query. The inner wrapper now marks its witness as reported when it emits a success or DB failure event, and only unreported activity is handed up to the enclosing witness, so the inner event is the one that survives. An outer helper that also queries Prisma directly or through undecorated callees still gets its own event. Tests: get_user_object and get_org_object cache hits emit no DB event; a get_user_object miss through the generated Prisma client emits exactly one postgres get_user_object event; decorator-level tests cover nested calls emitting only the inner event, outer calls with their own query, inner non-DB failures, bounded lookups and sibling-request isolation. Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/log_db_metrics.py | 53 +++++- litellm/proxy/db/prisma_client.py | 5 + .../test_log_db_redis_services.py | 10 ++ tests/unit/proxy/auth/test_auth_checks.py | 106 +++++++++++- tests/unit/proxy/db/test_log_db_metrics.py | 152 ++++++++++++++++++ 5 files changed, 321 insertions(+), 5 deletions(-) create mode 100644 tests/unit/proxy/db/test_log_db_metrics.py diff --git a/litellm/proxy/db/log_db_metrics.py b/litellm/proxy/db/log_db_metrics.py index 559caddbca0..8c3d757838d 100644 --- a/litellm/proxy/db/log_db_metrics.py +++ b/litellm/proxy/db/log_db_metrics.py @@ -6,6 +6,7 @@ ServiceLogger() then sends DB logs to Prometheus, OTEL, Datadog etc import asyncio from collections.abc import Callable +from contextvars import ContextVar from datetime import datetime from functools import wraps from typing import Final @@ -25,6 +26,40 @@ def _safe_db_event_metadata(kwargs: dict) -> dict[str, str] | None: return {"table_name": table_name} if isinstance(table_name, str) else None +class _DbIoWitness: + """Activity an inner decorated call already reported stays with it; only unreported activity reaches the enclosing call.""" + + __slots__ = ("_parent", "_reported", "_touched") + + def __init__(self, parent: "_DbIoWitness | None") -> None: + self._parent: Final = parent + self._touched = False + self._reported = False + + @property + def touched(self) -> bool: + return self._touched + + def mark(self) -> None: + self._touched = True + + def report(self) -> None: + self._reported = True + + def close(self) -> None: + if self._touched and not self._reported and self._parent is not None: + self._parent.mark() + + +_db_io_witness: Final[ContextVar["_DbIoWitness | None"]] = ContextVar("litellm_db_io_witness", default=None) + + +def record_db_io() -> None: + witness: Final = _db_io_witness.get() + if witness is not None: + witness.mark() + + def log_db_metrics(func): """ Decorator to log the duration of a DB related function to ServiceLogger() @@ -46,6 +81,8 @@ def log_db_metrics(func): @wraps(func) async def wrapper(*args, **kwargs): start_time: Final[datetime] = datetime.now() + witness: Final = _DbIoWitness(parent=_db_io_witness.get()) + witness_token: Final = _db_io_witness.set(witness) try: result: Final = await func(*args, **kwargs) @@ -53,6 +90,8 @@ def log_db_metrics(func): from litellm.proxy.proxy_server import proxy_logging_obj if "PROXY" not in func.__name__: + if not witness.touched: + return result asyncio.create_task( proxy_logging_obj.service_logging_obj.async_service_success_hook( service=ServiceTypes.DB, @@ -64,6 +103,7 @@ def log_db_metrics(func): event_metadata=_safe_db_event_metadata(kwargs), ) ) + witness.report() elif ( # in litellm custom callbacks kwargs is passed as arg[0] # https://docs.litellm.ai/docs/observability/custom_callback#callback-functions @@ -90,15 +130,19 @@ def log_db_metrics(func): return result except Exception as e: end_time: datetime = datetime.now() - await _handle_logging_db_exception( + if await _handle_logging_db_exception( e=e, func=func, kwargs=kwargs, args=args, start_time=start_time, end_time=end_time, - ) + ): + witness.report() raise e + finally: + witness.close() + _db_io_witness.reset(witness_token) return wrapper @@ -121,12 +165,12 @@ async def _handle_logging_db_exception( args: tuple, start_time: datetime, end_time: datetime, -) -> None: +) -> bool: from litellm.proxy.proxy_server import proxy_logging_obj # don't log this as a DB Service Failure, if the DB did not raise an exception if _is_exception_related_to_db(e) is not True: - return + return False await proxy_logging_obj.service_logging_obj.async_service_failure_hook( error=e, @@ -138,3 +182,4 @@ async def _handle_logging_db_exception( end_time=end_time, event_metadata=_safe_db_event_metadata(kwargs), ) + return True diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index e7c7102c98f..28e03d6dbb5 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -17,6 +17,7 @@ from typing import TYPE_CHECKING, Any, Final, Protocol from litellm._logging import verbose_proxy_logger from litellm.proxy.db.db_url_settings import add_missing_query_params, token_refresh_params_from_url +from litellm.proxy.db.log_db_metrics import record_db_io from litellm.proxy.db.token_auth import ( DEFAULT_POSTGRES_PORT, DatabaseTokenAuth, @@ -106,6 +107,7 @@ class _TrackedPrismaEngine: return getattr(self._engine, name) async def query(self, content: str, *, tx_id: str | None) -> object: + record_db_io() self.tracker.begin_operation() try: return await self._engine.query(content, tx_id=tx_id) @@ -113,6 +115,7 @@ class _TrackedPrismaEngine: self.tracker.end_operation() async def start_transaction(self, *, content: str) -> str: + record_db_io() self.tracker.begin_operation() try: transaction_id: Final = await self._engine.start_transaction(content=content) @@ -123,6 +126,7 @@ class _TrackedPrismaEngine: return transaction_id async def commit_transaction(self, tx_id: str) -> None: + record_db_io() self.tracker.begin_operation() try: await self._engine.commit_transaction(tx_id) @@ -131,6 +135,7 @@ class _TrackedPrismaEngine: self.tracker.transaction_finished(tx_id) async def rollback_transaction(self, tx_id: str) -> None: + record_db_io() self.tracker.begin_operation() try: await self._engine.rollback_transaction(tx_id) diff --git a/tests/logging_callback_tests/test_log_db_redis_services.py b/tests/logging_callback_tests/test_log_db_redis_services.py index e3bc8383c46..ba7b333e097 100644 --- a/tests/logging_callback_tests/test_log_db_redis_services.py +++ b/tests/logging_callback_tests/test_log_db_redis_services.py @@ -14,14 +14,22 @@ import litellm from litellm import completion from litellm._logging import verbose_logger from litellm.proxy.utils import log_db_metrics, ServiceTypes +from litellm.proxy.db.prisma_client import _PrismaDrainTracker, _TrackedPrismaEngine from datetime import datetime +from types import SimpleNamespace import httpx from prisma.errors import ClientNotConnectedError +async def _run_prisma_query() -> None: + engine = _TrackedPrismaEngine(SimpleNamespace(query=AsyncMock(return_value={})), _PrismaDrainTracker()) + await engine.query("{}", tx_id=None) + + # Test async function to decorate @log_db_metrics async def sample_db_function(*args, **kwargs): + await _run_prisma_query() return "success" @@ -71,6 +79,7 @@ async def test_log_db_metrics_event_metadata_is_safe(): @log_db_metrics async def db_call(**kwargs): + await _run_prisma_query() return "success" await db_call( @@ -99,6 +108,7 @@ async def test_log_db_metrics_duration(): # Add a delay to the function to test duration @log_db_metrics async def delayed_function(**kwargs): + await _run_prisma_query() await asyncio.sleep(1) # 1 second delay return "success" diff --git a/tests/unit/proxy/auth/test_auth_checks.py b/tests/unit/proxy/auth/test_auth_checks.py index 2538556d3b5..448211978d1 100644 --- a/tests/unit/proxy/auth/test_auth_checks.py +++ b/tests/unit/proxy/auth/test_auth_checks.py @@ -7,9 +7,18 @@ from dotenv import load_dotenv load_dotenv() +from collections.abc import Iterator +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, MagicMock, patch + import pytest, litellm import httpx -from litellm.proxy._types import UserAPIKeyAuth +from prisma import Prisma +from litellm._service_logger import ServiceTypes +from litellm.proxy._types import LiteLLM_OrganizationTable, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import get_org_object, get_user_object +from litellm.proxy.db.prisma_client import PrismaWrapper from litellm.proxy.auth.auth_checks import get_end_user_object from litellm.caching.caching import DualCache from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -1491,3 +1500,98 @@ async def test_key_access_group_grants_model_when_get_access_object_raises(): finally: for p in patches: p.stop() + + +@pytest.fixture +def db_success_hook() -> Iterator[AsyncMock]: + hook: Final = AsyncMock() + with patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + MagicMock(service_logging_obj=MagicMock(async_service_success_hook=hook)), + ): + yield hook + + +async def _db_service_call_types(hook: AsyncMock) -> tuple[str, ...]: + await asyncio.sleep(0) + return tuple(call.kwargs["call_type"] for call in hook.await_args_list if call.kwargs["service"] == ServiceTypes.DB) + + +def _prisma_client_serving(user_id: str) -> SimpleNamespace: + row: Final = { + "user_id": user_id, + "user_role": "internal_user", + "teams": [], + "spend": 0.0, + "models": [], + "metadata": "{}", + "allowed_cache_controls": [], + "policies": [], + "model_spend": "{}", + "model_max_budget": "{}", + "organization_memberships": [], + } + engine: Final = SimpleNamespace(query=AsyncMock(return_value={"data": {"result": row}}), stop=lambda: None) + generated_client: Final = Prisma() + generated_client._engine = engine + return SimpleNamespace(db=PrismaWrapper(original_prisma=generated_client, iam_token_db_auth=False)) + + +@pytest.mark.asyncio +async def test_get_user_object_cache_hit_emits_no_postgres_service_event(db_success_hook: AsyncMock) -> None: + user_id: Final = f"cached-user-{uuid.uuid4()}" + cache: Final = UserApiKeyCache() + await cache.async_set_cache(key=user_id, value=LiteLLM_UserTable(user_id=user_id, user_role="internal_user")) + + result: Final = await get_user_object( + user_id=user_id, + prisma_client=MagicMock(), + user_api_key_cache=cache, + user_id_upsert=False, + parent_otel_span="auth-span", + ) + + assert result is not None and result.user_id == user_id + assert await _db_service_call_types(db_success_hook) == () + + +@pytest.mark.asyncio +async def test_get_org_object_cache_hit_emits_no_postgres_service_event(db_success_hook: AsyncMock) -> None: + org_id: Final = f"cached-org-{uuid.uuid4()}" + cache: Final = UserApiKeyCache() + await cache.async_set_cache( + key=f"org_id:{org_id}", + value=LiteLLM_OrganizationTable( + organization_id=org_id, budget_id="b", models=[], created_by="t", updated_by="t" + ), + ) + + result: Final = await get_org_object( + org_id=org_id, + prisma_client=MagicMock(), + user_api_key_cache=cache, + parent_otel_span="auth-span", + ) + + assert result is not None and result.organization_id == org_id + assert await _db_service_call_types(db_success_hook) == () + + +@pytest.mark.asyncio +async def test_get_user_object_cache_miss_emits_exactly_one_postgres_get_user_object_event( + db_success_hook: AsyncMock, +) -> None: + user_id: Final = f"db-user-{uuid.uuid4()}" + prisma_client: Final = _prisma_client_serving(user_id) + + result: Final = await get_user_object( + user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=UserApiKeyCache(), + user_id_upsert=False, + parent_otel_span="auth-span", + ) + + assert result is not None and result.user_id == user_id + assert await _db_service_call_types(db_success_hook) == ("get_user_object",) + assert db_success_hook.await_args_list[0].kwargs["parent_otel_span"] == "auth-span" diff --git a/tests/unit/proxy/db/test_log_db_metrics.py b/tests/unit/proxy/db/test_log_db_metrics.py new file mode 100644 index 00000000000..1ec1afe6106 --- /dev/null +++ b/tests/unit/proxy/db/test_log_db_metrics.py @@ -0,0 +1,152 @@ +import asyncio +from collections.abc import Iterator +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm._service_logger import ServiceTypes +from litellm.proxy.db.db_lookup_gate import bounded_db_lookup +from litellm.proxy.db.log_db_metrics import log_db_metrics +from litellm.proxy.db.prisma_client import _PrismaDrainTracker, _TrackedPrismaEngine + + +def _tracked_engine() -> _TrackedPrismaEngine: + raw_engine: Final = SimpleNamespace(query=AsyncMock(return_value={"data": {}})) + return _TrackedPrismaEngine(raw_engine, _PrismaDrainTracker()) + + +@pytest.fixture +def success_hook() -> Iterator[AsyncMock]: + hook: Final = AsyncMock() + with patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + MagicMock(service_logging_obj=MagicMock(async_service_success_hook=hook)), + ): + yield hook + + +async def _db_call_types(hook: AsyncMock) -> tuple[str, ...]: + await asyncio.sleep(0) + return tuple(call.kwargs["call_type"] for call in hook.await_args_list if call.kwargs["service"] == ServiceTypes.DB) + + +@pytest.mark.asyncio +async def test_a_decorated_call_that_never_queries_the_engine_emits_no_db_event(success_hook: AsyncMock) -> None: + @log_db_metrics + async def cache_hit(**kwargs: object) -> str: + return "cached" + + assert await cache_hit(parent_otel_span="span") == "cached" + assert await _db_call_types(success_hook) == () + + +@pytest.mark.asyncio +async def test_a_decorated_call_that_queries_the_engine_emits_one_db_event_named_after_it( + success_hook: AsyncMock, +) -> None: + engine: Final = _tracked_engine() + + @log_db_metrics + async def read_user_row(**kwargs: object) -> object: + return await engine.query("{}", tx_id=None) + + await read_user_row(parent_otel_span="span", table_name="LiteLLM_UserTable") + + assert await _db_call_types(success_hook) == ("read_user_row",) + event: Final = success_hook.await_args_list[0].kwargs + assert (event["parent_otel_span"], event["event_metadata"]) == ("span", {"table_name": "LiteLLM_UserTable"}) + + +@pytest.mark.asyncio +async def test_a_query_behind_the_bounded_lookup_task_still_counts_for_the_enclosing_call( + success_hook: AsyncMock, +) -> None: + engine: Final = _tracked_engine() + + @log_db_metrics + async def read_through_gate(**kwargs: object) -> object: + return await bounded_db_lookup(engine.query("{}", tx_id=None), name="user") + + await read_through_gate() + + assert await _db_call_types(success_hook) == ("read_through_gate",) + + +@pytest.mark.asyncio +async def test_one_query_inside_a_nested_decorated_call_emits_only_the_inner_event(success_hook: AsyncMock) -> None: + engine: Final = _tracked_engine() + + @log_db_metrics + async def get_data(**kwargs: object) -> object: + return await engine.query("{}", tx_id=None) + + @log_db_metrics + async def get_key_object(**kwargs: object) -> object: + return await get_data() + + await get_key_object() + + assert await _db_call_types(success_hook) == ("get_data",) + + +@pytest.mark.asyncio +async def test_an_outer_call_that_also_queries_outside_the_inner_call_emits_its_own_event( + success_hook: AsyncMock, +) -> None: + engine: Final = _tracked_engine() + + @log_db_metrics + async def get_object_permission(**kwargs: object) -> object: + return await engine.query("{}", tx_id=None) + + @log_db_metrics + async def get_key_object(**kwargs: object) -> object: + await engine.query("{}", tx_id=None) + return await get_object_permission() + + await get_key_object() + + assert await _db_call_types(success_hook) == ("get_object_permission", "get_key_object") + + +@pytest.mark.asyncio +async def test_a_query_inside_an_inner_call_that_fails_without_a_db_error_is_reported_by_the_outer_call( + success_hook: AsyncMock, +) -> None: + engine: Final = _tracked_engine() + + @log_db_metrics + async def read_row(**kwargs: object) -> object: + await engine.query("{}", tx_id=None) + raise ValueError("row did not validate") + + @log_db_metrics + async def get_key_object(**kwargs: object) -> str: + try: + await read_row() + except ValueError: + return "fallback" + return "row" + + assert await get_key_object() == "fallback" + assert await _db_call_types(success_hook) == ("get_key_object",) + + +@pytest.mark.asyncio +async def test_a_cache_hit_after_a_sibling_db_read_emits_no_db_event(success_hook: AsyncMock) -> None: + engine: Final = _tracked_engine() + + @log_db_metrics + async def read_row(**kwargs: object) -> object: + return await engine.query("{}", tx_id=None) + + @log_db_metrics + async def cache_hit(**kwargs: object) -> str: + return "cached" + + await read_row() + await cache_hit() + + assert await _db_call_types(success_hook) == ("read_row",)