fix(otel): name postgres service spans by operation and table (#44240)

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:28 -07:00 • committed by GitHub
parent 564d236985
commit 797353f13a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
62 changed files with 2452 additions and 312 deletions

View file

@ -17,6 +17,7 @@ from litellm_enterprise.types.enterprise_callbacks.send_emails import (
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.db.db_span import db_span
router = APIRouter()
@ -94,16 +95,17 @@ async def _save_email_settings(prisma_client, settings: Dict[str, bool]):
json_settings = json.dumps(general_settings, default=str)
# Save updated general settings
await prisma_client.db.litellm_config.upsert(
where={"param_name": "general_settings"},
data={
"create": {
"param_name": "general_settings",
"param_value": json_settings,
async with db_span("save_email_settings", "LiteLLM_Config"):
await prisma_client.db.litellm_config.upsert(
where={"param_name": "general_settings"},
data={
"create": {
"param_name": "general_settings",
"param_value": json_settings,
},
"update": {"param_value": json_settings},
},
"update": {"param_value": json_settings},
},
)
)
except Exception as e:
raise HTTPException(
status_code=500,

View file

@ -11,15 +11,15 @@ A traced proxy request produces one trace with two kinds of spans:
```
SERVER span "POST /v1/chat/completions" ← FastAPI instrumentation
├── INTERNAL span "auth /v1/chat/completions" ← auth phase ┐
│ ├── CLIENT span "postgres get_key_object" ← datastore call │
│ └── CLIENT span "postgres get_team_membership" │
│ ├── CLIENT span "postgres.select LiteLLM_VerificationToken" │
│ └── CLIENT span "postgres.select LiteLLM_TeamMembership" │
├── INTERNAL span "execute_guardrail …" ← guardrail │ this package
├── INTERNAL span "cache.get llm_response" ← response cache │
│ └── CLIENT span "redis.get llm_response" │
├── INTERNAL span "route gpt-4o" ← deployment pick │
│ └── CLIENT span "redis.mget router_cooldowns" │
├── CLIENT span "chat gpt-4o" ← LLM call │
└── CLIENT span "batch_write_to_db …" ← spend write ┘
└── CLIENT span "postgres.update LiteLLM_UserTable" ← spend flush┘
```
The gen-ai spans are siblings under the server span. In particular the guardrail
@ -107,13 +107,26 @@ batch op settles on its own and one write-back of three auth objects shows as th
parallel `redis.set auth_objects` spans with the same caller, not one pipeline
span. Every `call_type` the Redis cache layer emits maps to a verb,
so the `{service} {call_type}` fallback is unreachable for Redis (a test asserts
it). Postgres spans are `postgres.{verb} {table}`, see
https://github.com/BerriAI/litellm/pull/44240; the other non-Redis services keep
the `"{service} {call_type}"` name (`"batch_write_to_db _PROXY_track_cost_callback"`):
one scheme, `{service}.{verb} {target}` when the method maps to a verb and
`{service} {call_type}` otherwise, and never a count, key or id in the name. Either way
the raw method name stays on `litellm.service.call_type` and `db.operation.name`
(and the bare `call_type` the metrics are keyed by), the target lands on
it). Postgres helpers are `postgres.{verb} {table}`
(`"postgres.select LiteLLM_VerificationToken"`, `"postgres.update LiteLLM_TeamTable"`,
`"postgres.insert LiteLLM_SpendLogs"`): the SQL verb comes from
`_POSTGRES_OPERATION_BY_CALL_TYPE` in `model/spans.py`, the table from that map when
the helper only touches one model and from the event's `table_name` metadata (the
`PrismaClient` CRUD literals, or a `LiteLLM_*` model name from `db_span`)
otherwise. A helper the map does not know, or one whose table did not resolve,
keeps `"postgres {call_type}"`, so a half-named `"postgres.select"` never ships;
every `PrismaClient` CRUD call site passes a literal `table_name` and a static scan
in the unit tests holds that line. A raw statement no producer wraps is
named by `_TrackedPrismaEngine` itself from the Prisma payload (`postgres.select
LiteLLM_UserTable` for `query_raw`, `postgres.set statement_timeout`, `postgres.ping`
for the health probes), see https://github.com/BerriAI/litellm/pull/44240. The verb
lands on `db.operation.name`, the table on `db.collection.name` and `"{VERB} {table}"`
on `db.query.summary`. Every other non-Redis service keeps `"{service} {call_type}"`
(`"reset_budget_job reset_budget"`): one scheme, `{service}.{verb} {target}`
when the method maps to a verb and `{service} {call_type}` otherwise, and never a
count, key or id in the name. Either way the raw method name stays on
`litellm.service.call_type` (and the bare `call_type` the metrics are keyed by; for
Redis it is also `db.operation.name`), the target lands on
`litellm.service.target`, and the litellm call chain that issued the call
(`_retrieve_from_cache <- _async_get_cache`) travels as
`ServiceLoggerPayload.caller` onto `litellm.service.caller`, with the forwarding

View file

@ -37,6 +37,7 @@ from litellm.integrations.otel.model.semconv import (
RpcSystem,
Server,
)
from litellm.integrations.otel.model.spans import postgres_operation
class GenAIMapper:
@ -196,7 +197,7 @@ class GenAIMapper:
# An outbound datastore call (DB_CALL / CLIENT span) also carries db.*
# semconv naming the server it reached. Internal services (router, budget
# jobs, …) have no db.system, so they get only the litellm.service.* keys.
attrs.update(db_span_attributes(data.service_name, data.call_type))
attrs.update(db_span_attributes(data.service_name, data.call_type, postgres_operation(data)))
attrs.update(
{
LiteLLM.REDIS_FAMILIES

View file

@ -19,7 +19,7 @@ from typing import Final
from urllib.parse import ParseResult, parse_qs, unquote, urlparse
from litellm.integrations.otel.model.semconv import DB, Server
from litellm.integrations.otel.model.spans import POSTGRESQL, db_system
from litellm.integrations.otel.model.spans import POSTGRESQL, PostgresOperation, db_system
_DATABASE_URL_ENV: Final = "DATABASE_URL"
_READ_REPLICA_ENV: Final = "DATABASE_URL_READ_REPLICA"
@ -140,23 +140,31 @@ def postgres_endpoint() -> DatabaseEndpoint | None:
return parse_database_endpoint(os.environ.get(_DATABASE_URL_ENV, ""))
def db_span_attributes(service_name: str, call_type: str | None = None) -> Mapping[str, str | int]:
def db_span_attributes(
service_name: str, call_type: str | None = None, operation: PostgresOperation | None = None
) -> Mapping[str, str | int]:
"""The ``db.*``/``server.*`` attributes for a datastore service call.
Empty for services that are not outbound datastore calls. Endpoint
attributes are PostgreSQL-only: ``DATABASE_URL`` says nothing about where
the redis-backed services point. ``db.system`` rides alongside the current
``db.system.name`` because Datadog's OTLP intake still types a database span
from the older key.
from the older key. A resolved Prisma ``operation`` puts the SQL verb on
``db.operation.name`` (the raw method stays on ``litellm.service.call_type``),
the table (or the declared ``collection`` list) on ``db.collection.name`` and ``"{VERB} {table}"`` on
``db.query.summary``; without one, ``db.operation.name`` is the call type.
"""
system: Final = db_system(service_name)
if system is None:
return _EMPTY_ATTRIBUTES
endpoint: Final = postgres_endpoint() if system == POSTGRESQL else None
table: Final = operation.table if operation is not None else None
pairs: Final[tuple[tuple[str, str | int | None], ...]] = (
(DB.SYSTEM_NAME, system),
(DB.SYSTEM_LEGACY, system),
(DB.OPERATION_NAME, call_type),
(DB.OPERATION_NAME, operation.verb if operation is not None else call_type),
(DB.COLLECTION_NAME, operation.collection or table if operation is not None else None),
(DB.QUERY_SUMMARY, f"{operation.verb.upper()} {table}" if operation is not None and table else None),
(Server.ADDRESS, endpoint.address if endpoint is not None else None),
(Server.PORT, endpoint.port if endpoint is not None else None),
(DB.NAMESPACE, endpoint.namespace if endpoint is not None else None),

View file

@ -266,6 +266,8 @@ class DB:
# still infers a span's database type from this key.
SYSTEM_LEGACY: Final = "db.system"
OPERATION_NAME: Final = "db.operation.name"
COLLECTION_NAME: Final = "db.collection.name"
QUERY_SUMMARY: Final = "db.query.summary"
NAMESPACE: Final = "db.namespace"

View file

@ -49,8 +49,11 @@ Management/admin endpoints are ordinary FastAPI routes — their SERVER spans ar
owned by the instrumentor too, so they don't appear as a role here.
"""
import re
from collections.abc import Mapping
from dataclasses import dataclass
from enum import Enum
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
if TYPE_CHECKING:
@ -230,9 +233,290 @@ _SERVICE_VERB_BY_CALL_TYPE: Final[dict[str, str]] = {
}
@dataclass(frozen=True, slots=True)
class PostgresOperation:
"""The SQL verb and primary table behind a Prisma helper, for ``postgres.{verb} {table}``;
``collection`` lists every relation on ``db.collection.name`` when one query joins several."""
verb: str
table: str | None
collection: str | None = None
_POSTGRES_SERVICE: Final = "postgres"
PG_CATALOG: Final = "pg_catalog"
_PRISMA_VIEWS: Final[frozenset[str]] = frozenset(
(
"LiteLLM_VerificationTokenView",
"MonthlyGlobalSpend",
"Last30dKeysBySpend",
"Last30dModelsBySpend",
"MonthlyGlobalSpendPerKey",
"MonthlyGlobalSpendPerUserPerKey",
"Last30dTopEndUsersSpend",
"DailyTagSpend",
)
)
_PRISMA_MODELS: Final[frozenset[str]] = frozenset(
(
"LiteLLM_BudgetTable",
"LiteLLM_CredentialsTable",
"LiteLLM_ProxyModelTable",
"LiteLLM_AgentsTable",
"LiteLLM_AgentIdentity",
"LiteLLM_RetiredAgentIdentity",
"LiteLLM_RetiredAgent",
"LiteLLM_VerifiedSubject",
"LiteLLM_OrganizationTable",
"LiteLLM_ModelTable",
"LiteLLM_TeamTable",
"LiteLLM_ProjectTable",
"LiteLLM_DeletedTeamTable",
"LiteLLM_UserTable",
"LiteLLM_ObjectPermissionTable",
"LiteLLM_MCPServerTable",
"LiteLLM_MCPToolsetTable",
"LiteLLM_MCPUserCredentials",
"LiteLLM_MCPUserEnvVars",
"LiteLLM_MCPServerOAuthClient",
"LiteLLM_SSOIdentityAssertion",
"LiteLLM_VerificationToken",
"LiteLLM_JWTKeyMapping",
"LiteLLM_DeprecatedVerificationToken",
"LiteLLM_DeletedVerificationToken",
"LiteLLM_EndUserTable",
"LiteLLM_ModelAccessGroupBudgetTable",
"LiteLLM_TagTable",
"LiteLLM_Config",
"LiteLLM_SpendLogs",
"LiteLLM_BudgetWindowSpend",
"LiteLLM_ErrorLogs",
"LiteLLM_UserNotifications",
"LiteLLM_TeamMembership",
"LiteLLM_OrganizationMembership",
"LiteLLM_InvitationLink",
"LiteLLM_AuditLog",
"LiteLLM_DailyUserSpend",
"LiteLLM_DailyGlobalSpend",
"LiteLLM_DailyOrganizationSpend",
"LiteLLM_DailyEndUserSpend",
"LiteLLM_DailyAgentSpend",
"LiteLLM_DailyTeamSpend",
"LiteLLM_DailyTagSpend",
"LiteLLM_ProxyWorkerHeartbeat",
"LiteLLM_CronJob",
"LiteLLM_ManagedFileTable",
"LiteLLM_ManagedObjectTable",
"LiteLLM_ManagedFileContentTable",
"LiteLLM_ManagedVectorStoreTable",
"LiteLLM_ManagedVectorStoresTable",
"LiteLLM_GuardrailsTable",
"LiteLLM_DailyGuardrailMetrics",
"LiteLLM_DailyGuardrailUsageUnits",
"LiteLLM_DailyPolicyMetrics",
"LiteLLM_SpendLogGuardrailIndex",
"LiteLLM_SpendLogToolIndex",
"LiteLLM_DailyToolSpend",
"LiteLLM_DailyModelUsage",
"LiteLLM_DailyGatewayRequests",
"LiteLLM_PromptTable",
"LiteLLM_HealthCheckTable",
"LiteLLM_SearchToolsTable",
"LiteLLM_SSOConfig",
"LiteLLM_ManagedVectorStoreIndexTable",
"LiteLLM_CacheConfig",
"LiteLLM_UISettings",
"LiteLLM_ConfigOverrides",
"LiteLLM_SkillsTable",
"LiteLLM_PolicyTable",
"LiteLLM_PolicyAttachmentTable",
"LiteLLM_ToolTable",
"LiteLLM_AccessGroupTable",
"LiteLLM_ClaudeCodePluginTable",
"LiteLLM_MemoryTable",
"LiteLLM_AdaptiveRouterState",
"LiteLLM_AdaptiveRouterSession",
"LiteLLM_AutoRouterBaselineComparison",
"LiteLLM_AutoRouterBaselineObservation",
"LiteLLM_AutoRouterSession",
"LiteLLM_AutoRouterUserSession",
"LiteLLM_AutoRouterDailySpend",
"LiteLLM_ShadowEvalJob",
"LiteLLM_ShadowEvalAttempt",
"LiteLLM_ShadowEvalFunnel",
"LiteLLM_WorkflowRun",
"LiteLLM_WorkflowEvent",
"LiteLLM_WorkflowMessage",
"LiteLLM_Lens",
"LiteLLM_LensRun",
"LiteLLM_LensWorker",
)
)
PRISMA_RELATIONS: Final[frozenset[str]] = _PRISMA_MODELS | _PRISMA_VIEWS
_TABLE_NAME_METADATA_KEY: Final = "table_name"
_PRISMA_MODEL_BY_TABLE_NAME: Final[Mapping[str, str]] = MappingProxyType(
{
"key": "LiteLLM_VerificationToken",
"keys": "LiteLLM_VerificationToken",
"combined_view": "LiteLLM_VerificationToken",
"user": "LiteLLM_UserTable",
"users": "LiteLLM_UserTable",
"team": "LiteLLM_TeamTable",
"config": "LiteLLM_Config",
"spend": "LiteLLM_SpendLogs",
"enduser": "LiteLLM_EndUserTable",
"budget": "LiteLLM_BudgetTable",
"user_notification": "LiteLLM_UserNotifications",
}
)
_AUTH_OBJECT_RELATIONS: Final = ",".join(
(
"LiteLLM_UserTable",
"LiteLLM_TeamTable",
"LiteLLM_TeamMembership",
"LiteLLM_OrganizationTable",
"LiteLLM_OrganizationMembership",
"LiteLLM_ProjectTable",
"LiteLLM_ModelTable",
"LiteLLM_BudgetTable",
"LiteLLM_ObjectPermissionTable",
)
)
_POSTGRES_OPERATION_BY_CALL_TYPE: Final[Mapping[str, PostgresOperation]] = MappingProxyType(
{
"get_data": PostgresOperation("select", None),
"get_generic_data": PostgresOperation("select", None),
"insert_data": PostgresOperation("insert", None),
"update_data": PostgresOperation("update", None),
"delete_data": PostgresOperation("delete", None),
"get_key_object": PostgresOperation("select", "LiteLLM_VerificationToken"),
"get_user_object": PostgresOperation("select", "LiteLLM_UserTable"),
"get_org_object": PostgresOperation("select", "LiteLLM_OrganizationTable"),
"get_org_object_by_alias": PostgresOperation("select", "LiteLLM_OrganizationTable"),
"_get_team_db_check": PostgresOperation("select", "LiteLLM_TeamTable"),
"get_team_object_by_alias": PostgresOperation("select", "LiteLLM_TeamTable"),
"_fetch_team_membership_from_db": PostgresOperation("select", "LiteLLM_TeamMembership"),
"get_team_member_default_budget": PostgresOperation("select", "LiteLLM_BudgetTable"),
"get_end_user_object": PostgresOperation("select", "LiteLLM_EndUserTable"),
"get_tag_object": PostgresOperation("select", "LiteLLM_TagTable"),
"get_tag_objects_batch": PostgresOperation("select", "LiteLLM_TagTable"),
"get_model_access_group_budgets_batch": PostgresOperation("select", "LiteLLM_ModelAccessGroupBudgetTable"),
"get_access_object": PostgresOperation("select", "LiteLLM_AccessGroupTable"),
"get_object_permission": PostgresOperation("select", "LiteLLM_ObjectPermissionTable"),
"get_jwt_key_mapping_object": PostgresOperation("select", "LiteLLM_JWTKeyMapping"),
"get_jwt_key_mapping_cache_keys_for_token": PostgresOperation("select", "LiteLLM_JWTKeyMapping"),
"get_managed_vector_store_rows_by_uuids": PostgresOperation("select", "LiteLLM_ManagedVectorStoresTable"),
"commit_spend_updates": PostgresOperation("update", None),
"update_end_user_spend": PostgresOperation("upsert", None),
"upsert_daily_spend": PostgresOperation("upsert", None),
"insert_spend_logs": PostgresOperation("insert", None),
"migrate_config_credentials": PostgresOperation("update", None),
"migrate_sso_credentials": PostgresOperation("update", None),
"backfill_mcp_oauth_issuer": PostgresOperation("update", None),
"auto_register_jwt_mapping": PostgresOperation("insert", None),
"delete_orphaned_jwt_key": PostgresOperation("delete", None),
"save_email_settings": PostgresOperation("upsert", None),
"reset_budget_cascade": PostgresOperation("transaction", None),
"reset_spend_rows": PostgresOperation("update", None),
"reset_budget_windows": PostgresOperation("select", None),
"write_budget_windows": PostgresOperation("update", None),
"roll_window_spend_row": PostgresOperation("update", None),
"seed_window_spend": PostgresOperation("select", None),
"select_window_spend_rows": PostgresOperation("select", None),
"commit_window_spend_updates": PostgresOperation("upsert", None),
"index_spend_log_tools": PostgresOperation("insert", None),
"commit_daily_tool_spend": PostgresOperation("upsert", None),
"flush_shadow_eval_funnel": PostgresOperation("upsert", None),
"commit_gateway_requests": PostgresOperation("upsert", None),
"cleanup_expired_rows": PostgresOperation("delete", None),
"count_expired_rows": PostgresOperation("select", None),
"check_spend_log_partitioning": PostgresOperation("select", None),
"list_spend_log_partitions": PostgresOperation("select", None),
"create_spend_log_partition": PostgresOperation("ddl", None),
"proxy_worker_heartbeat": PostgresOperation("upsert", None),
"prune_proxy_worker_heartbeats": PostgresOperation("delete", None),
"deregister_proxy_worker": PostgresOperation("delete", None),
"count_live_proxy_workers": PostgresOperation("select", None),
"recover_key_metadata": PostgresOperation("select", None),
"recover_user_details": PostgresOperation("select", None),
"sync_team_access_group_membership": PostgresOperation("transaction", None),
"latest_health_checks": PostgresOperation("select", None),
"prefetch_auth_objects": PostgresOperation("select", "auth_objects", _AUTH_OBJECT_RELATIONS),
"baseline_accounting": PostgresOperation("transaction", "LiteLLM_AutoRouterBaselineComparison"),
"write_autorouter_turn": PostgresOperation("upsert", None),
"team_user_spend": PostgresOperation("select", "LiteLLM_SpendLogs"),
"daily_activity_query": PostgresOperation("select", None),
"auto_router_report_query": PostgresOperation("select", None),
"create_view": PostgresOperation("ddl", None),
"health_check": PostgresOperation("ping", None),
"db_health_watchdog": PostgresOperation("ping", None),
"find_unique": PostgresOperation("select", None),
"find_first": PostgresOperation("select", None),
"find_many": PostgresOperation("select", None),
"count": PostgresOperation("select", None),
"group_by": PostgresOperation("select", None),
"create": PostgresOperation("insert", None),
"create_many": PostgresOperation("insert", None),
"update": PostgresOperation("update", None),
"update_many": PostgresOperation("update", None),
"delete": PostgresOperation("delete", None),
"delete_many": PostgresOperation("delete", None),
"upsert": PostgresOperation("upsert", None),
}
)
_RAW_PRISMA_CALL_TYPES: Final[frozenset[str]] = frozenset(("query_raw", "execute_raw"))
_DB_OPERATION_METADATA_KEY: Final = "db_operation"
_POSTGRES_VERBS: Final[frozenset[str]] = frozenset(
("select", "insert", "update", "delete", "upsert", "ddl", "set", "ping")
)
_TARGETLESS_VERBS: Final[frozenset[str]] = frozenset(("ping",))
_SETTING_NAME: Final = re.compile(r"[a-z_][a-z0-9_.]*")
def _postgres_table_from_metadata(data: "ServiceSpanData", verb: str) -> str | None:
"""The relation named by the event's ``table_name`` metadata, or ``None``.
Only the short ``PrismaClient`` literals, the relations declared in ``schema.prisma``
(plus the spend views), ``pg_catalog`` and, for a ``set`` verb, a Postgres setting name
resolve, so a free-form string can never become a span-name cardinality."""
table_name: Final = data.event_metadata.get(_TABLE_NAME_METADATA_KEY)
if not isinstance(table_name, str):
return None
if table_name in PRISMA_RELATIONS or table_name == PG_CATALOG:
return table_name
if verb == "set":
return table_name if _SETTING_NAME.fullmatch(table_name) else None
return _PRISMA_MODEL_BY_TABLE_NAME.get(table_name)
def _postgres_verb_from_metadata(data: "ServiceSpanData") -> str | None:
"""The SQL verb a raw-statement producer declared on ``db_operation``, bounded to the known verbs."""
verb: Final = data.event_metadata.get(_DB_OPERATION_METADATA_KEY)
return verb if isinstance(verb, str) and verb in _POSTGRES_VERBS else None
def postgres_operation(data: "ServiceSpanData") -> PostgresOperation | None:
"""The verb and table behind a ``postgres`` service event, else ``None``.
``None`` for every other service (Redis keeps its own verb table) and for a
Postgres call type this module does not know, which stays ``postgres {call_type}``."""
if data.service_name != _POSTGRES_SERVICE or not data.call_type:
return None
if data.call_type in _RAW_PRISMA_CALL_TYPES:
verb: Final = _postgres_verb_from_metadata(data)
return PostgresOperation(verb, _postgres_table_from_metadata(data, verb)) if verb is not None else None
operation: Final = _POSTGRES_OPERATION_BY_CALL_TYPE.get(data.call_type)
if operation is None:
return None
if operation.table is not None:
return operation
return PostgresOperation(operation.verb, _postgres_table_from_metadata(data, operation.verb))
def service_operation(data: "ServiceSpanData") -> str | None:
"""``"redis.get"`` when the call type is a known datastore verb, else ``None``
(Postgres helpers stay function-named until they get ``db.select {table}`` names)."""
"""``"redis.get"`` when the call type is a known datastore verb, else ``None``."""
if not data.call_type:
return None
verb: Final = _SERVICE_VERB_BY_CALL_TYPE.get(data.call_type)
@ -244,8 +528,19 @@ def service_operation(data: "ServiceSpanData") -> str | None:
def service_span_name(data: "ServiceSpanData") -> str:
"""``"{service}.{verb} {target}"`` (``"redis.get llm_response"``) for a known datastore
verb, ``"{service}.{verb}"`` (``"redis.pipeline"``) when the producer declared no
target, else ``"{service} {call_type}"`` (``"postgres get_data"``) — service name alone
target, ``"postgres.{verb} {table}"`` (``"postgres.select LiteLLM_UserTable"``) for a
known Prisma helper whose table resolved (from the helper or the event's ``table_name``,
never from the ambient ``service_target``, which names a cache key family), else
``"{service} {call_type}"`` (``"postgres some_helper"``, and a known helper whose table
did not resolve, so a half-named ``postgres.select`` never ships) — service name alone
when no call type is known, so identically-named calls stay distinguishable."""
postgres: Final = postgres_operation(data)
if postgres is not None and postgres.table is not None:
return f"{data.service_name}.{postgres.verb} {postgres.table}"
if postgres is not None and postgres.verb in _TARGETLESS_VERBS:
return f"{data.service_name}.{postgres.verb}"
if postgres is not None:
return f"{data.service_name} {data.call_type}"
operation: Final = service_operation(data)
if operation is None:
return f"{data.service_name} {data.call_type or ''}".strip()

View file

@ -41,6 +41,7 @@ from urllib.parse import urlparse
from litellm._logging import verbose_proxy_logger
from litellm.proxy._experimental.mcp_server.oauth_utils import canonicalize_url_identity
from litellm.proxy.db.db_span import db_span
from litellm.proxy.utils import PrismaClient
# The actor the removed discovery write-back stamped rows with.
@ -118,10 +119,11 @@ async def backfill_discovery_stamped_issuers(prisma_client: PrismaClient) -> int
healed = 0
for row in stamped:
try:
await prisma_client.db.litellm_mcpservertable.update(
where={"server_id": row.server_id},
data={"issuer": None, "updated_by": _BACKFILL_ACTOR},
)
async with db_span("backfill_mcp_oauth_issuer", "LiteLLM_MCPServerTable"):
await prisma_client.db.litellm_mcpservertable.update(
where={"server_id": row.server_id},
data={"issuer": None, "updated_by": _BACKFILL_ACTOR},
)
except Exception as exc: # noqa: BLE001 - per-row best effort; the next boot retries
verbose_proxy_logger.warning(
"MCP issuer stamp backfill: could not heal server_id=%s: %s", row.server_id, exc

View file

@ -30,6 +30,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
team_membership_auth_cache_key,
team_membership_reservation_cache_key,
)
from litellm.proxy.db.db_span import db_span
from litellm.proxy.utils import PrismaClient
_RowKind: TypeAlias = Literal["user_row", "team_row", "membership_row", "organization_row", "project_row"]
@ -264,14 +265,15 @@ def _validate_row(
async def _fetch_rows(
refs: AuthObjectRefs, kinds: frozenset[_RowKind], prisma_client: PrismaClient
) -> Mapping[str, object]:
row: Final[object] = await prisma_client.db.query_first( # pyright: ignore[reportAny] # prisma types query_first as Any
_SQL,
refs.user_id if "user_row" in kinds else None,
refs.team_id if kinds & _TEAM_BOUND_ROWS else None,
refs.membership_user_id if "membership_row" in kinds else None,
refs.organization_id if "organization_row" in kinds else None,
refs.project_id if "project_row" in kinds else None,
)
async with db_span("prefetch_auth_objects", AUTH_OBJECTS_TARGET):
row: Final[object] = await prisma_client.db.query_first( # pyright: ignore[reportAny] # prisma types query_first as Any
_SQL,
refs.user_id if "user_row" in kinds else None,
refs.team_id if kinds & _TEAM_BOUND_ROWS else None,
refs.membership_user_id if "membership_row" in kinds else None,
refs.organization_id if "organization_row" in kinds else None,
refs.project_id if "project_row" in kinds else None,
)
return _RowValues.validate_python(row) if row is not None else _NO_ROWS

View file

@ -132,6 +132,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
team_membership_auth_cache_key,
)
from litellm.proxy.db.db_lookup_gate import bounded_db_lookup
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.spend_tracking.carried_budget_state import carry_team_and_user_budget_state
@ -1035,16 +1036,17 @@ async def _auto_register_jwt_mapping(
token_hash = hash_token(key_data["token"])
try:
await prisma_client.db.litellm_jwtkeymapping.create(
data={
"jwt_issuer": jwt_issuer or "",
"jwt_claim_name": virtual_key_claim_field,
"jwt_claim_value": claim_value,
"token": token_hash,
"created_by": "auto_register",
"updated_by": "auto_register",
}
)
async with db_span("auto_register_jwt_mapping", "LiteLLM_JWTKeyMapping"):
await prisma_client.db.litellm_jwtkeymapping.create(
data={
"jwt_issuer": jwt_issuer or "",
"jwt_claim_name": virtual_key_claim_field,
"jwt_claim_value": claim_value,
"token": token_hash,
"created_by": "auto_register",
"updated_by": "auto_register",
}
)
except Exception as e:
error_str: Final = str(e).lower()
if "unique" in error_str or "p2002" in error_str:
@ -1061,7 +1063,8 @@ async def _auto_register_jwt_mapping(
)
if minted:
try:
await prisma_client.db.litellm_verificationtoken.delete(where={"token": token_hash})
async with db_span("delete_orphaned_jwt_key", "LiteLLM_VerificationToken"):
await prisma_client.db.litellm_verificationtoken.delete(where={"token": token_hash})
except Exception as delete_err:
# Don't fail the request if cleanup fails — the orphan is
# unmapped and inert. Log so an operator can prune it later.

View file

@ -47,6 +47,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
tag_cache_key,
)
from litellm.proxy.db.budget_window_spend_writer import roll_window_spend_row
from litellm.proxy.db.db_span import db_span, db_spanned
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import POD_LOCK_TARGET, PodLockManager
from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry
from litellm.proxy.spend_tracking.spend_counter_batch import SPEND_COUNTERS_TARGET
@ -383,17 +384,19 @@ class _Lease(Enum):
async def _write_key_windows(prisma_client: PrismaClient, row_id: str, payload: str) -> None:
await VerificationTokenRepository(prisma_client).table.update(
where={"token": row_id},
data={"budget_limits": payload},
)
async with db_span("write_budget_windows", "LiteLLM_VerificationToken"):
await VerificationTokenRepository(prisma_client).table.update(
where={"token": row_id},
data={"budget_limits": payload},
)
async def _write_team_windows(prisma_client: PrismaClient, row_id: str, payload: str) -> None:
await TeamRepository(prisma_client).table.update(
where={"team_id": row_id},
data={"budget_limits": payload},
)
async with db_span("write_budget_windows", "LiteLLM_TeamTable"):
await TeamRepository(prisma_client).table.update(
where={"team_id": row_id},
data={"budget_limits": payload},
)
@dataclass(frozen=True, slots=True)
@ -838,7 +841,10 @@ class ResetBudgetJob:
)
async def _commit_budget_cascade_once(self, cascade: _BudgetCascade) -> None:
async with budget_cascade_unit_of_work(self._new_batch) as uow:
async with (
db_span("reset_budget_cascade", "LiteLLM_BudgetTable"),
budget_cascade_unit_of_work(self._new_batch) as uow,
):
_queue_budget_linked_resets(uow.team_memberships, cascade)
_queue_budget_linked_resets(uow.keys, cascade, extra=_LINKED_KEYS_WHERE)
_queue_budget_linked_resets(uow.organizations, cascade, extra=_SPENT_ROWS_WHERE)
@ -960,7 +966,10 @@ class ResetBudgetJob:
)
async def _write_key_reset_updates_once(self, updated_keys: Sequence[_RowReset[LiteLLM_VerificationToken]]) -> None:
async with spend_reset_unit_of_work(self._new_batch) as uow:
async with (
db_span("reset_spend_rows", "LiteLLM_VerificationToken"),
spend_reset_unit_of_work(self._new_batch) as uow,
):
for k in updated_keys:
if k.row.token is None:
continue
@ -984,7 +993,10 @@ class ResetBudgetJob:
)
async def _write_user_reset_updates_once(self, updated_users: Sequence[_RowReset[LiteLLM_UserTable]]) -> None:
async with spend_reset_unit_of_work(self._new_batch) as uow:
async with (
db_span("reset_spend_rows", "LiteLLM_UserTable"),
spend_reset_unit_of_work(self._new_batch) as uow,
):
for u in updated_users:
uow.users.queue_spend_reset(
user_id=u.row.user_id,
@ -1006,7 +1018,10 @@ class ResetBudgetJob:
)
async def _write_team_reset_updates_once(self, updated_teams: Sequence[_RowReset[LiteLLM_TeamTable]]) -> None:
async with spend_reset_unit_of_work(self._new_batch) as uow:
async with (
db_span("reset_spend_rows", "LiteLLM_TeamTable"),
spend_reset_unit_of_work(self._new_batch) as uow,
):
for t in updated_teams:
uow.teams.queue_spend_reset(
team_id=t.row.team_id,
@ -1532,7 +1547,11 @@ class ResetBudgetJob:
) -> str | None:
"""Reset one page of windows; return the next cursor, or None when drained."""
rows: Final = await self._with_db_retry(
lambda: self.prisma_client.db.query_raw(source.page_query(), cursor, RESET_BUDGET_JOB_BATCH_SIZE),
lambda: db_spanned(
"reset_budget_windows",
source.table,
lambda: self.prisma_client.db.query_raw(source.page_query(), cursor, RESET_BUDGET_JOB_BATCH_SIZE),
),
reason=f"reset_budget_read_{source.retry_subject}_windows_failure",
)
for row in rows:

View file

@ -22,12 +22,14 @@ from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, NamedTuple
from litellm._logging import verbose_proxy_logger
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES
from litellm.proxy.db.create_views import SupportsExecuteRaw
from litellm.proxy.db.db_span import db_span
if TYPE_CHECKING:
from litellm.proxy._types import SpendLogsPayload
@ -468,6 +470,13 @@ WITH {_DAY_UPSERT_SQL}
{_session_upsert_sql(user_scoped=True)}
"""
_SESSION_TABLE_BY_STATEMENT: Final[Mapping[str, str]] = MappingProxyType(
{
UPSERT_AUTOROUTER_SESSION_SQL: "LiteLLM_AutoRouterSession",
UPSERT_AUTOROUTER_USER_SESSION_SQL: "LiteLLM_AutoRouterUserSession",
}
)
def _as_sql_param(value: str | float | bool | datetime | None) -> str | float | None:
if isinstance(value, bool):
@ -486,7 +495,8 @@ async def write_autorouter_turn(
transaction: AutoRouterTurnTransaction,
statement: str = UPSERT_AUTOROUTER_SESSION_SQL,
) -> None:
await db.execute_raw(statement, *_upsert_params(transaction))
async with db_span("write_autorouter_turn", _SESSION_TABLE_BY_STATEMENT.get(statement)):
await db.execute_raw(statement, *_upsert_params(transaction))
async def _upsert_turn_with_retry(

View file

@ -2,7 +2,8 @@ from __future__ import annotations
import asyncio
import json
from collections.abc import AsyncIterator, Callable, Sequence
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Sequence
from contextlib import asynccontextmanager
from datetime import datetime, timedelta
from functools import reduce
from itertools import groupby
@ -25,6 +26,7 @@ from litellm.proxy.db.daily_spend_bulk_upsert import (
build_bulk_upsert,
merge_by_conflict_key,
)
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper
from litellm.proxy.spend_tracking.baseline_accounting import (
BaselineEstimate,
@ -417,8 +419,13 @@ class BaselineAccountingStore:
@classmethod
def for_client(cls, client: PrismaClient) -> BaselineAccountingStore:
def transaction() -> _TransactionManager:
return _primary_transaction(client)
@asynccontextmanager
async def transaction() -> AsyncGenerator[SupportsRawQueries]:
async with (
db_span("baseline_accounting", "LiteLLM_AutoRouterBaselineComparison"),
_primary_transaction(client) as db,
):
yield db
return cls(transaction)

View file

@ -21,6 +21,7 @@ from typing import TYPE_CHECKING, Final, Protocol
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import Litellm_EntityType
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
WindowSpendTransaction,
to_naive_utc,
@ -146,16 +147,17 @@ async def spend_logs_seed_totals(
bounded_sql, unbounded_sql = _SEED_FROM_SPEND_LOGS_TEAM_SQL, _SEED_FROM_SPEND_LOGS_TEAM_UNBOUNDED_SQL
else:
return None
rows: Final = (
await prisma_client.db.query_raw(unbounded_sql, entity_id, window_start)
if batch_started_at is None
else await prisma_client.db.query_raw(
bounded_sql,
entity_id,
window_start,
_exclusion_upper_bound(batch_started_at),
async with db_span("seed_window_spend", "LiteLLM_SpendLogs"):
rows: Final = (
await prisma_client.db.query_raw(unbounded_sql, entity_id, window_start)
if batch_started_at is None
else await prisma_client.db.query_raw(
bounded_sql,
entity_id,
window_start,
_exclusion_upper_bound(batch_started_at),
)
)
)
if not rows:
return WindowSeedTotals(total=0.0, before_batch=0.0)
return WindowSeedTotals(
@ -182,12 +184,13 @@ async def _existing_primary_keys(
prisma_client: "PrismaClient",
transactions: tuple[WindowSpendTransaction, ...],
) -> frozenset[tuple[str, str, str]]:
rows: Final = await prisma_client.db.query_raw(
_SELECT_EXISTING_ROWS_SQL,
tuple(transaction["entity_type"] for transaction in transactions),
tuple(transaction["entity_id"] for transaction in transactions),
tuple(transaction["window_duration"] for transaction in transactions),
)
async with db_span("select_window_spend_rows", "LiteLLM_BudgetWindowSpend"):
rows: Final = await prisma_client.db.query_raw(
_SELECT_EXISTING_ROWS_SQL,
tuple(transaction["entity_type"] for transaction in transactions),
tuple(transaction["entity_id"] for transaction in transactions),
tuple(transaction["window_duration"] for transaction in transactions),
)
return frozenset((row["entity_type"], row["entity_id"], row["window_duration"]) for row in rows or ())
@ -300,6 +303,7 @@ async def commit_window_spend_updates(
len(existing_primary_keys),
)
async with (
db_span("commit_window_spend_updates", "LiteLLM_BudgetWindowSpend"),
prisma_client.db.tx(timeout=_UPSERT_TRANSACTION_TIMEOUT) as db_transaction,
db_transaction.batch_() as batcher,
):
@ -323,11 +327,12 @@ async def roll_window_spend_row(
pod that already rolled the row (or increments that arrived under the new
window) are not clobbered.
"""
await prisma_client.db.execute_raw(
_ROLL_WINDOW_SPEND_SQL,
entity_type,
entity_id,
window_duration,
to_naive_utc(new_window_start),
to_naive_utc(datetime.now(timezone.utc)),
)
async with db_span("roll_window_spend_row", "LiteLLM_BudgetWindowSpend"):
await prisma_client.db.execute_raw(
_ROLL_WINDOW_SPEND_SQL,
entity_type,
entity_id,
window_duration,
to_naive_utc(new_window_start),
to_naive_utc(datetime.now(timezone.utc)),
)

View file

@ -2,6 +2,7 @@ from collections.abc import Mapping, Sequence
from typing import Final, Protocol
from litellm import verbose_logger
from litellm.proxy.db.db_span import db_span
class SupportsExecuteRaw(Protocol):
@ -38,7 +39,8 @@ async def create_view_tolerating_race(db: SupportsExecuteRaw, view_name: str, dd
a detached startup task and the remaining views are never created.
"""
try:
await db.execute_raw(ddl)
async with db_span("create_view", view_name):
await db.execute_raw(ddl)
verbose_logger.debug("%s Created!", view_name)
except Exception as e:
if not any(marker in str(e).lower() for marker in _VIEW_ALREADY_EXISTS_MARKERS):

View file

@ -0,0 +1,99 @@
"""A ``ServiceTypes.DB`` event around Prisma I/O that ``@log_db_metrics`` cannot wrap.
The spend flush, the spend-log batch insert and the background jobs run raw
``prisma_client.db`` statements and transactions, often inside retry loops, so the
decorator (one event per decorated coroutine) cannot name the table each round
trip touches. ``db_span`` emits one success or failure event per round trip,
carrying the raw ``call_type`` for the metric labels and the Prisma model on
``table_name`` so OTel renders ``postgres.{verb} {table}``; ``db_spanned`` is the
same event around a thunk, for the retry helpers that take one. The outermost
producer owns the event: a ``db_span`` nested in another ``db_span`` or in a
decorated helper emits nothing, so one transaction stays one span, and a block
whose Prisma client never reached the engine emits nothing at all.
"""
from __future__ import annotations
import asyncio
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping
from contextlib import asynccontextmanager
from datetime import datetime
from typing import Final, TypeVar
from litellm._logging import verbose_proxy_logger
from litellm._service_logger import ServiceLogging, ServiceTypes
from litellm.proxy.db.log_db_metrics import _is_exception_related_to_db, claim_db_io, db_io_claimed
_T = TypeVar("_T")
def _service_logging() -> ServiceLogging | None:
try:
from litellm.proxy.proxy_server import proxy_logging_obj
except ImportError:
return None
return proxy_logging_obj.service_logging_obj
def _event_metadata(table: str | None, operation: str | None) -> dict[str, str]:
pairs: Final = (("table_name", table), ("db_operation", operation))
return {key: value for key, value in pairs if value is not None}
async def _emit_failure(
service_logging: ServiceLogging,
call_type: str,
event_metadata: Mapping[str, str],
start_time: datetime,
error: Exception,
) -> None:
end_time: Final = datetime.now()
try:
await service_logging.async_service_failure_hook(
error=error,
service=ServiceTypes.DB,
call_type=call_type,
parent_otel_span=None,
duration=(end_time - start_time).total_seconds(),
start_time=start_time,
end_time=end_time,
event_metadata=dict(event_metadata),
)
except Exception as hook_error:
verbose_proxy_logger.debug("db_span: failure hook raised for %s: %s", call_type, hook_error)
@asynccontextmanager
async def db_span(call_type: str, table: str | None, operation: str | None = None) -> AsyncGenerator[None]:
if db_io_claimed():
yield
return
service_logging: Final = _service_logging()
start_time: Final = datetime.now()
event_metadata: Final = _event_metadata(table, operation)
with claim_db_io() as witness:
try:
yield
except Exception as e:
if service_logging is not None and _is_exception_related_to_db(e):
await _emit_failure(service_logging, call_type, event_metadata, start_time, e)
raise
if service_logging is None or not witness.touched:
return
end_time_ok: Final = datetime.now()
asyncio.create_task(
service_logging.async_service_success_hook(
service=ServiceTypes.DB,
call_type=call_type,
parent_otel_span=None,
duration=(end_time_ok - start_time).total_seconds(),
start_time=start_time,
end_time=end_time_ok,
event_metadata=event_metadata,
)
)
async def db_spanned(call_type: str, table: str | None, load: Callable[[], Awaitable[_T]]) -> _T:
async with db_span(call_type, table):
return await load()

View file

@ -13,7 +13,8 @@ import os
import random
import time
import traceback
from collections.abc import Callable, Coroutine, Mapping, Sequence
from collections.abc import AsyncGenerator, Callable, Coroutine, Mapping, Sequence
from contextlib import asynccontextmanager
from contextvars import ContextVar
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
@ -58,6 +59,7 @@ from litellm.proxy.db.daily_spend_bulk_upsert import (
daily_spend_entity_ids,
merge_by_conflict_key,
)
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
DailySpendUpdateQueue,
)
@ -274,9 +276,23 @@ def _timed_request_duration_ms(
return duration_ms
def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager:
_ENTITY_SPEND_MODELS: Final[Mapping[_EntitySpendTable, str]] = MappingProxyType(
{
"litellm_tagtable": "LiteLLM_TagTable",
"litellm_agentstable": "LiteLLM_AgentsTable",
"litellm_modelaccessgroupbudgettable": "LiteLLM_ModelAccessGroupBudgetTable",
"litellm_projecttable": "LiteLLM_ProjectTable",
}
)
@asynccontextmanager
async def _spend_update_tx(
prisma_client: PrismaClient, table: str, call_type: str = "commit_spend_updates"
) -> AsyncGenerator[_SpendTransaction]:
tx: Final[_SpendTransactionManager] = prisma_client.db.tx(timeout=timedelta(seconds=60))
return tx
async with db_span(call_type, table), tx as transaction:
yield transaction
_daily_spend_commit_started: Final[ContextVar[asyncio.Event | None]] = ContextVar(
@ -2052,7 +2068,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with _spend_update_tx(prisma_client) as transaction:
async with _spend_update_tx(prisma_client, "LiteLLM_UserTable") as transaction:
async with transaction.batch_() as batcher:
# Sort by ID for consistent lock ordering across pods to prevent deadlocks.
# batch_() issues statements sequentially within the tx, so iteration
@ -2093,7 +2109,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with _spend_update_tx(prisma_client) as transaction:
async with _spend_update_tx(prisma_client, "LiteLLM_VerificationToken") as transaction:
async with transaction.batch_() as batcher:
# Sort by token for consistent lock ordering across pods to prevent deadlocks.
for token, response_cost in sorted(key_list_transactions.items()):
@ -2125,7 +2141,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with _spend_update_tx(prisma_client) as transaction:
async with _spend_update_tx(prisma_client, "LiteLLM_TeamTable") as transaction:
async with transaction.batch_() as batcher:
# Sort by team_id for consistent lock ordering across pods to prevent deadlocks.
for team_id, response_cost in sorted(team_list_transactions.items()):
@ -2163,7 +2179,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with _spend_update_tx(prisma_client) as transaction:
async with _spend_update_tx(prisma_client, "LiteLLM_TeamMembership") as transaction:
await _write_team_member_spend(transaction, team_member_list_transactions)
# Transaction succeeded, break out of retry loop
break
@ -2200,7 +2216,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with _spend_update_tx(prisma_client) as transaction:
async with _spend_update_tx(prisma_client, "LiteLLM_OrganizationTable") as transaction:
async with transaction.batch_() as batcher:
# Sort by org_id for consistent lock ordering across pods to prevent deadlocks.
for org_id, response_cost in sorted(org_list_transactions.items()):
@ -2226,7 +2242,10 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with _spend_update_tx(prisma_client) as transaction, transaction.batch_() as batcher:
async with (
_spend_update_tx(prisma_client, "LiteLLM_OrganizationMembership") as transaction,
transaction.batch_() as batcher,
):
for key, response_cost in sorted(org_member_list_transactions.items()):
_, quoted_org_id, _, quoted_user_id = key.split("::")
batcher.litellm_organizationmembership.update_many(
@ -2345,7 +2364,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with _spend_update_tx(prisma_client) as transaction:
async with _spend_update_tx(prisma_client, _ENTITY_SPEND_MODELS[table_accessor]) as transaction:
async with transaction.batch_() as batcher:
# Sort by entity_id for consistent lock ordering across pods to prevent deadlocks.
for entity_id, response_cost in sorted(transactions.items()):
@ -2512,7 +2531,7 @@ class DBSpendUpdateWriter:
table=table, transactions=tuple(transactions_to_process.values())
)
sql, params = build_bulk_upsert(table=table, batch=merged_batch)
async with _spend_update_tx(prisma_client) as transaction:
async with _spend_update_tx(prisma_client, table.name, "upsert_daily_spend") as transaction:
await transaction.execute_raw(sql, *params)
_mark_daily_spend_commit_started()
_mark_daily_spend_commit_finished()

View file

@ -20,6 +20,7 @@ from litellm.constants import (
SPEND_LOG_RUN_LOOPS,
)
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.db_transaction_queue.spend_log_cleanup_metrics import (
RunOutcome,
SpendLogCleanupMetrics,
@ -289,7 +290,7 @@ class SpendLogCleanup:
return remaining
async def _execute_delete_batch(
self, prisma_client: PrismaClient, delete_sql: str, cutoff_date: Cutoff, deadline: float
self, prisma_client: PrismaClient, delete_sql: str, cutoff_date: Cutoff, table_name: str, deadline: float
) -> int | None:
"""
Run one delete batch under a Postgres statement and lock timeout.
@ -305,7 +306,7 @@ class SpendLogCleanup:
fault, so the caller stops instead of retrying.
"""
timeout_ms: Final = self._timeout_ms(deadline)
async with prisma_client.db.tx() as tx:
async with db_span("cleanup_expired_rows", table_name), prisma_client.db.tx() as tx:
await tx.execute_raw(f"SET LOCAL statement_timeout = {timeout_ms}")
await tx.execute_raw(f"SET LOCAL lock_timeout = {timeout_ms}")
deleted_result: Final = await tx.execute_raw(delete_sql, cutoff_date, self.batch_size)
@ -330,7 +331,7 @@ class SpendLogCleanup:
) capped
"""
try:
async with prisma_client.db.tx() as tx:
async with db_span("count_expired_rows", table_name), prisma_client.db.tx() as tx:
await tx.execute_raw(f"SET LOCAL statement_timeout = {self._timeout_ms(deadline)}")
rows: Final = _REMAINING_ROWS.validate_python(
await tx.query_raw(count_sql, cutoff_date, SPEND_LOG_CLEANUP_REMAINING_COUNT_CAP)
@ -388,7 +389,9 @@ class SpendLogCleanup:
# Find rows and delete them in one go without fetching to application
batch_started_at = time.monotonic()
try:
batch_result = await self._execute_delete_batch(prisma_client, delete_sql, cutoff_date, deadline)
batch_result = await self._execute_delete_batch(
prisma_client, delete_sql, cutoff_date, table_name, deadline
)
except Exception as batch_exc:
if time.monotonic() >= deadline:
# The statement timeout was clamped to the budget that was

View file

@ -28,6 +28,7 @@ from litellm.constants import (
SPEND_LOG_PARTITION_INTERVAL,
SPEND_LOG_PARTITION_PRECREATE_AHEAD,
)
from litellm.proxy.db.db_span import db_span
if TYPE_CHECKING:
from prisma.client import TransactionManager
@ -159,7 +160,10 @@ class SpendLogsPartitionManager:
if budget_ms is None:
return False
try:
async with _bounded_tx(prisma_client, budget_ms) as tx:
async with (
db_span("check_spend_log_partitioning", "LiteLLM_SpendLogs"),
_bounded_tx(prisma_client, budget_ms) as tx,
):
await tx.execute_raw(f"SET LOCAL statement_timeout = {budget_ms}")
rows: Final = await tx.query_raw(
"""
@ -194,7 +198,10 @@ class SpendLogsPartitionManager:
wait for the lock and statement_timeout bounds the work itself, so a
partition this run cannot get is simply left for the next one.
"""
async with _bounded_tx(prisma_client, timeout_ms) as tx:
async with (
db_span("create_spend_log_partition", "LiteLLM_SpendLogs"),
_bounded_tx(prisma_client, timeout_ms) as tx,
):
await tx.execute_raw(f"SET LOCAL statement_timeout = {timeout_ms}")
await tx.execute_raw(f"SET LOCAL lock_timeout = {timeout_ms}")
await tx.execute_raw(statement)
@ -231,7 +238,10 @@ class SpendLogsPartitionManager:
async def _list_partitions(
self, prisma_client: "PrismaClient", timeout_ms: int
) -> list[tuple[str, datetime | None]]:
async with _bounded_tx(prisma_client, timeout_ms) as tx:
async with (
db_span("list_spend_log_partitions", "LiteLLM_SpendLogs"),
_bounded_tx(prisma_client, timeout_ms) as tx,
):
await tx.execute_raw(f"SET LOCAL statement_timeout = {timeout_ms}")
rows: Final = await tx.query_raw(
"""

View file

@ -32,6 +32,7 @@ from litellm._internal_context import with_service_target
from litellm._logging import verbose_proxy_logger
from litellm.caching import RedisCache
from litellm.constants import MAX_REDIS_BUFFER_DEQUEUE_COUNT, REDIS_GATEWAY_REQUESTS_BUFFER_KEY
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.middleware.billable_request_metrics_middleware import BillableCategory
from litellm.types.proxy.gateway_requests import (
@ -146,7 +147,8 @@ async def commit_gateway_requests_to_db(
return
sql, params = build_gateway_requests_upsert(snapshot)
await prisma_client.db.execute_raw(sql, *params) # pyright: ignore[reportAny] # untyped prisma client
async with db_span("commit_gateway_requests", "LiteLLM_DailyGatewayRequests"):
await prisma_client.db.execute_raw(sql, *params) # pyright: ignore[reportAny] # untyped prisma client
verbose_proxy_logger.debug(
"Gateway request tracking - committed %d aggregated rows in one statement", len(snapshot)

View file

@ -17,6 +17,7 @@ from typing import TYPE_CHECKING, Final
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, field_validator
from litellm._logging import verbose_proxy_logger
from litellm.proxy.db.db_span import db_span
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
@ -75,7 +76,8 @@ _ROWS_ADAPTER: Final = TypeAdapter(tuple[LatestHealthCheckRow, ...])
async def query_latest_health_checks(prisma_client: PrismaClient) -> tuple[LatestHealthCheckRow, ...]:
rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_SQL)
async with db_span("latest_health_checks", "LiteLLM_HealthCheckTable"):
rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_SQL)
return _ROWS_ADAPTER.validate_python(rows)
@ -93,7 +95,8 @@ async def fetch_latest_health_checks_for_models(
if not model_names:
return ()
try:
rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_FOR_MODELS_SQL, list(model_names))
async with db_span("latest_health_checks", "LiteLLM_HealthCheckTable"):
rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_FOR_MODELS_SQL, list(model_names))
return _ROWS_ADAPTER.validate_python(rows)
except Exception as query_err: # noqa: BLE001 # a paged model list must not fail on its health decoration
verbose_proxy_logger.error("Error getting latest health checks for models: %s", query_err)

View file

@ -5,48 +5,73 @@ ServiceLogger() then sends DB logs to Prometheus, OTEL, Datadog etc
"""
import asyncio
from collections.abc import Callable
from collections.abc import Callable, Generator, Mapping
from contextlib import contextmanager
from contextvars import ContextVar
from datetime import datetime
from functools import wraps
from types import MappingProxyType
from typing import Final
from litellm._service_logger import ServiceTypes
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
_PRISMA_CLIENT_CRUD: Final = frozenset({"get_data", "update_data", "delete_data"})
_DEFAULT_TABLE_BY_KWARG: Final[Mapping[str, str]] = MappingProxyType(
{"token": "key", "tokens": "key", "user_id": "user", "team_id": "team"}
)
def _safe_db_event_metadata(kwargs: dict) -> dict[str, str] | None:
def _table_metadata_reader(infers_table: bool) -> Callable[[Mapping[str, object]], dict[str, str] | None]:
"""Minimal, non-sensitive ``event_metadata`` for a DB service log.
The raw ``kwargs``/``args`` carry live objects (Prisma client, OTel spans)
and secrets (tokens), none of which belongs on a span — so we surface only
the table name when present. Everything else is dropped.
and secrets (tokens), none of which belongs on a span, so only the table name
surfaces. A ``PrismaClient`` CRUD method called without ``table_name`` picks
its table from the lookup key, in the same order the method dispatches on.
"""
table_name: Final = kwargs.get("table_name")
return {"table_name": table_name} if isinstance(table_name, str) else None
def read(kwargs: Mapping[str, object]) -> dict[str, str] | None:
table_name: Final = kwargs.get("table_name")
if isinstance(table_name, str):
return {"table_name": table_name}
if not infers_table:
return None
inferred: Final = next(
(table for key, table in _DEFAULT_TABLE_BY_KWARG.items() if kwargs.get(key) is not None), None
)
return {"table_name": inferred} if inferred is not None else None
return read
class _DbIoWitness:
"""Activity an inner decorated call already reported stays with it; only unreported activity reaches the enclosing call."""
__slots__ = ("_parent", "_reported", "_touched")
__slots__ = ("_open", "_parent", "_reported", "_touched")
def __init__(self, parent: "_DbIoWitness | None") -> None:
self._parent: Final = parent
self._touched = False
self._reported = False
self._open = True
@property
def touched(self) -> bool:
return self._touched
@property
def open(self) -> bool:
return self._open
def mark(self) -> None:
self._touched = True
if self._open:
self._touched = True
def report(self) -> None:
self._reported = True
def close(self) -> None:
self._open = False
if self._touched and not self._reported and self._parent is not None:
self._parent.mark()
@ -60,13 +85,36 @@ def record_db_io() -> None:
witness.mark()
def db_io_claimed() -> bool:
"""Whether an enclosing producer (``@log_db_metrics`` or ``db_span``) reports the Prisma I/O that runs now.
A task spawned inside a producer inherits its witness by context copy; once that producer has
finished, the copy is closed and the task's own Prisma I/O is nobody's to report but its own.
"""
witness: Final = _db_io_witness.get()
return witness is not None and witness.open
@contextmanager
def claim_db_io() -> Generator[_DbIoWitness]:
"""Own the DB event for the Prisma I/O inside: inner producers and the engine fallback stay quiet."""
witness: Final = _DbIoWitness(parent=_db_io_witness.get())
token: Final = _db_io_witness.set(witness)
try:
yield witness
finally:
witness.report()
witness.close()
_db_io_witness.reset(token)
def log_db_metrics(func):
"""
Decorator to log the duration of a DB related function to ServiceLogger()
Handles logging DB success/failure to ServiceLogger(), which logs to Prometheus, OTEL, Datadog
When logging Failure it checks if the Exception is a PrismaError, httpx.ConnectError or httpx.TimeoutException and then logs that as a DB Service Failure
When logging Failure it checks if the Exception is a PrismaError or an httpx.TransportError and then logs that as a DB Service Failure
Args:
func: The function to be decorated
@ -78,8 +126,10 @@ def log_db_metrics(func):
Exception: If the decorated function raises an exception
"""
metadata_of: Final = _table_metadata_reader(func.__name__ in _PRISMA_CLIENT_CRUD)
@wraps(func)
async def wrapper(*args, **kwargs):
async def wrapper(*args, **kwargs: object):
start_time: Final[datetime] = datetime.now()
witness: Final = _DbIoWitness(parent=_db_io_witness.get())
witness_token: Final = _db_io_witness.set(witness)
@ -89,44 +139,20 @@ def log_db_metrics(func):
end_time: datetime = datetime.now()
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,
call_type=func.__name__,
parent_otel_span=kwargs.get("parent_otel_span", None),
duration=(end_time - start_time).total_seconds(),
start_time=start_time,
end_time=end_time,
event_metadata=_safe_db_event_metadata(kwargs),
)
if not witness.touched:
return result
asyncio.create_task(
proxy_logging_obj.service_logging_obj.async_service_success_hook(
service=ServiceTypes.DB,
call_type=func.__name__,
parent_otel_span=kwargs.get("parent_otel_span", None),
duration=(end_time - start_time).total_seconds(),
start_time=start_time,
end_time=end_time,
event_metadata=metadata_of(kwargs),
)
witness.report()
elif (
# in litellm custom callbacks kwargs is passed as arg[0]
# https://docs.litellm.ai/docs/observability/custom_callback#callback-functions
args is not None and len(args) > 1 and isinstance(args[1], dict)
):
passed_kwargs: Final = args[1]
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs=passed_kwargs)
if parent_otel_span is not None:
# No metadata dump: identity rides on Baggage, and the full
# request metadata (auth blob, response headers, tokens) must
# not land on a span.
asyncio.create_task(
proxy_logging_obj.service_logging_obj.async_service_success_hook(
service=ServiceTypes.BATCH_WRITE_TO_DB,
call_type=func.__name__,
parent_otel_span=parent_otel_span,
duration=0.0,
start_time=start_time,
end_time=end_time,
event_metadata=None,
)
)
# end of logging to otel
)
witness.report()
return result
except Exception as e:
end_time: datetime = datetime.now()
@ -137,6 +163,7 @@ def log_db_metrics(func):
args=args,
start_time=start_time,
end_time=end_time,
metadata_of=metadata_of,
):
witness.report()
raise e
@ -155,16 +182,17 @@ def _is_exception_related_to_db(e: Exception) -> bool:
import httpx
from prisma.errors import PrismaError
return isinstance(e, (PrismaError, httpx.ConnectError, httpx.TimeoutException))
return isinstance(e, (PrismaError, httpx.TransportError))
async def _handle_logging_db_exception(
e: Exception,
func: Callable,
kwargs: dict,
kwargs: Mapping[str, object],
args: tuple,
start_time: datetime,
end_time: datetime,
metadata_of: Callable[[Mapping[str, object]], dict[str, str] | None],
) -> bool:
from litellm.proxy.proxy_server import proxy_logging_obj
@ -180,6 +208,6 @@ async def _handle_logging_db_exception(
duration=(end_time - start_time).total_seconds(),
start_time=start_time,
end_time=end_time,
event_metadata=_safe_db_event_metadata(kwargs),
event_metadata=metadata_of(kwargs),
)
return True

View file

@ -16,8 +16,10 @@ from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Final, Protocol
from litellm._logging import verbose_proxy_logger
from litellm.proxy.db.db_span import db_span
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.log_db_metrics import db_io_claimed, record_db_io
from litellm.proxy.db.prisma_query_span import parse_prisma_query
from litellm.proxy.db.token_auth import (
DEFAULT_POSTGRES_PORT,
DatabaseTokenAuth,
@ -107,10 +109,15 @@ 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)
if db_io_claimed():
record_db_io()
return await self._engine.query(content, tx_id=tx_id)
query: Final = parse_prisma_query(content)
async with db_span(query.call_type, query.table, query.operation):
record_db_io()
return await self._engine.query(content, tx_id=tx_id)
finally:
self.tracker.end_operation()

View file

@ -0,0 +1,136 @@
"""Name the Prisma round trips that no producer claims.
``_TrackedPrismaEngine.query`` sees every statement the proxy sends to the query
engine. When neither ``@log_db_metrics`` nor ``db_span`` encloses the call, the
engine names the event itself from the GraphQL payload Prisma built: the root
field (``findUniqueLiteLLM_VerificationToken``, ``createOneLiteLLM_SpendLogs``)
carries the method and the model, and for ``queryRaw``/``executeRaw`` the leading
SQL keyword gives the verb and the first ``schema.prisma`` relation the statement
names gives the table. Only bounded names ever leave this module: relations
declared in the schema, the spend views, ``pg_catalog`` for catalog probes and
the setting a ``SET`` statement targets. No SQL text or values.
"""
from __future__ import annotations
import json
import re
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
from litellm.integrations.otel.model.spans import PG_CATALOG, PRISMA_RELATIONS
@dataclass(frozen=True, slots=True)
class PrismaQuery:
"""What one engine round trip is, for the ``ServiceTypes.DB`` event: the raw method,
the SQL verb (``None`` when the statement is not one this module knows) and the relation."""
call_type: str
operation: str | None
table: str | None
UNKNOWN_PRISMA_QUERY: Final = PrismaQuery("prisma_query", None, None)
_ROOT_FIELD: Final = re.compile(r"result:\s*(\w+)")
_RAW_SQL: Final = re.compile(r'query:\s*"((?:[^"\\]|\\.)*)"')
_LEADING_KEYWORD: Final = re.compile(r"(?:\\[nrt]|\s|\()*(\w+)")
_SETTING: Final = re.compile(r"(?:\\[nrt]|\s)*SET\s+(?:LOCAL\s+|SESSION\s+)?([A-Za-z_.]+)", re.IGNORECASE)
_CATALOG: Final = re.compile(r"\bpg_\w+|\bto_regclass\b|\binformation_schema\b|\bcurrent_setting\s*\(|^\s*SHOW\b")
_PROBE: Final = re.compile(r"(?:\\[nrt]|\s)*SELECT\s+\d+\s*;?(?:\\[nrt]|\s)*$", re.IGNORECASE)
_CTE_WRITE: Final = re.compile(r"\b(UPDATE|INSERT|DELETE)\s+(?:INTO\s+|FROM\s+)?(?:\\?\")", re.IGNORECASE)
_RELATION: Final = re.compile(
r"\b(?:" + "|".join(sorted(map(re.escape, PRISMA_RELATIONS), key=len, reverse=True)) + r")\b"
)
_MODEL_ACTIONS: Final[Mapping[str, tuple[str, str]]] = MappingProxyType(
{
"findUnique": ("find_unique", "select"),
"findFirst": ("find_first", "select"),
"findMany": ("find_many", "select"),
"aggregate": ("count", "select"),
"groupBy": ("group_by", "select"),
"createOne": ("create", "insert"),
"createMany": ("create_many", "insert"),
"updateOne": ("update", "update"),
"updateMany": ("update_many", "update"),
"deleteOne": ("delete", "delete"),
"deleteMany": ("delete_many", "delete"),
"upsertOne": ("upsert", "upsert"),
}
)
_RAW_ACTIONS: Final[Mapping[str, str]] = MappingProxyType({"queryRaw": "query_raw", "executeRaw": "execute_raw"})
_VERB_BY_KEYWORD: Final[Mapping[str, str]] = MappingProxyType(
{
"SELECT": "select",
"WITH": "select",
"INSERT": "insert",
"UPDATE": "update",
"DELETE": "delete",
"CREATE": "ddl",
"ALTER": "ddl",
"DROP": "ddl",
"REFRESH": "ddl",
"TRUNCATE": "delete",
"SET": "set",
}
)
def sql_relation(sql: str) -> str | None:
"""The first schema relation (model or spend view) the statement names, ``pg_catalog``
for a statement that only reads the system catalog, else ``None``."""
relation: Final = _RELATION.search(sql)
if relation is not None:
return relation.group(0)
return PG_CATALOG if _CATALOG.search(sql) else None
def sql_operation(sql: str) -> tuple[str | None, str | None]:
"""``(verb, target)`` for a raw statement: the SQL verb from its leading keyword and the
relation it names, or for ``SET`` the setting it changes."""
if _PROBE.match(sql):
return "ping", None
keyword: Final = _LEADING_KEYWORD.match(sql)
leading: Final = keyword.group(1).upper() if keyword is not None else ""
cte_write: Final = _CTE_WRITE.search(sql) if leading == "WITH" else None
verb: Final = _VERB_BY_KEYWORD[cte_write.group(1).upper()] if cte_write else _VERB_BY_KEYWORD.get(leading)
if verb != "set":
return verb, sql_relation(sql)
setting: Final = _SETTING.match(sql)
return verb, setting.group(1).lower() if setting is not None else None
def _query_text(content: str) -> str:
try:
payload: Final[object] = json.loads(content)
except ValueError:
return content
query: Final = payload.get("query") if isinstance(payload, dict) else None
return query if isinstance(query, str) else content
def _model_query(root_field: str) -> PrismaQuery | None:
action: Final = next((prefix for prefix in _MODEL_ACTIONS if root_field.startswith(prefix)), None)
if action is None:
return None
call_type, verb = _MODEL_ACTIONS[action]
model: Final = root_field.removeprefix(action).removesuffix("OrThrow")
return PrismaQuery(call_type, verb, model) if model in PRISMA_RELATIONS else None
def parse_prisma_query(content: str) -> PrismaQuery:
"""The round trip behind one query-engine payload, ``UNKNOWN_PRISMA_QUERY`` when the
payload is not a shape this module knows (which renders ``postgres prisma_query``)."""
query: Final = _query_text(content)
root: Final = _ROOT_FIELD.search(query)
if root is None:
return UNKNOWN_PRISMA_QUERY
raw_call_type: Final = _RAW_ACTIONS.get(root.group(1))
if raw_call_type is None:
return _model_query(root.group(1)) or UNKNOWN_PRISMA_QUERY
sql: Final = _RAW_SQL.search(query, root.end())
verb, target = sql_operation(sql.group(1)) if sql is not None else (None, None)
return PrismaQuery(raw_call_type, verb, target)

View file

@ -20,6 +20,7 @@ from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
if TYPE_CHECKING:
@ -65,14 +66,17 @@ class ProxyWorkerHeartbeat:
async def beat(self) -> None:
try:
await self.prisma_client.db.execute_raw(BEAT_SQL, self.worker_id, self.hostname)
await self.prisma_client.db.execute_raw(PRUNE_SQL, STALE_ROW_RETENTION_SECONDS)
async with db_span("proxy_worker_heartbeat", "LiteLLM_ProxyWorkerHeartbeat"):
await self.prisma_client.db.execute_raw(BEAT_SQL, self.worker_id, self.hostname)
async with db_span("prune_proxy_worker_heartbeats", "LiteLLM_ProxyWorkerHeartbeat"):
await self.prisma_client.db.execute_raw(PRUNE_SQL, STALE_ROW_RETENTION_SECONDS)
except Exception as beat_err: # noqa: BLE001 # a missed heartbeat must never take down the worker
verbose_proxy_logger.debug("Proxy worker heartbeat write failed: %s", beat_err)
async def deregister(self) -> None:
try:
await self.prisma_client.db.execute_raw(DEREGISTER_SQL, self.worker_id)
async with db_span("deregister_proxy_worker", "LiteLLM_ProxyWorkerHeartbeat"):
await self.prisma_client.db.execute_raw(DEREGISTER_SQL, self.worker_id)
except Exception as deregister_err: # noqa: BLE001 # best-effort cleanup; the liveness window ages the row out anyway
verbose_proxy_logger.debug("Proxy worker heartbeat deregister failed: %s", deregister_err)
@ -86,7 +90,8 @@ async def count_live_proxy_workers(prisma_client: PrismaClient) -> int | None:
try:
db: Final = prisma_client.db
primary_db: Final = db.writer if isinstance(db, RoutingPrismaWrapper) else db
rows: Final = await primary_db.query_raw(COUNT_SQL, PROXY_WORKER_LIVENESS_WINDOW_SECONDS)
async with db_span("count_live_proxy_workers", "LiteLLM_ProxyWorkerHeartbeat"):
rows: Final = await primary_db.query_raw(COUNT_SQL, PROXY_WORKER_LIVENESS_WINDOW_SECONDS)
return _COUNT_ROWS_ADAPTER.validate_python(rows)[0]["live_workers"]
except Exception as count_err: # noqa: BLE001 # an unknown count must degrade to "warn", never to a 503
verbose_proxy_logger.debug("Live proxy worker count unavailable: %s", count_err)

View file

@ -11,6 +11,7 @@ than an undercount (same call as the auto-router session rollup flush).
from typing import TYPE_CHECKING, Final, Literal
from litellm._logging import verbose_proxy_logger
from litellm.proxy.db.db_span import db_span
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
@ -51,11 +52,12 @@ async def flush_shadow_eval_funnel(prisma_client: "PrismaClient") -> None:
_pending.clear()
for job_id, counters in batch.items():
try:
await prisma_client.db.execute_raw(
_UPSERT_FUNNEL_SQL,
job_id,
*(counters[stage] for stage in FUNNEL_STAGES),
)
async with db_span("flush_shadow_eval_funnel", "LiteLLM_ShadowEvalFunnel"):
await prisma_client.db.execute_raw(
_UPSERT_FUNNEL_SQL,
job_id,
*(counters[stage] for stage in FUNNEL_STAGES),
)
except Exception as flush_err: # noqa: BLE001 # drop this leg's batch: a repeated increment is worse than an undercount
verbose_proxy_logger.error(
"Spend tracking - shadow eval funnel flush failed for job %s, %s dropped: %s",

View file

@ -23,6 +23,7 @@ from typing import TYPE_CHECKING, Any, Final
from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS
from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.spend_log_batching import spend_log_write_batches
from litellm.repositories.table_repositories import SpendLogToolIndexRepository
@ -135,8 +136,12 @@ async def flush_tool_usage_transactions(
for statement_rows in spend_log_write_batches(
index_rows, SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS
):
await index_table.create_many(data=statement_rows, skip_duplicates=True)
async with prisma_client.db.batch_() as batcher:
async with db_span("index_spend_log_tools", "LiteLLM_SpendLogToolIndex"):
await index_table.create_many(data=statement_rows, skip_duplicates=True)
async with (
db_span("commit_daily_tool_spend", "LiteLLM_DailyToolSpend"),
prisma_client.db.batch_() as batcher,
):
for (date_key, tool_name), grouped in groupby(per_tool_day, key=lambda entry: (entry[0], entry[1])):
entries = tuple(grouped)
spend = sum(entry[2] for entry in entries)

View file

@ -21,7 +21,6 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import (
get_key_object,
get_team_object,
log_db_metrics,
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded
@ -289,7 +288,6 @@ class _ProxyDBLogger(CustomLogger):
project_id=user_api_key_dict.project_id,
)
@log_db_metrics
async def _PROXY_track_cost_callback(
self,
kwargs, # kwargs to completion

View file

@ -37,6 +37,8 @@ from litellm.proxy.db.autorouter_session_rollup import (
AUTOROUTER_BENCHMARKS_SQL,
bounded_session_id,
)
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.prisma_query_span import sql_relation
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
refresh_proxy_server_request_body_snapshot,
@ -209,7 +211,8 @@ def _shadow_eval_attempts(prisma_client: "PrismaClient") -> _ShadowEvalAttemptTa
async def _query_raw(prisma_client: "PrismaClient", query: str, *args: object) -> Sequence[Mapping[str, object]]:
return await prisma_client.db.query_raw(query, *args)
async with db_span("auto_router_report_query", sql_relation(query)):
return await prisma_client.db.query_raw(query, *args)
async def _authorize_router_dry_run(user_api_key_dict: UserAPIKeyAuth, team_id: str | None) -> LiteLLM_TeamTable | None:

View file

@ -35,6 +35,7 @@ from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Final, Literal, cast
from litellm._logging import verbose_proxy_logger
from litellm.proxy.db.db_span import db_span
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
@ -261,10 +262,11 @@ async def _migrate_config_settings_row(
report.plaintext += 1
if changed and not dry_run:
await prisma_client.db.litellm_config.update(
where={"param_name": param_name},
data={"param_value": json.dumps(settings)},
)
async with db_span("migrate_config_credentials", "LiteLLM_Config"):
await prisma_client.db.litellm_config.update(
where={"param_name": param_name},
data={"param_value": json.dumps(settings)},
)
return report
@ -313,10 +315,11 @@ async def _migrate_sso_config(prisma_client: object, dry_run: bool) -> LocationR
report.plaintext += 1
if changed and not dry_run:
await prisma_client.db.litellm_ssoconfig.update(
where={"id": "sso_config"},
data={"sso_settings": json.dumps(new_settings)},
)
async with db_span("migrate_sso_credentials", "LiteLLM_SSOConfig"):
await prisma_client.db.litellm_ssoconfig.update(
where={"id": "sso_config"},
data={"sso_settings": json.dumps(new_settings)},
)
return report

View file

@ -117,6 +117,7 @@ from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.db.db_span import db_span
from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks
from litellm.proxy.hooks.model_max_budget_limiter import (
build_model_max_budget_usage,
@ -6794,13 +6795,14 @@ async def get_team_spend_by_user(
own_user_only: Final = scope.api_key_filter is not None
user_param: Final = (user_api_key_dict.user_id or "",) if own_user_only else ()
rows: Final[Sequence[_TeamUserSpendDbRow]] = await prisma_client.db.query_raw(
_team_user_spend_sql(team_count=len(scoped_team_ids), restrict_to_user=own_user_only),
start_date,
end_date,
*scoped_team_ids,
*user_param,
)
async with db_span("team_user_spend", "LiteLLM_SpendLogs"):
rows: Final[Sequence[_TeamUserSpendDbRow]] = await prisma_client.db.query_raw(
_team_user_spend_sql(team_count=len(scoped_team_ids), restrict_to_user=own_user_only),
start_date,
end_date,
*scoped_team_ids,
*user_param,
)
results: Final = tuple(
TeamUserSpendRow(
team_id=row["team_id"],

View file

@ -21,6 +21,7 @@ from typing import Final, Protocol
from pydantic import BaseModel, TypeAdapter
from litellm.proxy.auth.auth_checks import _delete_cache_access_object
from litellm.proxy.db.db_span import db_span
# hashtext collisions only cost two unrelated teams a little serialization, and the
# lock is never taken by the access-group endpoints as a SELECT ... FOR UPDATE row lock,
@ -151,7 +152,7 @@ async def reconcile_team_access_group_membership(tx: AccessGroupSyncTx, team_id:
async def sync_team_access_group_membership(prisma_client: _PrismaClient, team_id: str) -> None:
"""Reconcile the mirror for an already committed team write, in its own transaction."""
async with prisma_client.db.tx() as tx:
async with db_span("sync_team_access_group_membership", "LiteLLM_AccessGroupTable"), prisma_client.db.tx() as tx:
affected: Final = await reconcile_team_access_group_membership(tx, team_id)
await invalidate_access_group_caches(affected)

View file

@ -20,6 +20,7 @@ from litellm.constants import (
SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE,
)
from litellm.litellm_core_utils.litellm_logging import is_valid_sha256_hash
from litellm.proxy.db.db_span import db_span, db_spanned
from litellm.proxy.utils import PrismaClient
from litellm.repositories.chunked_in import find_many_in
from litellm.repositories.user_repository import UserRepository
@ -184,9 +185,13 @@ async def _rows_within_the_statement_timeout(
prisma_client: PrismaClient,
sql: str,
*params: object,
table: str,
planner_settings: tuple[str, ...] = (),
) -> Sequence[Mapping[str, object]]:
async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction:
async with (
db_span("recover_key_metadata", table),
prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction,
):
await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL)
for setting in planner_settings:
await transaction.execute_raw(setting)
@ -198,10 +203,11 @@ async def _reverse_hash_key_metadata(
sql: str,
wanted: AbstractSet[str],
*,
table: str,
warning: str,
) -> Mapping[str, KeyMetadataDict]:
rows: Final = await _db_or_empty(
lambda: prisma_client.db.query_raw(sql, sorted(wanted)),
lambda: db_spanned("recover_key_metadata", table, lambda: prisma_client.db.query_raw(sql, sorted(wanted))),
warning,
len(wanted),
)
@ -223,7 +229,9 @@ async def recover_key_owner_from_daily_spend(
if not keys:
return _EMPTY_KEY_OWNERS
rows: Final = await _db_or_empty(
lambda: _rows_within_the_statement_timeout(prisma_client, _DAILY_USER_SPEND_OWNER_SQL, sorted(keys)),
lambda: _rows_within_the_statement_timeout(
prisma_client, _DAILY_USER_SPEND_OWNER_SQL, sorted(keys), table="LiteLLM_DailyUserSpend"
),
"Failed daily-spend key owner recovery for %d keys: %s",
len(keys),
)
@ -255,7 +263,11 @@ async def _details_for_user_ids(
if not user_ids:
return _EMPTY_USER_DETAILS
users: Final = await _db_or_empty(
lambda: find_many_in(UserRepository(prisma_client).table, "user_id", user_ids),
lambda: db_spanned(
"recover_user_details",
"LiteLLM_UserTable",
lambda: find_many_in(UserRepository(prisma_client).table, "user_id", user_ids),
),
"Failed user detail recovery for %d user ids: %s",
len(user_ids),
)
@ -358,6 +370,7 @@ async def recover_double_hashed_key_metadata(
prisma_client,
_ACTIVE_TOKEN_DIGEST_SQL,
sha_missing,
table="LiteLLM_VerificationToken",
warning="Failed reverse-hash recovery against active keys for %d missing keys: %s",
)
still_missing: Final = sha_missing - frozenset(from_active)
@ -367,6 +380,7 @@ async def recover_double_hashed_key_metadata(
prisma_client,
_DELETED_TOKEN_DIGEST_SQL,
still_missing,
table="LiteLLM_DeletedVerificationToken",
warning="Failed reverse-hash recovery against deleted keys for %d missing keys: %s",
)
return MappingProxyType({**from_active, **from_deleted})
@ -409,6 +423,7 @@ async def _query_spend_log_metadata(
sorted(digests),
start,
end,
table="LiteLLM_SpendLogs",
planner_settings=(_SPEND_LOG_NO_BITMAP_SCAN_SQL,),
),
"Failed spend-log alias recovery for %d missing keys: %s",

View file

@ -171,6 +171,7 @@ from litellm.proxy.db.create_views import (
create_view_tolerating_race,
should_create_missing_views,
)
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
from litellm.proxy.db.db_url_settings import (
DatabaseURLSettings,
@ -186,7 +187,7 @@ from litellm.proxy.db.health_check_latest import (
fetch_latest_health_checks,
fetch_latest_health_checks_for_models,
)
from litellm.proxy.db.log_db_metrics import log_db_metrics
from litellm.proxy.db.log_db_metrics import _is_exception_related_to_db, log_db_metrics
from litellm.proxy.db.pgbouncer import database_url_is_pooled
from litellm.proxy.db.prisma_client import (
PrismaWrapper,
@ -1082,6 +1083,7 @@ def _call_type_for_route(route: str | None) -> str | None:
_PROXY_ONLY_LLM_API_ERRORS: Final = (HTTPException, ProxyException, GuardrailRaisedException)
_LOG_DB_METRICS_CALL_TYPES: Final = frozenset(("get_data", "insert_data", "update_data", "delete_data"))
def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, object]:
@ -3184,7 +3186,10 @@ class ProxyLogging:
)
)
if hasattr(self, "service_logging_obj"):
logged_by_decorator: Final = call_type in _LOG_DB_METRICS_CALL_TYPES and _is_exception_related_to_db(
original_exception
)
if hasattr(self, "service_logging_obj") and not logged_by_decorator:
await self.service_logging_obj.async_service_failure_hook(
service=ServiceTypes.DB,
duration=duration,
@ -5289,6 +5294,7 @@ class PrismaClient:
max_time=10, # maximum total time to retry for
on_backoff=on_backoff, # specifying the function to call on backoff
)
@log_db_metrics
async def insert_data(
self,
data: Mapping[str, object],
@ -5438,6 +5444,7 @@ class PrismaClient:
max_time=10, # maximum total time to retry for
on_backoff=on_backoff, # specifying the function to call on backoff
)
@log_db_metrics
async def update_data(
self,
token: str | None = None,
@ -5677,6 +5684,7 @@ class PrismaClient:
max_time=10, # maximum total time to retry for
on_backoff=on_backoff, # specifying the function to call on backoff
)
@log_db_metrics
async def delete_data(
self,
tokens: Sequence[str | None] | None = None,
@ -6655,10 +6663,11 @@ class PrismaClient:
while True:
try:
await asyncio.sleep(self._db_health_watchdog_interval_seconds)
await asyncio.wait_for(
self.db.query_raw("SELECT 1"),
timeout=self._db_health_watchdog_probe_timeout_seconds,
)
async with db_span("db_health_watchdog", None):
await asyncio.wait_for(
self.db.query_raw("SELECT 1"),
timeout=self._db_health_watchdog_probe_timeout_seconds,
)
if isinstance(self.db, RoutingPrismaWrapper) and self.db.writer_unavailable:
await self.attempt_db_reconnect(
reason="db_health_watchdog_writer_unavailable",
@ -6742,7 +6751,8 @@ class PrismaClient:
about to check, and attribute the failure to the wrong replacement.
"""
sql_query: Final = "SELECT 1"
response: Final[object] = await wrapper.query_raw(sql_query)
async with db_span("health_check", None):
response: Final[object] = await wrapper.query_raw(sql_query)
return response
async def _probe_answers_now(self, wrapper: PrismaWrapper) -> bool:
@ -7287,7 +7297,10 @@ class ProxyUpdateSpend:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with prisma_client.db.tx(timeout=timedelta(seconds=60)) as transaction:
async with (
db_span("update_end_user_spend", "LiteLLM_EndUserTable"),
prisma_client.db.tx(timeout=timedelta(seconds=60)) as transaction,
):
batcher: _EndUserSpendBatch
async with transaction.batch_() as batcher:
# Sort by end_user_id for consistent lock ordering across pods to prevent deadlocks.
@ -7364,11 +7377,12 @@ class ProxyUpdateSpend:
SPEND_LOG_WRITE_BATCH_MAX_BYTES,
SPEND_LOG_WRITE_BATCH_MAX_ROWS,
):
isolation_budget = await _create_spend_logs_with_poison_isolation(
SpendLogsRepository(prisma_client),
statement_rows,
isolation_budget,
)
async with db_span("insert_spend_logs", "LiteLLM_SpendLogs"):
isolation_budget = await _create_spend_logs_with_poison_isolation(
SpendLogsRepository(prisma_client),
statement_rows,
isolation_budget,
)
verbose_proxy_logger.debug("Flushed %s logs to the DB.", len(batch))
# Explicitly clear batch memory
del batch, batch_with_dates

View file

@ -10,6 +10,8 @@ from typing_extensions import assert_never
from litellm import constants
from litellm._logging import verbose_proxy_logger
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.prisma_query_span import sql_relation
from litellm.repositories.chunked_in import find_many_in
from litellm.repositories.daily_activity_sql import (
ExportCursor,
@ -146,7 +148,10 @@ class DailyActivityRepository:
async def _query(self, query: SqlQuery) -> tuple[Mapping[str, object], ...]:
first_line: Final = query.sql.lstrip().splitlines()[0].lstrip("(").strip()
verbose_proxy_logger.debug("DailyActivityRepository query: %s", first_line)
result: Sequence[Mapping[str, object]] | None = await self._prisma_client.db.query_raw(query.sql, *query.params)
async with db_span("daily_activity_query", sql_relation(query.sql)):
result: Sequence[Mapping[str, object]] | None = await self._prisma_client.db.query_raw(
query.sql, *query.params
)
if result is None:
return ()
return tuple(result)

View file

@ -3,8 +3,7 @@
Covers logging.otel.success.exports_metric: a successful non-streaming call must
land at the OTEL destination as ONE connected trace - a single root SERVER span
with the auth phase and db lookups under it, the gen-AI CLIENT span parented
into the same tree, and the cost write either under it or as the root of its
own trace linked back to the request span. The regression this pins: the proxy publishing
into the same tree. The regression this pins: the proxy publishing
the global TracerProvider before callbacks init made server spans export through
a different provider than the preset's gen-AI spans, so the destination received
the gen-AI span alone, dangling (fixed in #30590; verified failing at its parent
@ -29,14 +28,13 @@ from e2e_config import CHEAP_ANTHROPIC_MODEL, CHEAP_OPENAI_MODEL, OTEL_EXPORTER_
from lifecycle import ResourceManager
from logging_client import INVALID_UPSTREAM_API_KEY, LoggingClient, first_ok, readiness_details_body
from models import LiteLLMParamsBody
from otel_client import CallTraces, JaegerSpan, JaegerTrace, OtelReader, root_span
from otel_client import CallTraces, JaegerSpan, JaegerTrace, OtelReader
from pydantic import BaseModel, ConfigDict, ValidationError
pytestmark = pytest.mark.e2e
MODEL = CHEAP_ANTHROPIC_MODEL
COST_SPAN = "batch_write_to_db _PROXY_track_cost_callback"
DB_SPAN_PREFIX = "postgres "
DB_SPAN_PREFIX = "postgres."
#: The active OTEL v2 logger's name in /health/readiness/details success_callbacks.
OTEL_V2_LOGGER_NAME = "OpenTelemetryV2"
@ -79,12 +77,12 @@ def _chain_reaches(span_id: str, root_id: str, trace: JaegerTrace) -> bool:
return False
def _assert_complete_trace(traces: CallTraces, *, route: str, genai_span: str, require_cost_span: bool = True) -> None:
def _assert_complete_trace(traces: CallTraces, *, route: str, genai_span: str) -> None:
"""The enforced behavior: the destination holds exactly one call-id-tagged
trace for the call, rooted at the SERVER span, with auth/db children and
the gen-AI span all connected into that one tree - no dangling parent
references - and the cost write either in that trace or as the root of
its own trace linked FOLLOWS_FROM to the request SERVER span."""
references. The spend enqueue after the response does no I/O, so it emits
no span; the flush that writes spend is its own background trace."""
hits = traces.hits
assert hits, (
"no trace for this call arrived at the destination within the deadline "
@ -121,21 +119,6 @@ def _assert_complete_trace(traces: CallTraces, *, route: str, genai_span: str, r
assert any(name.startswith(DB_SPAN_PREFIX) for name in names), (
f"no db ('{DB_SPAN_PREFIX}*') span in the trace; spans: {names}"
)
if require_cost_span and COST_SPAN not in names:
cost_traces = [t for t in traces.linked if (r := root_span(t)) is not None and r.operation_name == COST_SPAN]
assert len(cost_traces) == 1, (
f"cost write span {COST_SPAN!r} reached neither the request trace nor its own "
f"trace linked to the request SERVER span; request spans: {names}; "
f"linked traces: {[(t.trace_id, t.span_names()) for t in traces.linked]}"
)
cost_root = root_span(cost_traces[0])
assert cost_root is not None, f"cost write trace has no single root; spans: {cost_traces[0].span_names()}"
link = next(ref for ref in cost_root.references if ref.span_id == root.span_id)
assert link.ref_type == "FOLLOWS_FROM" and link.trace_id == trace.trace_id, (
f"the cost write trace's root must reference the request SERVER span FOLLOWS_FROM, "
f"got refType={link.ref_type!r} traceID={link.trace_id!r} (request trace {trace.trace_id})"
)
genai = next((span for span in trace.spans if span.operation_name == genai_span), None)
assert genai is not None, f"gen-AI span {genai_span!r} missing; spans: {names}"
assert genai.kind == "client", f"gen-AI span must have kind=client, got {genai.kind!r}"
@ -145,14 +128,11 @@ def _assert_complete_trace(traces: CallTraces, *, route: str, genai_span: str, r
)
def _poll(
otel_reader: OtelReader, *, call_id: str, route: str, genai_span: str, require_cost_span: bool = True
) -> CallTraces:
def _poll(otel_reader: OtelReader, *, call_id: str, route: str, genai_span: str) -> CallTraces:
return otel_reader.poll_traces_for_call(
call_id=call_id,
settled_names={f"POST {route}", f"auth {route}", genai_span},
settled_prefixes={DB_SPAN_PREFIX},
linked_names=frozenset({COST_SPAN}) if require_cost_span else frozenset(),
)
@ -308,8 +288,7 @@ class TestOtelTraceCompleteness:
The trace should have a single server root span for the incoming request, with
the authentication and database work beneath it. The span for the actual model
call must also belong to that same trace, rather than being exported separately
with a missing parent, and the cost-recording work must land either in that
trace or in its own trace linked to it.
with a missing parent.
This matters because a split trace is easy to miss: all of the spans may still
arrive, but the model call appears without the surrounding request context.
@ -367,8 +346,6 @@ class TestOtelTraceCompleteness:
The trace must have a single root span named "POST /v1/messages". The
authentication, database, and model-call spans must all belong to
the same trace and have valid parent relationships leading back to that root.
The cost-writing span must land in the request trace or in its own trace
linked to it.
The model-call span is expected to be named "chat <model>". The test fails if
the request is split across multiple traces, if any span references a missing
@ -397,13 +374,10 @@ class TestOtelTraceCompleteness:
The trace must have a single root span named "POST /v1/responses". The
authentication, database, and model-call spans must all belong to the same
trace and have valid parent relationships leading back to that root. The cost
write finishes after the response, so it lands as the root of its own trace
linked FOLLOWS_FROM to the request SERVER span.
trace and have valid parent relationships leading back to that root.
The model-call span is expected to be named "chat <model>". The test fails on
a split request trace, a dangling parent, a disconnected model-call span, or
a cost write that is neither in the request trace nor linked to it."""
a split request trace, a dangling parent or a disconnected model-call span."""
route = "/v1/responses"
_assert_otel_destination_configured(client)
@ -428,8 +402,7 @@ class TestOtelTraceCompleteness:
"""A successful streamed `/chat/completions` request should export one
complete OTEL trace. The trace must contain a single root `SERVER`
span, with the auth, database, and gen-AI `CLIENT` spans all
connected back to that root, and the cost write in that trace or in
its own trace linked to it.
connected back to that root.
Streaming has an additional lifecycle risk because the gen-AI span is
closed by the stream-consumption path after the final chunk has
@ -477,8 +450,7 @@ class TestOtelTraceCompleteness:
"""A successful streamed `/v1/messages` request should export one
complete OTEL trace. The trace must contain a single root `SERVER`
span, with the auth, database, and gen-AI `CLIENT` spans all
connected back to that root, and the cost write in that trace or in
its own trace linked to it.
connected back to that root.
This endpoint has the same streaming lifecycle risk as
`/chat/completions`: the gen-AI span is closed by the
@ -558,18 +530,14 @@ class TestOtelTraceCompleteness:
)
genai_span = f"chat {CHEAP_OPENAI_MODEL}"
traces = _poll(
otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False
)
_assert_complete_trace(traces, route=route, genai_span=genai_span, require_cost_span=False)
traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span)
_assert_complete_trace(traces, route=route, genai_span=genai_span)
one_served_genai_span(traces.hits[0], genai_span)
spend_row = client.poll_proxy_spend_for_key(key)
assert spend_row is not None and spend_row.spend is not None and spend_row.spend > 0, (
"a successful streamed responses call must record a positive-spend row in /spend/logs "
"(the cost-write SPAN is knowingly absent on this surface, LIT-4428, but the spend "
f"itself must land); got {spend_row!r}"
f"a successful streamed responses call must record a positive-spend row in /spend/logs; got {spend_row!r}"
)
assert spend_row.call_type == "aresponses", (
f"the spend row must be attributed to the responses call type, got {spend_row.call_type!r}"
@ -686,9 +654,7 @@ class TestOtelTraceCompleteness:
)
genai_span = f"chat {CHEAP_OPENAI_MODEL}"
traces = _poll(
otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False
)
traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span)
_assert_real_ttft(traces.hits, genai_span=genai_span)
@pytest.mark.covers("logging.otel.failure.exports_metric", exercised_on=["chat_completions"])
@ -704,9 +670,7 @@ class TestOtelTraceCompleteness:
The test uses a deployment with an invalid upstream API key. This
allows the request to pass LiteLLM’s proxy authentication and fail at
the provider, which is necessary to generate a model-call error span.
There should be no cost-write span because failed requests are not
billed."""
the provider, which is necessary to generate a model-call error span."""
route = "/chat/completions"
_assert_otel_destination_configured(client)
@ -736,10 +700,8 @@ class TestOtelTraceCompleteness:
assert outcome.call_id is not None, "failed responses must still carry x-litellm-call-id"
genai_span = f"chat {model_name}"
traces = _poll(
otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False
)
_assert_complete_trace(traces, route=route, genai_span=genai_span, require_cost_span=False)
traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span)
_assert_complete_trace(traces, route=route, genai_span=genai_span)
root = next(span for span in traces.hits[0].spans if not span.references)
assert str(_tag(root, "http.status_code")) == "401", (
@ -760,8 +722,7 @@ class TestOtelTraceCompleteness:
litellm.provider.error.llm_provider attribute.
Same setup as the chat sibling: a deployment with an invalid upstream
API key passes proxy auth and fails at the provider with a real 401,
and failed requests are not billed, so no cost-write span."""
API key passes proxy auth and fails at the provider with a real 401."""
route = "/v1/messages"
_assert_otel_destination_configured(client)
@ -792,10 +753,8 @@ class TestOtelTraceCompleteness:
assert outcome.call_id is not None, "failed responses must still carry x-litellm-call-id"
genai_span = f"chat {model_name}"
traces = _poll(
otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False
)
_assert_complete_trace(traces, route=route, genai_span=genai_span, require_cost_span=False)
traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span)
_assert_complete_trace(traces, route=route, genai_span=genai_span)
root = next(span for span in traces.hits[0].spans if not span.references)
assert str(_tag(root, "http.status_code")) == "401", (

View file

@ -1,11 +1,10 @@
import asyncio
import json
import unittest.mock as mock
import pytest
from fastapi import HTTPException
from fastapi.testclient import TestClient
from litellm_enterprise.enterprise_callbacks.send_emails.endpoints import (
_get_email_settings,
_save_email_settings,
@ -21,6 +20,9 @@ from litellm_enterprise.types.enterprise_callbacks.send_emails import (
EmailEventSettingsUpdateRequest,
)
from litellm._service_logger import ServiceTypes
from tests.unit.proxy.db.fake_prisma_engine import engine_call
# Mock user_api_key_auth dependency
@pytest.fixture
@ -347,3 +349,21 @@ async def test_reset_event_settings_surfaces_the_config_owned_refusal(mock_user_
assert refused.value.status_code == 400
assert refused.value.detail["keys"] == ["email_settings"]
assert upserts == []
@pytest.mark.asyncio
async def test_save_email_settings_emits_a_postgres_upsert_event_for_litellm_config(mock_prisma_client):
mock_prisma_client.db.litellm_config.upsert = engine_call()
success = mock.AsyncMock()
service_logging = mock.MagicMock(async_service_success_hook=success, async_service_failure_hook=mock.AsyncMock())
with mock.patch("litellm.proxy.proxy_server.proxy_logging_obj", mock.MagicMock(service_logging_obj=service_logging)):
await _save_email_settings(mock_prisma_client, {"send_key_created_email": True})
await asyncio.sleep(0)
event = success.await_args.kwargs
assert (event["service"], event["call_type"], event["event_metadata"]) == (
ServiceTypes.DB,
"save_email_settings",
{"table_name": "LiteLLM_Config"},
)

View file

@ -243,6 +243,27 @@ def test_batch_write_service_is_also_attributed_to_postgres():
assert attrs["server.address"] == "litellm-prod.abc123.us-east-1.rds.amazonaws.com"
def test_resolved_prisma_operation_puts_verb_table_and_summary_on_the_db_keys():
from litellm.integrations.otel.model.spans import PostgresOperation
with patch.dict(os.environ, {"DATABASE_URL": LOCAL_DSN}, clear=False):
os.environ.pop("DATABASE_URL_READ_REPLICA", None)
attrs = dict(db_span_attributes("postgres", "update_data", PostgresOperation("update", "LiteLLM_TeamTable")))
bare = dict(db_span_attributes("postgres", "update_data", PostgresOperation("update", None)))
assert attrs == {
"db.system.name": "postgresql",
"db.system": "postgresql",
"db.operation.name": "update",
"db.collection.name": "LiteLLM_TeamTable",
"db.query.summary": "UPDATE LiteLLM_TeamTable",
"server.address": "localhost",
"server.port": 5432,
"db.namespace": "litellm",
}
assert bare["db.operation.name"] == "update"
assert {"db.collection.name", "db.query.summary"}.isdisjoint(bare)
def test_redis_service_never_borrows_the_postgres_endpoint():
assert _resolve("redis", "set", database_url=REMOTE_DSN) == {
"db.system.name": "redis",

View file

@ -6,10 +6,15 @@ hooks, proxy SERVER span lifecycle (start + setters), parent-context resolution
(ambient context), and Baggage promotion onto child spans.
"""
import ast
import asyncio
import contextlib
import os
import re
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Final
from unittest.mock import patch
import pytest
@ -1745,15 +1750,58 @@ def test_service_span_verb_follows_the_cache_method_behind_the_call(call_type, t
assert service_span_name(ServiceSpanData(service_name="redis", call_type=call_type)) == untargeted
def test_postgres_service_span_keeps_its_function_name_inside_a_targeted_phase():
"""A DB helper that runs inside ``service_target("auth_objects")`` (the whole auth phase
does) is still ``postgres get_data``: the verb scheme is for cache methods, the Postgres
rename to ``db.select {table}`` is a separate change."""
@pytest.mark.parametrize(
("call_type", "event_metadata", "expected"),
[
("get_data", {"table_name": "combined_view"}, "postgres.select LiteLLM_VerificationToken"),
("get_data", {"table_name": "team"}, "postgres.select LiteLLM_TeamTable"),
("get_data", {}, "postgres get_data"),
("get_generic_data", {"table_name": "users"}, "postgres.select LiteLLM_UserTable"),
("insert_data", {"table_name": "key"}, "postgres.insert LiteLLM_VerificationToken"),
("update_data", {"table_name": "team"}, "postgres.update LiteLLM_TeamTable"),
("delete_data", {"table_name": "user"}, "postgres.delete LiteLLM_UserTable"),
("get_user_object", {}, "postgres.select LiteLLM_UserTable"),
("get_key_object", {}, "postgres.select LiteLLM_VerificationToken"),
("_get_team_db_check", {}, "postgres.select LiteLLM_TeamTable"),
("get_org_object", {}, "postgres.select LiteLLM_OrganizationTable"),
("get_end_user_object", {}, "postgres.select LiteLLM_EndUserTable"),
("get_object_permission", {}, "postgres.select LiteLLM_ObjectPermissionTable"),
("commit_spend_updates", {"table_name": "LiteLLM_UserTable"}, "postgres.update LiteLLM_UserTable"),
("upsert_daily_spend", {"table_name": "LiteLLM_DailyTeamSpend"}, "postgres.upsert LiteLLM_DailyTeamSpend"),
("insert_spend_logs", {"table_name": "LiteLLM_SpendLogs"}, "postgres.insert LiteLLM_SpendLogs"),
("update_end_user_spend", {"table_name": "LiteLLM_EndUserTable"}, "postgres.upsert LiteLLM_EndUserTable"),
("migrate_config_credentials", {"table_name": "LiteLLM_Config"}, "postgres.update LiteLLM_Config"),
("migrate_sso_credentials", {"table_name": "LiteLLM_SSOConfig"}, "postgres.update LiteLLM_SSOConfig"),
(
"backfill_mcp_oauth_issuer",
{"table_name": "LiteLLM_MCPServerTable"},
"postgres.update LiteLLM_MCPServerTable",
),
("auto_register_jwt_mapping", {"table_name": "LiteLLM_JWTKeyMapping"}, "postgres.insert LiteLLM_JWTKeyMapping"),
(
"delete_orphaned_jwt_key",
{"table_name": "LiteLLM_VerificationToken"},
"postgres.delete LiteLLM_VerificationToken",
),
("save_email_settings", {"table_name": "LiteLLM_Config"}, "postgres.upsert LiteLLM_Config"),
("get_data", {"table_name": "DROP TABLE x"}, "postgres get_data"),
("get_data", {"table_name": "LiteLLM_NotInSchema"}, "postgres get_data"),
("some_new_helper", {"table_name": "key"}, "postgres some_new_helper"),
],
)
def test_postgres_service_span_is_named_by_sql_verb_and_prisma_table(call_type, event_metadata, expected):
"""A Postgres helper renders as ``postgres.{verb} {table}``: the verb comes from the helper,
the table from the helper when it only ever touches one model and from the event's
``table_name`` metadata otherwise (only the bounded ``PrismaClient`` literals and
``LiteLLM_*`` model names resolve, so a stray string cannot become a span name). The
ambient ``service_target`` names a cache key family and never leaks into the name."""
from litellm.integrations.otel.model.payloads import ServiceSpanData
from litellm.integrations.otel.model.spans import service_span_name
data = ServiceSpanData(service_name="postgres", call_type="get_data", target="auth_objects")
assert service_span_name(data) == "postgres get_data"
data = ServiceSpanData(
service_name="postgres", call_type=call_type, target="auth_objects", event_metadata=event_metadata
)
assert service_span_name(data) == expected
def test_service_target_declared_by_the_producer_rides_the_service_logger_payload():
@ -1876,7 +1924,9 @@ def test_async_service_success_hook_emits_service_span():
def test_postgres_db_span_names_the_database_server_not_the_prisma_engine():
"""Prisma reaches Postgres over loopback, so without server.address the
backend attributes the wait to localhost."""
backend attributes the wait to localhost. The span is named by SQL verb and
table, with the raw helper name kept on ``litellm.service.call_type`` for the
metric labels and the verb, table and summary on the ``db.*`` semconv keys."""
dsn = "postgresql://llmproxy:dbpassword9090@litellm-prod.abc123.us-east-1.rds.amazonaws.com:6432/litellm?schema=reporting"
logger, exporter = _logger()
parent = _service_parent(logger)
@ -1887,14 +1937,19 @@ def test_postgres_db_span_names_the_database_server_not_the_prisma_engine():
logger.async_service_success_hook(
payload=_ServicePayload("postgres", "get_data"),
parent_otel_span=parent,
event_metadata={"table_name": "combined_view"},
)
)
finally:
parent.end()
span = {s.name: s for s in exporter.get_finished_spans()}["postgres get_data"]
span = {s.name: s for s in exporter.get_finished_spans()}["postgres.select LiteLLM_VerificationToken"]
assert span.kind is SpanKind.CLIENT
assert span.parent.span_id == parent.get_span_context().span_id
assert span.attributes[LiteLLM.SERVICE_CALL_TYPE] == "get_data"
assert span.attributes["db.system.name"] == "postgresql"
assert span.attributes["db.operation.name"] == "get_data"
assert span.attributes["db.operation.name"] == "select"
assert span.attributes["db.collection.name"] == "LiteLLM_VerificationToken"
assert span.attributes["db.query.summary"] == "SELECT LiteLLM_VerificationToken"
assert span.attributes["server.address"] == "litellm-prod.abc123.us-east-1.rds.amazonaws.com"
assert span.attributes["server.port"] == 6432
assert span.attributes["db.namespace"] == "litellm|reporting"
@ -1904,6 +1959,44 @@ def test_postgres_db_span_names_the_database_server_not_the_prisma_engine():
assert "llmproxy" not in exported
def test_postgres_helper_without_a_known_table_keeps_the_legacy_name_and_the_verb_attribute():
"""A ``get_data`` event with no resolvable ``table_name`` must not ship as a half-named
``postgres.select``: it keeps the legacy ``postgres get_data`` name so the gap is visible,
while ``db.operation.name`` still says SELECT and no ``db.collection.name`` is made up."""
logger, exporter = _logger()
parent = _service_parent(logger)
try:
asyncio.run(
logger.async_service_success_hook(payload=_ServicePayload("postgres", "get_data"), parent_otel_span=parent)
)
finally:
parent.end()
span = {s.name: s for s in exporter.get_finished_spans()}["postgres get_data"]
assert span.attributes["db.operation.name"] == "select"
assert "db.collection.name" not in span.attributes
assert "db.query.summary" not in span.attributes
def test_redis_service_span_attributes_keep_the_raw_method_on_db_operation_name():
"""The Postgres verb table must not reach Redis: a Redis call keeps its raw method on
``db.operation.name`` and never grows a ``db.collection.name``."""
logger, exporter = _logger()
parent = _service_parent(logger)
try:
asyncio.run(
logger.async_service_success_hook(
payload=_ServicePayload("redis", "async_get_cache", target="llm_response"),
parent_otel_span=parent,
event_metadata={"table_name": "key"},
)
)
finally:
parent.end()
span = {s.name: s for s in exporter.get_finished_spans()}["redis.get llm_response"]
assert span.attributes["db.operation.name"] == "async_get_cache"
assert "db.collection.name" not in span.attributes
def test_async_service_failure_hook_marks_error_status():
logger, exporter = _logger()
parent = _service_parent(logger)
@ -2065,7 +2158,7 @@ def test_service_call_that_outlives_the_request_roots_its_own_trace_linked_to_th
logger, exporter = _logger()
server = _ended_request_span(logger)
hook = logger.async_service_success_hook(
payload=_ServicePayload("batch_write_to_db", "_PROXY_track_cost_callback"),
payload=_ServicePayload("postgres", "get_key_object"),
parent_otel_span=server if parent_source == "threaded" else None,
start_time=_REQUEST_END + 0.1,
end_time=_REQUEST_END + 0.5,
@ -2076,7 +2169,7 @@ def test_service_call_that_outlives_the_request_roots_its_own_trace_linked_to_th
else:
asyncio.run(hook)
by_name = {s.name: s for s in exporter.get_finished_spans()}
span = by_name["batch_write_to_db _PROXY_track_cost_callback"]
span = by_name["postgres.select LiteLLM_VerificationToken"]
request_ctx = server.get_span_context()
assert span.parent is None
assert span.context.trace_id != request_ctx.trace_id
@ -2094,13 +2187,13 @@ def test_service_call_that_finished_before_the_response_stays_in_the_request_tra
server = _ended_request_span(logger)
asyncio.run(
logger.async_service_success_hook(
payload=_ServicePayload("postgres", "get_data"),
payload=_ServicePayload("postgres", "get_user_object"),
parent_otel_span=server,
start_time=_REQUEST_END - 0.5,
end_time=_REQUEST_END - 0.1,
)
)
span = {s.name: s for s in exporter.get_finished_spans()}["postgres get_data"]
span = {s.name: s for s in exporter.get_finished_spans()}["postgres.select LiteLLM_UserTable"]
assert span.parent.span_id == server.get_span_context().span_id
assert span.context.trace_id == server.get_span_context().trace_id
assert list(span.links) == []
@ -3322,3 +3415,89 @@ def test_a_classifier_closed_without_a_carrier_still_nests_under_the_route_phase
assert by_name["chat gpt-4o-mini"].parent.span_id == route.get_span_context().span_id
assert by_name["chat gpt-4o-mini"].attributes[LiteLLM.REQUEST_PURPOSE] == AUTOROUTER_CLASSIFIER_CALL_ORIGIN
assert by_name["chat gpt-4o"].parent.span_id == root.get_span_context().span_id
@dataclass(frozen=True, slots=True)
class _CrudCallSite:
location: str
call_type: str
table_name: str | None
_GENERIC_CRUD_HELPERS: Final = frozenset({"get_data", "get_generic_data", "insert_data", "update_data", "delete_data"})
_PRISMA_RECEIVERS: Final = frozenset({"self", "db"})
_SOURCE_ROOTS: Final = ("litellm", "enterprise", "litellm-proxy-extras")
def _declared_table(call: ast.Call, call_type: str) -> str | None:
"""The table the call names: a literal ``table_name``, else the lookup key the
``PrismaClient`` CRUD helpers and ``@log_db_metrics`` both infer it from."""
from litellm.proxy.db.log_db_metrics import _DEFAULT_TABLE_BY_KWARG, _PRISMA_CLIENT_CRUD
by_arg: Final = {keyword.arg: keyword.value for keyword in call.keywords}
literal: Final = by_arg.get("table_name")
if isinstance(literal, ast.Constant) and isinstance(literal.value, str):
return literal.value
if call_type not in _PRISMA_CLIENT_CRUD:
return None
return next((table for key, table in _DEFAULT_TABLE_BY_KWARG.items() if key in by_arg), None)
def _is_prisma_crud_call(node: ast.AST) -> bool:
if not isinstance(node, ast.Call) or not isinstance(node.func, ast.Attribute):
return False
if node.func.attr not in _GENERIC_CRUD_HELPERS:
return False
receiver: Final = ast.unparse(node.func.value)
return receiver in _PRISMA_RECEIVERS or receiver.endswith("prisma_client")
def _crud_call_sites_in(path: Path, repo: Path) -> tuple[_CrudCallSite, ...]:
tree: Final = ast.parse(path.read_text(encoding="utf-8"))
return tuple(
_CrudCallSite(
location=f"{path.relative_to(repo)}:{node.lineno}",
call_type=node.func.attr,
table_name=_declared_table(node, node.func.attr),
)
for node in ast.walk(tree)
if _is_prisma_crud_call(node) and isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute)
)
def _source_files(repo: Path) -> tuple[Path, ...]:
roots: Final = (repo / root for root in _SOURCE_ROOTS)
files: Final = (path for root in roots for path in root.rglob("*.py")) # comprehension-ok: flatten
return tuple(path for path in files if "tests" not in path.parts and "node_modules" not in path.parts)
def test_every_prisma_crud_call_site_names_a_known_table_so_no_half_named_postgres_span_ships():
"""``PrismaClient.get_data`` and friends dispatch on ``table_name``, and so does the span
name. A call site that leaves it out would render the legacy ``postgres get_data`` with no
table, so every direct call in the proxy sources must name one, as a literal the schema
knows or through a lookup key the helper infers it from, and render ``postgres.{verb} LiteLLM_*``."""
from litellm.integrations.otel.model.payloads import ServiceSpanData
from litellm.integrations.otel.model.spans import _PRISMA_MODEL_BY_TABLE_NAME, service_span_name
repo = Path(__file__).resolve().parents[4]
per_file = (_crud_call_sites_in(path, repo) for path in _source_files(repo))
sites = tuple(site for sites_in_file in per_file for site in sites_in_file) # comprehension-ok: flatten
assert len(sites) >= 40, f"the scan lost the PrismaClient call sites: {sites}"
unresolved = [site for site in sites if site.table_name not in _PRISMA_MODEL_BY_TABLE_NAME]
assert unresolved == [], f"PrismaClient CRUD calls whose table the schema cannot resolve: {unresolved}"
rendered = {
site.location: service_span_name(
ServiceSpanData(
service_name="postgres", call_type=site.call_type, event_metadata={"table_name": site.table_name}
)
)
for site in sites
}
half_named = {
location: name
for location, name in rendered.items()
if re.fullmatch(r"postgres\.(select|insert|update|delete) LiteLLM_\w+", name) is None
}
assert half_named == {}, half_named

View file

@ -1,10 +1,15 @@
"""Tests for the one-time heal of issuer values a released version's discovery write-back stamped."""
import asyncio
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm._service_logger import ServiceTypes
from litellm.proxy import proxy_server
from tests.unit.proxy.db.fake_prisma_engine import engine_call
from litellm.proxy._experimental.mcp_server.oauth_issuer_stamp_backfill import (
backfill_discovery_stamped_issuers,
)
@ -127,3 +132,23 @@ async def test_a_failed_row_does_not_abort_the_rest():
assert await backfill_discovery_stamped_issuers(prisma_client) == 1
assert prisma_client.db.litellm_mcpservertable.update.await_count == 2
@pytest.mark.asyncio
async def test_each_healed_row_emits_a_postgres_update_event_for_the_mcp_server_table(monkeypatch):
prisma_client = _prisma([_row(server_id="a"), _row(server_id="b")])
prisma_client.db.litellm_mcpservertable.update = engine_call()
success: Final = AsyncMock()
service_logging: Final = MagicMock(async_service_success_hook=success, async_service_failure_hook=AsyncMock())
monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(service_logging_obj=service_logging))
assert await backfill_discovery_stamped_issuers(prisma_client) == 2
await asyncio.sleep(0)
assert success.await_count == 2
event: Final = success.await_args.kwargs
assert (event["service"], event["call_type"], event["event_metadata"]) == (
ServiceTypes.DB,
"backfill_mcp_oauth_issuer",
{"table_name": "LiteLLM_MCPServerTable"},
)

View file

@ -64,6 +64,7 @@ from litellm.proxy.auth.user_api_key_auth import (
user_api_key_auth_websocket_for_model,
)
from litellm.proxy.spend_tracking.carried_budget_state import carried_budget_metadata
from tests.unit.proxy.db.fake_prisma_engine import engine_call
class _RoutingRequest:
@ -10142,3 +10143,63 @@ async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(mon
)
assert admitted.authenticated_by_custom_auth is False
assert admitted.via_virtual_key is True
@pytest.mark.asyncio
async def test_auto_register_mapping_insert_emits_a_postgres_insert_event_for_the_jwt_key_mapping_table():
from litellm._service_logger import ServiceTypes
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.resolvers.models import CredentialRef
from litellm.proxy.auth.resolvers.store import IdentityStore
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy.proxy_server import hash_token
plaintext = "sk-auto-registered-span"
token_hash = hash_token(plaintext)
principal = IdentityStore._principal_from_key(
UserAPIKeyAuth(token=token_hash, user_id="validated-user", team_id="validated-team"),
auth_method=AuthMethod.API_KEY,
credential_ref=CredentialRef(token_id=token_hash),
)
prisma_client = MagicMock()
prisma_client.db.litellm_jwtkeymapping.create = engine_call()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(virtual_key_mapping_cache_ttl=300)
success = AsyncMock()
service_logging = MagicMock(async_service_success_hook=success, async_service_failure_hook=AsyncMock())
with (
patch( # test-quality-ok: key creation is an inline import inside the helper; no dependency injection seam exists
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new_callable=AsyncMock,
return_value={"token": plaintext},
),
patch( # test-quality-ok: the helper constructs IdentityStore itself; no dependency injection seam exists
"litellm.proxy.auth.resolvers.store.IdentityStore.resolve",
new_callable=AsyncMock,
return_value=principal,
),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=service_logging)),
):
await _auto_register_jwt_mapping(
virtual_key_claim_field="sub",
claim_value="user1",
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=MagicMock(),
cache_key="jwt_key_mapping:sub:user1",
team_id="validated-team",
user_id="validated-user",
)
await asyncio.sleep(0)
event = success.await_args.kwargs
assert (event["service"], event["call_type"], event["event_metadata"]) == (
ServiceTypes.DB,
"auto_register_jwt_mapping",
{"table_name": "LiteLLM_JWTKeyMapping"},
)

View file

@ -2,6 +2,7 @@ import asyncio
import json
import sys
import types
from collections.abc import Awaitable, Callable
from datetime import datetime, timedelta, timezone
from datetime import time as dt_time
from typing import Any, Dict, Final, List, Optional
@ -20,8 +21,15 @@ from litellm.constants import (
RESET_BUDGET_JOB_LOCK_TTL_SECONDS,
RESET_BUDGET_JOB_NAME,
)
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob, _RowReset
from litellm.proxy.common_utils.reset_budget_job import (
ResetBudgetJob,
_RowReset,
_write_key_windows,
_write_team_windows,
)
from litellm.proxy.common_utils.timezone_utils import BudgetResetSettings
from litellm.proxy.utils import PrismaClient
from tests.unit.proxy.db.fake_prisma_engine import engine_call
# Mock classes for testing
@ -3578,3 +3586,27 @@ def test_reset_deletes_spend_counter_instead_of_seeding(reset_budget_job, mock_p
counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:user:carol")
counter_cache.in_memory_cache.set_cache.assert_not_called()
counter_cache.redis_cache.async_set_cache.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
("write_windows", "prisma_table", "span_name"),
[
(_write_key_windows, "litellm_verificationtoken", "postgres.update LiteLLM_VerificationToken"),
(_write_team_windows, "litellm_teamtable", "postgres.update LiteLLM_TeamTable"),
],
)
async def test_a_budget_window_write_renders_a_postgres_update_span_for_its_table(
postgres_span_names: Callable[[], Awaitable[tuple[str, ...]]],
write_windows: Callable[[PrismaClient, str, str], Awaitable[None]],
prisma_table: str,
span_name: str,
) -> None:
prisma = MagicMock()
update = engine_call()
setattr(prisma.db, prisma_table, MagicMock(update=update))
await write_windows(prisma, "row-1", "{}")
assert update.await_count == 1
assert await postgres_span_names() == (span_name,)

View file

@ -6,17 +6,20 @@ import inspect
import os
import tempfile
import warnings
from collections.abc import Iterator
from typing import Dict, Optional
from collections.abc import Awaitable, Callable, Iterator
from typing import Dict, Final, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import yaml
from fastapi.testclient import TestClient
from prisma.errors import ClientNotConnectedError
import litellm
import litellm.proxy.proxy_server
from litellm._service_logger import ServiceTypes
from litellm.integrations.otel.model.payloads import ServiceSpanData
from litellm.integrations.otel.model.spans import service_span_name
from tests.unit.litellm_core_utils.fake_secret_vault import FakeSecretVault
@ -416,3 +419,28 @@ def fresh_agent_read_through(monkeypatch):
)
monkeypatch.setattr(registry_read_through, "agent_registry_read_through", read_through)
return read_through
@pytest.fixture
def postgres_span_names() -> Iterator[Callable[[], Awaitable[tuple[str, ...]]]]:
"""The ``postgres.{verb} {table}`` names OTel would render for every DB service event
the code under test emits, in emission order, once the hook tasks have run."""
success: Final = AsyncMock()
service_logging: Final = MagicMock(async_service_success_hook=success, async_service_failure_hook=AsyncMock())
async def rendered() -> tuple[str, ...]:
await asyncio.sleep(0)
return tuple(
service_span_name(
ServiceSpanData(
service_name="postgres",
call_type=call.kwargs["call_type"],
event_metadata=call.kwargs["event_metadata"] or {},
)
)
for call in success.await_args_list
if call.kwargs["service"] == ServiceTypes.DB
)
with patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=service_logging)):
yield rendered

View file

@ -4,6 +4,7 @@ selection, the non-partitioned no-op safety path, and the drop/ensure SQL flow.
"""
from contextlib import asynccontextmanager
from collections.abc import Awaitable, Callable
from datetime import date, datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock
@ -18,6 +19,7 @@ from litellm.proxy.db.db_transaction_queue.spend_logs_partition_manager import (
select_partitions_to_drop,
upcoming_partitions,
)
from tests.unit.proxy.db.fake_prisma_engine import engine_call
DDL_TIMEOUT_MS = 30000
@ -411,3 +413,15 @@ async def test_drop_partitions_continues_when_one_drop_fails():
# both were eligible; the first drop failed so only the second is reported
assert dropped == ["LiteLLM_SpendLogs_p20260602"]
@pytest.mark.asyncio
async def test_the_partitioning_probe_renders_a_postgres_select_span(
postgres_span_names: Callable[[], Awaitable[tuple[str, ...]]],
) -> None:
client = MagicMock()
client.db.query_raw = engine_call([{"partitioned": True}])
_wire_tx(client.db)
assert await SpendLogsPartitionManager().is_partitioned(client, _budget()) is True
assert await postgres_span_names() == ("postgres.select LiteLLM_SpendLogs",)

View file

@ -0,0 +1,18 @@
"""An ``AsyncMock`` standing in for a ``prisma_client.db`` method that reached the engine,
marking the DB I/O witness the way ``_TrackedPrismaEngine`` does, so the producer under test
emits its service event."""
from typing import TypeVar
from unittest.mock import AsyncMock
from litellm.proxy.db.log_db_metrics import record_db_io
_T = TypeVar("_T")
def engine_call(return_value: _T | None = None) -> AsyncMock:
async def run(*args: object, **kwargs: object) -> _T | None:
record_db_io()
return return_value
return AsyncMock(side_effect=run)

View file

@ -17,10 +17,13 @@ import pytest
from litellm.proxy.db.autorouter_session_rollup import (
UPSERT_AUTOROUTER_SESSION_SQL,
UPSERT_AUTOROUTER_USER_SESSION_SQL,
AutoRouterTurnTransaction,
build_autorouter_turn_transaction,
flush_autorouter_turn_transactions,
write_autorouter_turn,
)
from tests.unit.proxy.db.fake_prisma_engine import engine_call
ROUTING_DECISION = {"router_model_name": "live-auto", "router_type": "complexity", "routed_model": "haiku"}
@ -485,3 +488,21 @@ def test_internal_call_origin_never_reaches_the_rollup():
gate alone would count it; the internal_call_origin stamp must exclude it."""
assert _build(metadata=_metadata(internal_call_origin="shadow_eval_router")) is None
assert _build() is not None
@pytest.mark.asyncio
@pytest.mark.parametrize(
("statement", "span_name"),
(
(UPSERT_AUTOROUTER_SESSION_SQL, "postgres.upsert LiteLLM_AutoRouterSession"),
(UPSERT_AUTOROUTER_USER_SESSION_SQL, "postgres.upsert LiteLLM_AutoRouterUserSession"),
),
)
async def test_the_turn_upsert_span_names_the_session_table_its_statement_writes(
statement: str, span_name: str, postgres_span_names
) -> None:
db: Final = SimpleNamespace(execute_raw=engine_call())
await write_autorouter_turn(db, _transaction(user_id="u1"), statement)
assert await postgres_span_names() == (span_name,)

View file

@ -1,5 +1,6 @@
import math
from contextlib import asynccontextmanager
from collections.abc import Awaitable, Callable
from datetime import datetime, timedelta, timezone
from typing import Any
@ -14,6 +15,7 @@ from litellm.proxy.db.budget_window_spend_writer import (
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
build_window_spend_transaction,
)
from litellm.proxy.db.log_db_metrics import record_db_io
WINDOW_A = datetime(2026, 8, 1, tzinfo=timezone.utc)
WINDOW_B = datetime(2026, 8, 31, tzinfo=timezone.utc)
@ -42,10 +44,12 @@ class _FakeDB:
self.committed = False
async def query_raw(self, query: str, *args: Any) -> list[dict[str, str]]:
record_db_io()
self.query_raw_calls.append((query, args))
return self.existing_rows
async def execute_raw(self, query: str, *args: Any) -> int:
record_db_io()
self.execute_raw_calls.append((query, args))
return 1
@ -59,6 +63,7 @@ class _FakeDB:
@asynccontextmanager
async def _batch(self):
yield self.batcher
record_db_io()
self.committed = True
def batch_(self):
@ -592,3 +597,17 @@ async def test_seed_aggregate_treats_an_entity_with_no_rows_as_zero():
)
assert totals == WindowSeedTotals(total=0.0, before_batch=0.0)
@pytest.mark.asyncio
async def test_rolling_a_window_row_renders_a_postgres_update_span(
postgres_span_names: Callable[[], Awaitable[tuple[str, ...]]],
) -> None:
await roll_window_spend_row(
prisma_client=_FakePrismaClient(_FakeDB()),
entity_type="team",
entity_id="t1",
window_duration="30d",
new_window_start=WINDOW_B,
)
assert await postgres_span_names() == ("postgres.update LiteLLM_BudgetWindowSpend",)

View file

@ -0,0 +1,124 @@
import asyncio
from collections.abc import Iterator
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from prisma.errors import PrismaError
from litellm._service_logger import ServiceTypes
from litellm.proxy.db.db_span import db_span
from litellm.proxy.db.log_db_metrics import record_db_io
@pytest.fixture
def service_hooks() -> Iterator[tuple[AsyncMock, AsyncMock]]:
success: Final = AsyncMock()
failure: Final = AsyncMock()
service_logging: Final = MagicMock(async_service_success_hook=success, async_service_failure_hook=failure)
with patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=service_logging)):
yield success, failure
@pytest.mark.asyncio
async def test_a_completed_write_emits_one_db_event_named_for_the_call_and_table(
service_hooks: tuple[AsyncMock, AsyncMock],
) -> None:
success, failure = service_hooks
async with db_span("commit_spend_updates", "LiteLLM_UserTable"):
record_db_io()
await asyncio.sleep(0)
event: Final = success.await_args.kwargs
assert (event["service"], event["call_type"], event["event_metadata"]) == (
ServiceTypes.DB,
"commit_spend_updates",
{"table_name": "LiteLLM_UserTable"},
)
assert event["duration"] == pytest.approx((event["end_time"] - event["start_time"]).total_seconds())
assert failure.await_count == 0
@pytest.mark.asyncio
async def test_a_prisma_error_inside_the_write_emits_a_db_failure_event_and_propagates(
service_hooks: tuple[AsyncMock, AsyncMock],
) -> None:
success, failure = service_hooks
with pytest.raises(PrismaError):
async with db_span("insert_spend_logs", "LiteLLM_SpendLogs"):
raise PrismaError("connection reset")
await asyncio.sleep(0)
event: Final = failure.await_args.kwargs
assert (event["service"], event["call_type"], event["event_metadata"], str(event["error"])) == (
ServiceTypes.DB,
"insert_spend_logs",
{"table_name": "LiteLLM_SpendLogs"},
"connection reset",
)
assert success.await_count == 0
@pytest.mark.asyncio
async def test_a_dropped_query_engine_connection_emits_a_db_failure_event(
service_hooks: tuple[AsyncMock, AsyncMock],
) -> None:
success, failure = service_hooks
with pytest.raises(httpx.ReadError):
async with db_span("write_tool_spend", "LiteLLM_DailyToolSpend"):
raise httpx.ReadError("peer closed connection")
await asyncio.sleep(0)
event: Final = failure.await_args.kwargs
assert (event["call_type"], event["event_metadata"], str(event["error"])) == (
"write_tool_spend",
{"table_name": "LiteLLM_DailyToolSpend"},
"peer closed connection",
)
assert success.await_count == 0
@pytest.mark.asyncio
async def test_a_non_database_error_inside_the_write_emits_no_db_event(
service_hooks: tuple[AsyncMock, AsyncMock],
) -> None:
success, failure = service_hooks
with pytest.raises(ValueError, match="bad row"):
async with db_span("insert_spend_logs", "LiteLLM_SpendLogs"):
raise ValueError("bad row")
await asyncio.sleep(0)
assert (success.await_count, failure.await_count) == (0, 0)
@pytest.mark.asyncio
async def test_a_raising_failure_hook_never_replaces_the_prisma_error(
service_hooks: tuple[AsyncMock, AsyncMock],
) -> None:
success, failure = service_hooks
failure.side_effect = RuntimeError("exporter down")
with pytest.raises(PrismaError):
async with db_span("commit_spend_updates", "LiteLLM_UserTable"):
raise PrismaError("connection reset")
assert failure.await_count == 1
assert success.await_count == 0
@pytest.mark.asyncio
async def test_a_block_whose_prisma_client_never_reached_the_engine_emits_no_db_event(
service_hooks: tuple[AsyncMock, AsyncMock],
) -> None:
success, failure = service_hooks
async with db_span("team_user_spend", "LiteLLM_SpendLogs"):
await asyncio.sleep(0)
await asyncio.sleep(0)
assert (success.await_count, failure.await_count) == (0, 0)

View file

@ -3,8 +3,6 @@ import copy
import json
import logging
import re
from collections.abc import AsyncIterator, Callable
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from datetime import datetime, timedelta, timezone
@ -20,13 +18,14 @@ from redis.exceptions import DataError
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._service_logger import ServiceTypes
from litellm.proxy._types import DailyTagSpendTransaction, Litellm_EntityType, SpendUpdateQueueItem
from litellm.proxy.db.db_spend_update_writer import (
_TEAM_ADVISORY_LOCK_SQL,
_TEAM_MEMBER_SPEND_SQL,
DBSpendUpdateWriter,
_SpendTableName,
_spend_tables_left_to_send,
_SpendTableName,
)
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import DailySpendUpdateQueue
from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer
@ -34,6 +33,7 @@ from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdate
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
build_window_spend_transaction,
)
from tests.unit.proxy.db.fake_prisma_engine import engine_call
@pytest.mark.asyncio
@ -1144,6 +1144,50 @@ async def test_org_spend_increments_organization_membership_row_for_the_calling_
)
@pytest.mark.asyncio
async def test_commit_spend_updates_reports_one_db_event_per_table_it_wrote():
"""The spend flush is the proxy's main Postgres write path. Each per-table
transaction must surface as a ``ServiceTypes.DB`` event naming the table,
so the trace shows ``postgres.update LiteLLM_UserTable`` and friends instead
of nothing at all."""
db_writer: Final = DBSpendUpdateWriter()
await db_writer._update_org_db(
response_cost=0.75,
org_id="org-abc",
user_id="user-xyz",
prisma_client=MagicMock(),
)
transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions()
transactions["user_list_transactions"] = {"user-xyz": 0.75}
transactions["key_list_transactions"] = {"hash": 0.75}
mock_prisma_client: Final = MagicMock()
mock_prisma_client.db.tx = MagicMock(return_value=_good_tx(MagicMock()))
proxy_logging: Final = MagicMock()
proxy_logging.call_details = {}
success_hook: Final = AsyncMock()
with patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
MagicMock(service_logging_obj=MagicMock(async_service_success_hook=success_hook)),
):
await db_writer._commit_spend_updates_to_db(
prisma_client=mock_prisma_client,
n_retry_times=0,
proxy_logging_obj=proxy_logging,
db_spend_update_transactions=transactions,
)
await asyncio.sleep(0)
events: Final = [c.kwargs for c in success_hook.await_args_list if c.kwargs["service"] == ServiceTypes.DB]
assert sorted((e["call_type"], e["event_metadata"]["table_name"]) for e in events) == [
("commit_spend_updates", "LiteLLM_OrganizationMembership"),
("commit_spend_updates", "LiteLLM_OrganizationTable"),
("commit_spend_updates", "LiteLLM_UserTable"),
("commit_spend_updates", "LiteLLM_VerificationToken"),
]
@pytest.mark.asyncio
async def test_org_spend_without_user_id_leaves_organization_membership_untouched():
db_writer: Final = DBSpendUpdateWriter()
@ -3773,8 +3817,8 @@ def _empty_spend_transactions(**overrides):
def _good_tx(mock_batcher):
tx = AsyncMock()
tx.__aenter__ = AsyncMock(return_value=tx)
tx.__aexit__ = AsyncMock(return_value=False)
tx.query_raw = AsyncMock(return_value=[])
tx.__aexit__ = engine_call(False)
tx.query_raw = engine_call([])
tx.batch_ = MagicMock(
return_value=AsyncMock(
__aenter__=AsyncMock(return_value=mock_batcher),

View file

@ -4,6 +4,7 @@ LiteLLM_DailyGatewayRequests.
"""
import asyncio
from collections.abc import Awaitable, Callable
from datetime import datetime, timezone
import pytest
@ -18,6 +19,7 @@ from litellm.proxy.db.gateway_request_tracking import (
)
from litellm.proxy.middleware.billable_request_metrics_middleware import BillableCategory
from litellm.types.proxy.gateway_requests import GatewayRequestCounts, GatewayRequestKey
from litellm.proxy.db.log_db_metrics import record_db_io
def _today() -> str:
@ -91,6 +93,7 @@ class FakeDB:
self.statements: list[tuple[str, tuple[object, ...]]] = []
async def execute_raw(self, query: str, *args: object) -> int:
record_db_io()
self.statements.append((query, args))
return len(args) // 5
@ -512,3 +515,19 @@ def test_failed_redis_push_keeps_counts_locally_for_the_next_flush():
GatewayRequestCounts(successful_requests=1, failed_requests=1)
)
}
@pytest.mark.asyncio
async def test_a_gateway_request_flush_renders_a_postgres_upsert_span(
postgres_span_names: Callable[[], Awaitable[tuple[str, ...]]],
) -> None:
prisma = FakePrismaClient()
snapshot = {
GatewayRequestKey(date="2026-08-01", category="llm", route="/chat/completions"): (
GatewayRequestCounts(successful_requests=7, failed_requests=2)
)
}
await commit_gateway_requests_to_db(prisma_client=prisma, snapshot=snapshot)
assert await postgres_span_names() == ("postgres.upsert LiteLLM_DailyGatewayRequests",)

View file

@ -1,3 +1,4 @@
from collections.abc import Awaitable, Callable
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock
@ -10,11 +11,12 @@ from litellm.proxy.db.health_check_latest import (
fetch_latest_health_checks_for_models,
query_latest_health_checks,
)
from tests.unit.proxy.db.fake_prisma_engine import engine_call
def _prisma(rows):
prisma = MagicMock()
prisma.db.query_raw = AsyncMock(return_value=rows)
prisma.db.query_raw = engine_call(rows)
return prisma
@ -119,3 +121,11 @@ async def test_fetch_for_models_degrades_to_no_rows_when_the_query_fails():
prisma = _prisma([])
prisma.db.query_raw.side_effect = RuntimeError("db down")
assert await fetch_latest_health_checks_for_models(prisma, ("gpt-4",)) == ()
@pytest.mark.asyncio
async def test_the_latest_health_check_read_renders_a_postgres_select_span(
postgres_span_names: Callable[[], Awaitable[tuple[str, ...]]],
) -> None:
assert await fetch_latest_health_checks(_prisma([])) == ()
assert await postgres_span_names() == ("postgres.select LiteLLM_HealthCheckTable",)

View file

@ -150,3 +150,101 @@ async def test_a_cache_hit_after_a_sibling_db_read_emits_no_db_event(success_hoo
await cache_hit()
assert await _db_call_types(success_hook) == ("read_row",)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("lookup", "table_name"),
[
({"token": "sk-hashed"}, "key"),
({"tokens": ["sk-hashed"]}, "key"),
({"user_id": "u-1"}, "user"),
({"team_id": "t-1"}, "team"),
({"token": "sk-hashed", "user_id": "u-1"}, "key"),
({"table_name": "spend", "token": "sk-hashed"}, "spend"),
],
)
async def test_a_crud_method_called_without_table_name_reports_the_table_its_lookup_key_selects(
success_hook: AsyncMock, lookup: dict[str, object], table_name: str
) -> None:
engine: Final = _tracked_engine()
@log_db_metrics
async def get_data(*, table_name: str | None = None, **kwargs: object) -> object:
return await engine.query("{}", tx_id=None)
await get_data(**lookup)
await asyncio.sleep(0)
assert success_hook.await_args_list[0].kwargs["event_metadata"] == {"table_name": table_name}
@pytest.mark.asyncio
async def test_a_helper_without_a_table_name_parameter_gets_no_inferred_table(success_hook: AsyncMock) -> None:
engine: Final = _tracked_engine()
@log_db_metrics
async def get_team_member_default_budget(*, team_id: str, user_id: str) -> object:
return await engine.query("{}", tx_id=None)
await get_team_member_default_budget(team_id="t-1", user_id="u-1")
await asyncio.sleep(0)
assert success_hook.await_args_list[0].kwargs["event_metadata"] is None
_FIND_UNIQUE_KEY_PAYLOAD: Final = (
'{"query": "query { result: findUniqueLiteLLM_VerificationToken(where: {token: \\"h\\"}) { token } }"}'
)
_RAW_SELECT_PAYLOAD: Final = '{"query": "mutation { result: queryRaw(query: \\"SELECT 1 FROM \\\\\\"LiteLLM_UserTable\\\\\\"\\", parameters: \\"[]\\") }"}'
@pytest.mark.asyncio
async def test_an_undecorated_prisma_query_emits_one_db_event_named_from_the_engine_payload(
success_hook: AsyncMock,
) -> None:
engine: Final = _tracked_engine()
await engine.query(_RAW_SELECT_PAYLOAD, tx_id=None)
await engine.query(_FIND_UNIQUE_KEY_PAYLOAD, tx_id=None)
assert await _db_call_types(success_hook) == ("query_raw", "find_unique")
raw, model = (call.kwargs["event_metadata"] for call in success_hook.await_args_list)
assert raw == {"table_name": "LiteLLM_UserTable", "db_operation": "select"}
assert model == {"table_name": "LiteLLM_VerificationToken", "db_operation": "select"}
@pytest.mark.asyncio
async def test_a_decorated_call_owns_its_query_so_the_engine_fallback_stays_silent(success_hook: AsyncMock) -> None:
engine: Final = _tracked_engine()
@log_db_metrics
async def read_key_row(**kwargs: object) -> object:
return await engine.query(_FIND_UNIQUE_KEY_PAYLOAD, tx_id=None)
await read_key_row(parent_otel_span="span", token="h")
assert await _db_call_types(success_hook) == ("read_key_row",)
@pytest.mark.asyncio
async def test_a_task_spawned_by_a_decorated_call_that_queries_after_it_returned_emits_its_own_event(
success_hook: AsyncMock,
) -> None:
engine: Final = _tracked_engine()
released: Final = asyncio.Event()
async def write_after_the_caller_returned() -> object:
await released.wait()
return await engine.query(_FIND_UNIQUE_KEY_PAYLOAD, tx_id=None)
@log_db_metrics
async def read_key_row(**kwargs: object) -> asyncio.Task[object]:
await engine.query(_FIND_UNIQUE_KEY_PAYLOAD, tx_id=None)
return asyncio.create_task(write_after_the_caller_returned())
background: Final = await read_key_row(token="h")
released.set()
await background
assert await _db_call_types(success_hook) == ("read_key_row", "find_unique")

View file

@ -0,0 +1,471 @@
import ast
import re
from collections.abc import Iterator, Mapping
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import pytest
from litellm.integrations.otel.model.payloads import ServiceSpanData
from litellm.integrations.otel.model import spans as spans_mod
from litellm.integrations.otel.model.spans import (
_POSTGRES_OPERATION_BY_CALL_TYPE,
PRISMA_RELATIONS,
service_span_name,
)
from litellm.proxy.db.prisma_query_span import UNKNOWN_PRISMA_QUERY, parse_prisma_query, sql_operation
_REPO: Final = Path(__file__).resolve().parents[4]
_SOURCE_ROOTS: Final = ("litellm", "enterprise", "litellm-proxy-extras")
_RAW_METHODS: Final = frozenset({"query_first", "query_raw", "execute_raw"})
_MODEL_METHODS: Final = frozenset(
{
"find_unique",
"find_unique_or_raise",
"find_first",
"find_first_or_raise",
"find_many",
"count",
"group_by",
"create",
"create_many",
"update",
"update_many",
"delete",
"delete_many",
"upsert",
}
)
_MODEL_BY_ACCESSOR: Final[Mapping[str, str]] = {relation.lower(): relation for relation in PRISMA_RELATIONS}
_GENERIC_CRUD_HELPERS: Final = frozenset({"get_data", "get_generic_data", "insert_data", "update_data", "delete_data"})
_TRANSACTION_BODIES: Final[Mapping[str, str]] = {"litellm/proxy/db/baseline_accounting.py": "baseline_accounting"}
_RENDERED_NAME: Final = re.compile(
r"postgres\.(select|insert|update|delete|upsert|ddl|set|transaction) .+|postgres\.ping"
)
def _engine_payload(root_field: str, sql: str | None = None) -> str:
selection: Final = f'queryRaw(query: "{sql}", parameters: "[]")' if sql is not None else root_field
return (
f'{{"query": "mutation {{ result: {selection} }}"}}'
if sql is not None
else f'{{"query": "query {{ result: {root_field}(where: {{token: \\"x\\"}}) {{ token }} }}"}}'
)
@pytest.mark.parametrize(
("content", "expected"),
[
(
_engine_payload("findUniqueLiteLLM_VerificationToken"),
("find_unique", "select", "LiteLLM_VerificationToken"),
),
(_engine_payload("createOneLiteLLM_SpendLogs"), ("create", "insert", "LiteLLM_SpendLogs")),
(_engine_payload("findFirstLiteLLM_UserTableOrThrow"), ("find_first", "select", "LiteLLM_UserTable")),
(
'{"query": "mutation { result: queryRaw(query: \\"SELECT * FROM \\\\\\"LiteLLM_UserTable\\\\\\" WHERE user_id = $1\\", parameters: \\"[]\\") }"}',
("query_raw", "select", "LiteLLM_UserTable"),
),
(
'{"query": "mutation { result: executeRaw(query: \\"SET LOCAL statement_timeout = 5000\\", parameters: \\"[]\\") }"}',
("execute_raw", "set", "statement_timeout"),
),
(
'{"query": "mutation { result: queryRaw(query: \\"SELECT to_regclass($1) IS NOT NULL AS present\\", parameters: \\"[]\\") }"}',
("query_raw", "select", "pg_catalog"),
),
(
'{"query": "mutation { result: queryRaw(query: \\"SELECT 1\\", parameters: \\"[]\\") }"}',
("query_raw", "ping", None),
),
],
)
def test_the_engine_names_a_round_trip_from_its_payload_without_copying_sql_text(
content: str, expected: tuple[str, str | None, str | None]
) -> None:
query = parse_prisma_query(content)
assert (query.call_type, query.operation, query.table) == expected
assert query.table is None or " " not in query.table
def test_a_payload_the_parser_does_not_know_stays_the_legacy_function_named_span() -> None:
assert parse_prisma_query("not json at all") is UNKNOWN_PRISMA_QUERY
assert parse_prisma_query('{"query": "mutation { result: somethingNew(x: 1) }"}') is UNKNOWN_PRISMA_QUERY
rendered = service_span_name(ServiceSpanData(service_name="postgres", call_type=UNKNOWN_PRISMA_QUERY.call_type))
assert rendered == "postgres prisma_query"
@pytest.mark.parametrize(
("sql", "expected"),
[
('SELECT 1 FROM "LiteLLM_VerificationTokenView" LIMIT 1', ("select", "LiteLLM_VerificationTokenView")),
(
'\n WITH keys AS (SELECT * FROM "LiteLLM_VerificationToken") SELECT 1',
("select", "LiteLLM_VerificationToken"),
),
('INSERT INTO "LiteLLM_DailyUserSpend" (id) VALUES ($1)', ("insert", "LiteLLM_DailyUserSpend")),
("SET LOCAL lock_timeout = 1000", ("set", "lock_timeout")),
("SELECT COUNT(*) FROM pg_stat_activity", ("select", "pg_catalog")),
("SELECT 1", ("ping", None)),
("SELECT current_setting('transaction_read_only') AS transaction_read_only", ("select", "pg_catalog")),
('REFRESH MATERIALIZED VIEW "MonthlyGlobalSpend"', ("ddl", "MonthlyGlobalSpend")),
(
'WITH team_rows AS (UPDATE "LiteLLM_TeamTable" SET models = $1 RETURNING team_id) SELECT team_id FROM team_rows',
("update", "LiteLLM_TeamTable"),
),
("BEGIN", (None, None)),
],
)
def test_sql_operation_is_the_leading_verb_and_the_first_schema_relation(
sql: str, expected: tuple[str | None, str | None]
) -> None:
assert sql_operation(sql) == expected
@dataclass(frozen=True, slots=True)
class _PrismaCallSite:
location: str
method: str
owner: str
rendered: str | None
@dataclass(frozen=True, slots=True)
class _Module:
path: Path
tree: ast.Module
constants: Mapping[str, ast.expr]
def ancestors(self, node: ast.AST) -> tuple[ast.AST, ...]:
parent_of: Final = _parent_map(self.tree)
chain: Final = [node]
while (parent := parent_of.get(id(chain[-1]))) is not None:
chain.append(parent)
return tuple(chain[1:])
_PARENTS: Final[dict[int, Mapping[int, ast.AST]]] = {} # mutable-ok: per-tree parent map memo
def _parent_map(tree: ast.Module) -> Mapping[int, ast.AST]:
if id(tree) not in _PARENTS:
_PARENTS[id(tree)] = {
id(child): node for node in ast.walk(tree) for child in ast.iter_child_nodes(node)
} # comprehension-ok: parent links
return _PARENTS[id(tree)]
def _modules() -> Iterator[_Module]:
for root in _SOURCE_ROOTS:
for path in sorted((_REPO / root).rglob("*.py")):
if "tests" in path.parts or "node_modules" in path.parts:
continue
tree: Final = ast.parse(path.read_text(encoding="utf-8"))
yield _Module(path, tree, _assignments(tree.body))
def _imported_module(module: _Module, name: str) -> Path | None:
for node in module.tree.body:
if isinstance(node, ast.ImportFrom) and node.module and any(alias.name == name for alias in node.names):
return _REPO / (node.module.replace(".", "/") + ".py")
return None
def _assignments(body: list[ast.stmt]) -> Mapping[str, ast.expr]:
return {
target.id: node.value
for node in ast.walk(ast.Module(body=body, type_ignores=[]))
if isinstance(node, (ast.Assign, ast.AnnAssign)) and node.value is not None
for target in (node.targets if isinstance(node, ast.Assign) else (node.target,))
if isinstance(target, ast.Name)
} # comprehension-ok: constants by name
def _mapping_values(expr: ast.expr | None) -> ast.expr | None:
"""The dict a ``Mapping`` constant was built from, through ``MappingProxyType(...)``."""
if (
isinstance(expr, ast.Call)
and isinstance(expr.func, ast.Name)
and expr.func.id == "MappingProxyType"
and expr.args
):
return expr.args[0]
return expr if isinstance(expr, (ast.Dict, ast.DictComp)) else None
def _returned_text(function_name: str, module: _Module) -> ast.expr | None:
"""What a module-level SQL builder returns, when its body is one ``return`` of a string expression."""
for node in module.tree.body:
if isinstance(node, ast.FunctionDef) and node.name == function_name:
returns: Final = [stmt for stmt in ast.walk(node) if isinstance(stmt, ast.Return)]
return returns[0].value if len(returns) == 1 else None
return None
_DYNAMIC: Final = " ? "
def _fragment(value: ast.expr, module: _Module, depth: int) -> str:
"""One f-string piece: literal text, a module constant spliced in, or a runtime placeholder."""
spliced: Final = (
_sql_text(value.value, module, depth + 1)
if isinstance(value, ast.FormattedValue)
else _sql_text(value, module, depth)
)
return _DYNAMIC if spliced is None or _ALTERNATIVE in spliced else spliced
def _sql_text(expr: ast.expr | None, module: _Module, depth: int = 0) -> str | None:
if expr is None or depth > 3:
return None
if isinstance(expr, ast.Constant) and isinstance(expr.value, str):
return expr.value
if isinstance(expr, ast.JoinedStr):
return "".join(_fragment(value, module, depth) for value in expr.values)
if isinstance(expr, ast.BinOp) and isinstance(expr.op, ast.Add):
left: Final = _sql_text(expr.left, module, depth)
return left if left is not None else _sql_text(expr.right, module, depth)
if (
isinstance(expr, ast.Call)
and isinstance(expr.func, ast.Attribute)
and expr.func.attr in {"format", "strip", "lstrip"}
):
return _sql_text(expr.func.value, module, depth)
if isinstance(expr, ast.Call) and isinstance(expr.func, ast.Attribute) and expr.func.attr == "dedent":
return _sql_text(expr.args[0], module, depth) if expr.args else None
if isinstance(expr, ast.IfExp):
branches: Final = (_sql_text(expr.body, module, depth), _sql_text(expr.orelse, module, depth))
return branches[0] if branches[0] == branches[1] or None in branches else _multi(branches)
if isinstance(expr, ast.Subscript) and isinstance(expr.value, ast.Name):
return _sql_text(_mapping_values(module.constants.get(expr.value.id)), module, depth + 1)
if (
isinstance(expr, ast.Call)
and isinstance(expr.func, ast.Name)
and expr.func.id == "MappingProxyType"
and expr.args
):
return _sql_text(expr.args[0], module, depth)
if isinstance(expr, ast.Dict):
values: Final = tuple(_sql_text(value, module, depth) for value in expr.values)
return _multi(values) if values and None not in values else None
if isinstance(expr, ast.DictComp):
return _sql_text(expr.value, module, depth)
if isinstance(expr, ast.Call) and isinstance(expr.func, ast.Name):
returned: Final = _returned_text(expr.func.id, module)
return _sql_text(returned, module, depth + 1) if returned is not None else None
if isinstance(expr, ast.Name):
if expr.id in module.constants:
return _sql_text(module.constants[expr.id], module, depth + 1)
source: Final = _imported_module(module, expr.id)
if source is None or not source.exists():
return None
imported: Final = ast.parse(source.read_text(encoding="utf-8"))
imported_module: Final = _Module(source, imported, _assignments(imported.body))
return _sql_text(imported_module.constants.get(expr.id), imported_module, depth + 1)
return None
_ALTERNATIVE: Final = "\x1f"
def _multi(texts: tuple[str | None, ...]) -> str:
return _ALTERNATIVE.join(text for text in texts if text is not None)
def _parameters(parents: tuple[ast.AST, ...]) -> frozenset[str]:
function: Final = next((p for p in parents if isinstance(p, (ast.AsyncFunctionDef, ast.FunctionDef))), None)
if function is None:
return frozenset()
return frozenset(arg.arg for arg in (*function.args.args, *function.args.kwonlyargs))
def _argument_for(call: ast.Call, function: ast.AsyncFunctionDef | ast.FunctionDef, parameter: str) -> ast.expr | None:
positional: Final = tuple(arg.arg for arg in function.args.args)
by_keyword: Final = next((k.value for k in call.keywords if k.arg == parameter), None)
if by_keyword is not None or parameter not in positional:
return by_keyword
index: Final = positional.index(parameter)
return call.args[index] if index < len(call.args) else None
def _parameter_site(
parameter: str, module: _Module, parents: tuple[ast.AST, ...], method: str, location: str
) -> _PrismaCallSite:
"""A statement that arrives as a parameter: a ``query_raw`` forwarder adds no round trip of its
own, any other helper is named by what its callers in the module hand it."""
function: Final = next(p for p in parents if isinstance(p, (ast.AsyncFunctionDef, ast.FunctionDef)))
if function.name in _RAW_METHODS:
return _PrismaCallSite(location, method, "forwarder", f"(callers of {function.name})")
callers: Final = tuple(
node
for node in ast.walk(module.tree)
if isinstance(node, ast.Call) and ast.unparse(node.func).endswith(function.name)
)
sites: Final = tuple(_caller_site(call, function, parameter, module, method, location) for call in callers)
names: Final = tuple(site.rendered for site in sites)
owners: Final = ", ".join(sorted({site.owner for site in sites}))
rendered: Final = " | ".join(sorted(set(names))) if names and None not in names else None # pyright: ignore[reportArgumentType] # None filtered above
return _PrismaCallSite(location, method, f"{owners} via {function.name} callers", rendered)
def _caller_site(
call: ast.Call,
function: ast.AsyncFunctionDef | ast.FunctionDef,
parameter: str,
module: _Module,
method: str,
location: str,
) -> _PrismaCallSite:
parents: Final = module.ancestors(call)
wrapped: Final = _wrapper_site(call, parents, method, location, module)
if wrapped is not None:
return wrapped
text: Final = _sql_text(_argument_for(call, function, parameter), _scope(module, parents))
return _PrismaCallSite(location, method, "engine", _render_statements(method, text) if text is not None else None)
def _render_statements(method: str, text: str) -> str | None:
names: Final = tuple(_render_statement(method, alternative) for alternative in text.split(_ALTERNATIVE))
return " | ".join(sorted(set(names))) if None not in names else None # pyright: ignore[reportArgumentType] # None filtered above
def _render_statement(method: str, text: str) -> str | None:
verb, target = sql_operation(text)
if verb is not None and target is None and verb != "ping" and _DYNAMIC in text:
return f"postgres.{verb} {{relation built at runtime}}"
return _render(method, target, verb) if verb is not None else None
def _render(call_type: str, table: str | None, operation: str | None = None) -> str:
metadata: Final = {
key: value for key, value in (("table_name", table), ("db_operation", operation)) if value is not None
}
return service_span_name(ServiceSpanData(service_name="postgres", call_type=call_type, event_metadata=metadata))
def _wrapper_site(
call: ast.Call, parents: tuple[ast.AST, ...], method: str, location: str, module: _Module
) -> _PrismaCallSite | None:
for parent in parents:
items: Final = parent.items if isinstance(parent, (ast.AsyncWith, ast.With)) else ()
for item in items:
context: Final = item.context_expr
if isinstance(context, ast.Call) and isinstance(context.func, ast.Name) and context.func.id == "db_span":
return _wrapped_by(context, method, location, "db_span", _scope(module, parents))
if (
isinstance(context, ast.Call)
and isinstance(context.func, ast.Name)
and context.func.id == "_spend_update_tx"
):
call_type: Final = context.args[2] if len(context.args) > 2 else ast.Constant("commit_spend_updates")
spend_tx: Final = ast.Call(func=ast.Name("db_span"), args=[call_type, context.args[1]], keywords=[])
return _wrapped_by(spend_tx, method, location, "_spend_update_tx", _scope(module, parents))
if isinstance(parent, ast.Call) and isinstance(parent.func, ast.Name) and parent.func.id == "db_spanned":
return _wrapped_by(parent, method, location, "db_spanned", _scope(module, parents))
if isinstance(parent, (ast.AsyncFunctionDef, ast.FunctionDef)):
decorators: Final = tuple(
decorator.id for decorator in parent.decorator_list if isinstance(decorator, ast.Name)
)
if "log_db_metrics" in decorators:
return _decorated_site(parent.name, method, location)
return None
def _wrapped_by(wrapper: ast.Call, method: str, location: str, owner: str, scope: _Module) -> _PrismaCallSite:
call_type: Final = _sql_text(wrapper.args[0], scope)
table_expr: Final = wrapper.args[1] if len(wrapper.args) > 1 else None
if call_type is None:
return _PrismaCallSite(location, method, owner, None)
if isinstance(table_expr, ast.Constant) and table_expr.value is None:
return _PrismaCallSite(location, method, owner, _render(call_type, None))
table: Final = _sql_text(table_expr, scope)
if table is not None:
return _PrismaCallSite(location, method, owner, _render(call_type, table))
operation: Final = _POSTGRES_OPERATION_BY_CALL_TYPE.get(call_type)
rendered: Final = f"postgres.{operation.verb} {{relation}}" if operation is not None else None
return _PrismaCallSite(location, method, f"{owner}(bounded)", rendered)
def _decorated_site(function: str, method: str, location: str) -> _PrismaCallSite:
if function in _GENERIC_CRUD_HELPERS:
return _PrismaCallSite(location, method, "log_db_metrics(crud)", "postgres.{verb} {table_name}")
operation: Final = _POSTGRES_OPERATION_BY_CALL_TYPE.get(function)
return _PrismaCallSite(
location, method, "log_db_metrics", _render(function, None) if operation is not None else None
)
def _scope(module: _Module, parents: tuple[ast.AST, ...]) -> _Module:
function: Final = next((p for p in parents if isinstance(p, (ast.AsyncFunctionDef, ast.FunctionDef))), None)
if function is None:
return module
return _Module(module.path, module.tree, {**module.constants, **_assignments(function.body)})
def _engine_site(
call: ast.Call, module: _Module, parents: tuple[ast.AST, ...], method: str, accessor: str | None, location: str
) -> _PrismaCallSite:
if accessor is not None:
return _PrismaCallSite(location, method, "engine", _render(method, _MODEL_BY_ACCESSOR.get(accessor)))
scope: Final = _scope(module, parents)
statement: Final = call.args[0] if call.args else next((k.value for k in call.keywords if k.arg == "query"), None)
if isinstance(statement, ast.Name) and statement.id in _parameters(parents):
return _parameter_site(statement.id, module, parents, method, location)
rendered: Final = _render_statements(method, _sql_text(statement, scope) or "")
relative: Final = str(module.path.relative_to(_REPO))
if _RENDERED_NAME.fullmatch(rendered or "") is None and relative in _TRANSACTION_BODIES:
owner: Final = _TRANSACTION_BODIES[relative]
return _PrismaCallSite(location, method, f"transaction({owner})", _render(owner, None))
return _PrismaCallSite(location, method, "engine", rendered)
def _accessor(receiver: ast.expr) -> str | None:
if isinstance(receiver, ast.Attribute) and receiver.attr in _MODEL_BY_ACCESSOR:
return receiver.attr
return None
def _call_sites(module: _Module) -> Iterator[_PrismaCallSite]:
def walk(node: ast.AST, parents: tuple[ast.AST, ...]) -> Iterator[_PrismaCallSite]:
for child in ast.iter_child_nodes(node):
if isinstance(child, ast.Call) and isinstance(child.func, ast.Attribute):
method: Final = child.func.attr
accessor: Final = _accessor(child.func.value)
if method in _RAW_METHODS or (method in _MODEL_METHODS and accessor is not None):
location: Final = f"{module.path.relative_to(_REPO)}:{child.lineno}"
yield _wrapper_site(child, parents, method, location, module) or _engine_site(
child, module, parents, method, accessor, location
)
yield from walk(child, (child, *parents))
yield from walk(module.tree, ())
def prisma_call_sites() -> tuple[_PrismaCallSite, ...]:
return tuple(site for module in _modules() for site in _call_sites(module)) # comprehension-ok: flatten
def test_every_prisma_call_site_in_the_proxy_renders_a_bounded_postgres_span_name() -> None:
"""A raw ``query_raw``/``execute_raw``/``query_first`` or a direct model call that no producer
wraps is named by the engine from its payload; this scan replays that naming (and the wrappers')
statically so a new statement that would ship as a bare ``postgres.select`` or an unnamed
``postgres query_raw`` fails here rather than in a trace."""
sites = prisma_call_sites()
assert len(sites) >= 120, f"the scan lost the Prisma call sites: {len(sites)}"
unresolved = [site for site in sites if site.rendered is None]
assert unresolved == [], f"Prisma call sites whose span name cannot be resolved: {unresolved}"
half_named = [
site
for site in sites
if site.owner.startswith("engine")
and any(_RENDERED_NAME.fullmatch(name) is None for name in (site.rendered or "").split(" | "))
]
assert half_named == [], f"Prisma call sites that would ship a half-named or legacy span: {half_named}"
def test_every_model_in_the_prisma_schema_is_a_renderable_span_table() -> None:
schema: Final = (_REPO / "schema.prisma").read_text()
declared: Final = frozenset(re.findall(r"^model (\w+) \{", schema, re.MULTILINE))
assert declared == spans_mod._PRISMA_MODELS

View file

@ -1,3 +1,4 @@
from collections.abc import Awaitable, Callable
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -13,6 +14,7 @@ from litellm.proxy.db.proxy_worker_heartbeat import (
count_live_proxy_workers,
)
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
from tests.unit.proxy.db.fake_prisma_engine import engine_call
def _prisma():
@ -92,3 +94,21 @@ async def test_count_returns_unknown_for_a_malformed_row():
prisma = _prisma()
prisma.db.query_raw.return_value = [{"unexpected": "shape"}]
assert await count_live_proxy_workers(prisma) is None
@pytest.mark.asyncio
async def test_a_heartbeat_tick_renders_one_postgres_span_per_round_trip(
postgres_span_names: Callable[[], Awaitable[tuple[str, ...]]],
) -> None:
prisma = _prisma()
prisma.db.execute_raw = engine_call()
prisma.db.query_raw = engine_call([{"live_workers": 2}])
await ProxyWorkerHeartbeat(prisma_client=prisma, worker_id="worker-1").beat()
assert await count_live_proxy_workers(prisma) == 2
assert await postgres_span_names() == (
"postgres.upsert LiteLLM_ProxyWorkerHeartbeat",
"postgres.delete LiteLLM_ProxyWorkerHeartbeat",
"postgres.select LiteLLM_ProxyWorkerHeartbeat",
)

View file

@ -1,3 +1,4 @@
from collections.abc import Awaitable, Callable
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -7,6 +8,7 @@ from litellm.proxy.db.shadow_eval_funnel import (
flush_shadow_eval_funnel,
record_shadow_eval_funnel_event,
)
from tests.unit.proxy.db.fake_prisma_engine import engine_call
@pytest.fixture(autouse=True)
@ -93,3 +95,17 @@ def test_pending_count_feeds_the_drain_census():
record_shadow_eval_funnel_event("leg-1", "shed")
record_shadow_eval_funnel_event("leg-2", "unjudgeable")
assert pending_shadow_eval_funnel_events() == 3
@pytest.mark.asyncio
async def test_a_funnel_flush_renders_one_postgres_upsert_span_per_job(
postgres_span_names: Callable[[], Awaitable[tuple[str, ...]]],
) -> None:
record_shadow_eval_funnel_event("job-a", "not_sampled")
record_shadow_eval_funnel_event("job-b", "not_sampled")
prisma = MagicMock()
prisma.db.execute_raw = engine_call(1)
await flush_shadow_eval_funnel(prisma)
assert await postgres_span_names() == ("postgres.upsert LiteLLM_ShadowEvalFunnel",) * 2

View file

@ -5,6 +5,7 @@ plus the LiteLLM_DailyToolSpend rollup in one transaction.
"""
from types import SimpleNamespace
from collections.abc import Awaitable, Callable
from typing import Any
from unittest.mock import AsyncMock, MagicMock
@ -18,6 +19,8 @@ from litellm.proxy.db.spend_log_tool_index import (
flush_tool_usage_transactions,
response_tool_call_names,
)
from litellm.proxy.db.log_db_metrics import record_db_io
from tests.unit.proxy.db.fake_prisma_engine import engine_call
def _response_with_tool_calls(*names: str) -> SimpleNamespace:
@ -34,13 +37,13 @@ class _FakeBatcher:
return self
async def __aexit__(self, *args: Any) -> None:
return None
record_db_io()
def _prisma(batch_: MagicMock) -> MagicMock:
prisma = MagicMock()
prisma.db.batch_ = batch_
prisma.db.litellm_spendlogtoolindex.create_many = AsyncMock()
prisma.db.litellm_spendlogtoolindex.create_many = engine_call()
return prisma
@ -377,3 +380,18 @@ class TestFlushToolUsageTransactions:
with pytest.raises((httpx.ReadTimeout, httpx.ReadError)):
await flush_tool_usage_transactions(prisma_client=prisma, transactions=[_transaction("r1")])
prisma.db.batch_.assert_called_once()
@pytest.mark.asyncio
async def test_a_tool_usage_flush_renders_one_postgres_span_per_table_written(
postgres_span_names: Callable[[], Awaitable[tuple[str, ...]]],
) -> None:
prisma, _ = _prisma_with_batcher()
await flush_tool_usage_transactions(
prisma_client=prisma,
transactions=[_transaction("r1", tool_names=("tool_a",), spend=0.10, total_tokens=100)],
)
assert await postgres_span_names() == (
"postgres.insert LiteLLM_SpendLogToolIndex",
"postgres.upsert LiteLLM_DailyToolSpend",
)

View file

@ -1,7 +1,7 @@
import asyncio
import json
import logging
from datetime import datetime
from datetime import datetime, timedelta, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
@ -2877,3 +2877,42 @@ def test_autonomous_agent_cost_tracking_needs_no_human_or_virtual_key(agent_id:
assert _should_track_cost_callback(
user_api_key=None, user_id=None, team_id=None, end_user_id=None, call_type="acompletion", agent_id=agent_id
) is expected
_CALL_START: Final = datetime(2026, 1, 1, tzinfo=timezone.utc)
@pytest.mark.asyncio
async def test_track_cost_callback_enqueue_emits_no_service_span(): # test-quality-ok: no event is the behaviour
"""Spend tracking only enqueues into the in-memory spend queues here, no Postgres round
trip happens, so neither a ``batch_write_to_db`` nor a ``postgres`` service event may be
emitted; the flush that writes the queue emits its own table-named spans."""
from litellm.proxy.proxy_server import proxy_logging_obj
logger = _ProxyDBLogger()
kwargs = {
"model": "gpt-4",
"call_type": "acompletion",
"litellm_params": {
"metadata": {
"user_api_key": "hashed-key",
"user_api_key_user_id": "user-1",
"litellm_parent_otel_span": MagicMock(name="server-span"),
},
},
"standard_logging_object": {"response_cost": 0.1, "request_tags": None},
"stream": False,
}
success_hook = AsyncMock()
update_database = AsyncMock()
with (
patch.object(proxy_logging_obj.service_logging_obj, "async_service_success_hook", success_hook),
patch.object(proxy_logging_obj.db_spend_update_writer, "update_database", update_database),
):
await logger._PROXY_track_cost_callback(
kwargs=kwargs, completion_response=None, start_time=_CALL_START, end_time=_CALL_START + timedelta(seconds=1)
)
await asyncio.sleep(0)
assert update_database.await_count == 1, "the spend enqueue itself must still run"
assert success_hook.await_count == 0, [call.kwargs for call in success_hook.await_args_list]

View file

@ -6,6 +6,7 @@ DB walkers are tested against an AsyncMock Prisma client. Live end-to-end
proof-of-fix (real proxy + DB) is performed separately on the repro server.
"""
import asyncio
import json
from types import SimpleNamespace
from typing import Final
@ -13,12 +14,14 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm._service_logger import ServiceTypes
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
_V2_GCM_PREFIX,
encrypt_value_helper,
)
from litellm.proxy.management_endpoints import credential_migration as cm
from tests.unit.proxy.db.fake_prisma_engine import engine_call
@pytest.fixture
@ -134,7 +137,7 @@ def _config_prisma(record):
"""Build an AsyncMock prisma client whose litellm_config returns `record`."""
client = MagicMock()
client.db.litellm_config.find_unique = AsyncMock(return_value=record)
client.db.litellm_config.update = AsyncMock()
client.db.litellm_config.update = engine_call()
return client
@ -596,3 +599,49 @@ async def test_migrate_covered_tables_reports_real_counts(salt_key, monkeypatch)
assert by_loc["model_table"].migrated == 1 # was legacy pre, v2 post
assert by_loc["model_table"].legacy == 0 # residual zero after rotation
assert by_loc["model_table"].already_v2 == 1
def _db_service_hooks() -> tuple[AsyncMock, MagicMock]:
success: Final = AsyncMock()
service_logging: Final = MagicMock(async_service_success_hook=success, async_service_failure_hook=AsyncMock())
return success, MagicMock(service_logging_obj=service_logging)
@pytest.mark.asyncio
async def test_config_walker_write_emits_a_postgres_update_event_for_litellm_config(salt_key, monkeypatch):
_enable_aes(monkeypatch)
client = _config_prisma(SimpleNamespace(param_value={"api_key": _legacy_ct("vantage-secret", monkeypatch)}))
success, proxy_logging = _db_service_hooks()
monkeypatch.setattr(proxy_server, "proxy_logging_obj", proxy_logging)
await cm._migrate_config_settings_row(client, "vantage_settings", cm._VANTAGE_SENSITIVE, dry_run=False)
await asyncio.sleep(0)
event: Final = success.await_args.kwargs
assert (event["service"], event["call_type"], event["event_metadata"]) == (
ServiceTypes.DB,
"migrate_config_credentials",
{"table_name": "LiteLLM_Config"},
)
@pytest.mark.asyncio
async def test_sso_walker_write_emits_a_postgres_update_event_for_litellm_ssoconfig(salt_key, monkeypatch):
_enable_aes(monkeypatch)
client = MagicMock()
client.db.litellm_ssoconfig.find_unique = AsyncMock(
return_value=SimpleNamespace(sso_settings={"client_secret": _legacy_ct("client-secret", monkeypatch)})
)
client.db.litellm_ssoconfig.update = engine_call()
success, proxy_logging = _db_service_hooks()
monkeypatch.setattr(proxy_server, "proxy_logging_obj", proxy_logging)
await cm._migrate_sso_config(client, dry_run=False)
await asyncio.sleep(0)
event: Final = success.await_args.kwargs
assert (event["service"], event["call_type"], event["event_metadata"]) == (
ServiceTypes.DB,
"migrate_sso_credentials",
{"table_name": "LiteLLM_SSOConfig"},
)

View file

@ -1,6 +1,6 @@
import asyncio
import time
from collections.abc import Sequence
from collections.abc import Awaitable, Callable, Sequence
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import Final
@ -24,6 +24,7 @@ from litellm.proxy.spend_tracking.key_metadata_recovery import (
recover_key_owner_from_daily_spend,
)
from litellm.proxy.utils import hash_token
from litellm.proxy.db.log_db_metrics import record_db_io
def _digest_row(digest: str, key_alias: str | None, team_id: str | None, user_id: str | None) -> dict[str, str | None]:
@ -64,6 +65,7 @@ def _query_raw_by_table(
deleted_rows: Sequence[dict[str, str | None]],
) -> AsyncMock:
async def query_raw(sql: str, *params: object) -> list[dict[str, str | None]]:
record_db_io()
if '"LiteLLM_VerificationToken"' in sql:
return list(active_rows)
if '"LiteLLM_DeletedVerificationToken"' in sql:
@ -796,3 +798,19 @@ async def test_recover_key_owner_from_daily_spend_bounds_the_lookup_with_a_state
assert mock_prisma.db.tx.call_args.kwargs["timeout"] == timedelta(
milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS
)
@pytest.mark.asyncio
async def test_reverse_hash_recovery_renders_a_postgres_select_span_for_the_table_it_read(
postgres_span_names: Callable[[], Awaitable[tuple[str, ...]]],
) -> None:
double_hashed = hash_token("a" * 64)
mock_prisma = MagicMock()
mock_prisma.db.query_raw = _query_raw_by_table(
active_rows=[_digest_row(double_hashed, "batch-worker", "team-1", "alice")],
deleted_rows=[],
)
await recover_double_hashed_key_metadata(mock_prisma, {double_hashed})
assert await postgres_span_names() == ("postgres.select LiteLLM_VerificationToken",)

View file

@ -7,6 +7,7 @@ import logging
import math
import time
from contextlib import asynccontextmanager
from collections.abc import Awaitable, Callable
from datetime import datetime, timedelta, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock
@ -23,6 +24,7 @@ from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import (
SpendLogCleanup,
TableCleanupResult,
)
from tests.unit.proxy.db.fake_prisma_engine import engine_call
from litellm.proxy.db.db_transaction_queue.spend_log_cleanup_metrics import (
SpendLogCleanupMetrics,
)
@ -1293,7 +1295,9 @@ async def test_a_statement_timeout_is_clamped_to_the_budget_that_is_left():
)
# Only 2s of budget left against a 30s batch timeout.
await cleaner._execute_delete_batch(client, "DELETE FROM x", datetime.now(timezone.utc), time.monotonic() + 2)
await cleaner._execute_delete_batch(
client, "DELETE FROM x", datetime.now(timezone.utc), "LiteLLM_SpendLogs", time.monotonic() + 2
)
timeouts = [sql for sql in recorded if "statement_timeout" in sql]
assert timeouts, f"no statement timeout was issued: {recorded}"
@ -1667,3 +1671,18 @@ async def test_run_that_drains_every_table_logs_the_summary_at_info_not_warning(
assert len(summaries) == 1
assert summaries[0].levelno == logging.INFO
assert "outcome=completed" in summaries[0].getMessage()
@pytest.mark.asyncio
async def test_a_cleanup_delete_batch_renders_a_postgres_delete_span_for_its_table(
postgres_span_names: Callable[[], Awaitable[tuple[str, ...]]],
) -> None:
client = MagicMock()
_wire_tx(client.db)
client.db.execute_raw = engine_call(5)
await SpendLogCleanup(general_settings={})._execute_delete_batch(
client, "DELETE FROM x", datetime(2026, 1, 1, tzinfo=timezone.utc), "LiteLLM_SpendLogs", time.monotonic() + 2
)
assert await postgres_span_names() == ("postgres.delete LiteLLM_SpendLogs",)

View file

@ -8,16 +8,18 @@ Symbols pinned here:
from __future__ import annotations
import asyncio
import hashlib
import json
import logging
from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from litellm._service_logger import ServiceTypes
from litellm.proxy.db.log_db_metrics import record_db_io
from litellm.proxy.utils import PrismaClient
@ -296,3 +298,32 @@ async def test_delete_data_logs_and_raises_on_error(
)
with pytest.raises(RuntimeError, match="delete fail"):
await prisma_client.delete_data(tokens=["sk-x"])
@pytest.mark.asyncio
@pytest.mark.parametrize(
("method", "table_name", "model", "prisma_method", "kwargs"),
[
("insert_data", "key", "litellm_verificationtoken", "upsert", {"data": {"token": "sk-1"}}),
("update_data", "team", "litellm_teamtable", "upsert", {"team_id": "t1", "data": {"spend": 1.0}}),
("delete_data", "key", "litellm_verificationtoken", "delete_many", {"tokens": ["sk-1"]}),
],
)
async def test_a_write_that_reaches_the_engine_reports_the_table_to_the_service_logger(
prisma_client: PrismaClient, method: str, table_name: str, model: str, prisma_method: str, kwargs: dict[str, object]
) -> None:
async def _queried(*args: object, **kwds: object) -> SimpleNamespace:
record_db_io()
return SimpleNamespace(token="h", team_id="t1", spend=1.0)
setattr(getattr(prisma_client.db, model), prisma_method, AsyncMock(side_effect=_queried))
success_hook = AsyncMock()
with patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
MagicMock(service_logging_obj=MagicMock(async_service_success_hook=success_hook)),
):
await getattr(prisma_client, method)(table_name=table_name, **kwargs)
await asyncio.sleep(0)
events = [c.kwargs for c in success_hook.await_args_list if c.kwargs["service"] == ServiceTypes.DB]
assert [(e["call_type"], e["event_metadata"]) for e in events] == [(method, {"table_name": table_name})]

View file

@ -12,8 +12,10 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from prisma.errors import PrismaError
import litellm
from litellm._service_logger import ServiceTypes
from litellm.proxy._types import AlertType, CallInfo
@ -252,6 +254,40 @@ async def test_failure_handler_logs_db_error_and_calls_service_logging(proxy_log
}
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", ["get_data", "insert_data", "update_data", "delete_data"])
async def test_failure_handler_alerts_but_leaves_prisma_error_event_to_log_db_metrics(
proxy_logging, monkeypatch, call_type
):
proxy_logging.alert_types = [AlertType.db_exceptions]
proxy_logging.alerting_handler = AsyncMock()
proxy_logging.service_logging_obj = MagicMock(async_service_failure_hook=AsyncMock())
monkeypatch.setattr(litellm.utils, "capture_exception", None)
await proxy_logging.failure_handler(
original_exception=PrismaError("connection reset"), duration=1.0, call_type=call_type
)
snapshot = {
"alerting_handler_scheduled": proxy_logging.alerting_handler.called,
"service_failure_called": proxy_logging.service_logging_obj.async_service_failure_hook.called,
}
assert snapshot == {"alerting_handler_scheduled": True, "service_failure_called": False}
@pytest.mark.asyncio
async def test_failure_handler_still_emits_db_event_for_wrapped_insert_error(proxy_logging, monkeypatch):
proxy_logging.alert_types = [AlertType.db_exceptions]
proxy_logging.alerting_handler = AsyncMock()
proxy_logging.service_logging_obj = MagicMock(async_service_failure_hook=AsyncMock())
monkeypatch.setattr(litellm.utils, "capture_exception", None)
await proxy_logging.failure_handler(
original_exception=HTTPException(status_code=400, detail={"error": "Foreign Key Constraint failed"}),
duration=1.0,
call_type="insert_data",
)
call_kwargs = proxy_logging.service_logging_obj.async_service_failure_hook.call_args.kwargs
assert (call_kwargs["service"], call_kwargs["call_type"]) == (ServiceTypes.DB, "insert_data")
@pytest.mark.asyncio
async def test_failure_handler_with_capture_exception_invoked(proxy_logging, monkeypatch):
proxy_logging.alert_types = [AlertType.db_exceptions]

View file

@ -238,7 +238,7 @@ async def test_service_span_not_duplicated_for_string_and_instance(monkeypatch):
parent.end()
db_spans = [
s for s in exporter.get_finished_spans() if s.name == "postgres get_user_object"
s for s in exporter.get_finished_spans() if s.name == "postgres.select LiteLLM_UserTable"
]
assert len(db_spans) == 1
@ -274,7 +274,7 @@ async def test_service_failure_span_not_duplicated_for_string_and_instance(
parent.end()
db_spans = [
s for s in exporter.get_finished_spans() if s.name == "postgres get_user_object"
s for s in exporter.get_finished_spans() if s.name == "postgres.select LiteLLM_UserTable"
]
assert len(db_spans) == 1