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 <fn> 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 <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-03 09:20:27 -07:00 • committed by GitHub
parent 66a422ea50
commit 53d2ab6b05
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 321 additions and 5 deletions

View file

@ -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

View file

@ -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)

View file

@ -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"

View file

@ -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"

View file

@ -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",)