diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py index 1ab173a915a..cf22488edcb 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py @@ -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, diff --git a/litellm/integrations/otel/README.md b/litellm/integrations/otel/README.md index cdd1b9ef95b..38f96f954f1 100644 --- a/litellm/integrations/otel/README.md +++ b/litellm/integrations/otel/README.md @@ -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 diff --git a/litellm/integrations/otel/mappers/genai.py b/litellm/integrations/otel/mappers/genai.py index 7947dbaae29..75a1098819d 100644 --- a/litellm/integrations/otel/mappers/genai.py +++ b/litellm/integrations/otel/mappers/genai.py @@ -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 diff --git a/litellm/integrations/otel/model/db_endpoint.py b/litellm/integrations/otel/model/db_endpoint.py index 562162a8f31..7a9c9fedbb7 100644 --- a/litellm/integrations/otel/model/db_endpoint.py +++ b/litellm/integrations/otel/model/db_endpoint.py @@ -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), diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index 2d05754d37b..38cf7266593 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -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" diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index 6ddc824235b..2cc8e035ebd 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -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() diff --git a/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py b/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py index 62f4489a784..eb02ed64e91 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py @@ -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 diff --git a/litellm/proxy/auth/auth_object_prefetch.py b/litellm/proxy/auth/auth_object_prefetch.py index d4a91d1f194..fd5b99951f4 100644 --- a/litellm/proxy/auth/auth_object_prefetch.py +++ b/litellm/proxy/auth/auth_object_prefetch.py @@ -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 diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 950aab723d4..b7620c5f8bd 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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. diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 72cac1224f6..61330d12ac3 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -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: diff --git a/litellm/proxy/db/autorouter_session_rollup.py b/litellm/proxy/db/autorouter_session_rollup.py index 350e4da231c..21aa0ca4d29 100644 --- a/litellm/proxy/db/autorouter_session_rollup.py +++ b/litellm/proxy/db/autorouter_session_rollup.py @@ -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( diff --git a/litellm/proxy/db/baseline_accounting.py b/litellm/proxy/db/baseline_accounting.py index 9536f8d740a..625f5d03125 100644 --- a/litellm/proxy/db/baseline_accounting.py +++ b/litellm/proxy/db/baseline_accounting.py @@ -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) diff --git a/litellm/proxy/db/budget_window_spend_writer.py b/litellm/proxy/db/budget_window_spend_writer.py index 8cf2f737063..ca7c3341197 100644 --- a/litellm/proxy/db/budget_window_spend_writer.py +++ b/litellm/proxy/db/budget_window_spend_writer.py @@ -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)), + ) diff --git a/litellm/proxy/db/create_views.py b/litellm/proxy/db/create_views.py index f7131091c0b..d51e39596c6 100644 --- a/litellm/proxy/db/create_views.py +++ b/litellm/proxy/db/create_views.py @@ -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): diff --git a/litellm/proxy/db/db_span.py b/litellm/proxy/db/db_span.py new file mode 100644 index 00000000000..0cbd7c8db10 --- /dev/null +++ b/litellm/proxy/db/db_span.py @@ -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() diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index d9a60742c9d..5fdde118fb9 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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() diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py index 679e286cedd..4ff6d926a99 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py +++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py @@ -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 diff --git a/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py b/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py index b756bfeb6f6..f2fa2e4a124 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py +++ b/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py @@ -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( """ diff --git a/litellm/proxy/db/gateway_request_tracking.py b/litellm/proxy/db/gateway_request_tracking.py index 2dbcd1ccdf0..e3483ca3215 100644 --- a/litellm/proxy/db/gateway_request_tracking.py +++ b/litellm/proxy/db/gateway_request_tracking.py @@ -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) diff --git a/litellm/proxy/db/health_check_latest.py b/litellm/proxy/db/health_check_latest.py index 21438f095bb..a825a23841e 100644 --- a/litellm/proxy/db/health_check_latest.py +++ b/litellm/proxy/db/health_check_latest.py @@ -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) diff --git a/litellm/proxy/db/log_db_metrics.py b/litellm/proxy/db/log_db_metrics.py index 8c3d757838d..5895610dadf 100644 --- a/litellm/proxy/db/log_db_metrics.py +++ b/litellm/proxy/db/log_db_metrics.py @@ -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 diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index 28e03d6dbb5..d73234b6de3 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -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() diff --git a/litellm/proxy/db/prisma_query_span.py b/litellm/proxy/db/prisma_query_span.py new file mode 100644 index 00000000000..70c8e2ec3f9 --- /dev/null +++ b/litellm/proxy/db/prisma_query_span.py @@ -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) diff --git a/litellm/proxy/db/proxy_worker_heartbeat.py b/litellm/proxy/db/proxy_worker_heartbeat.py index 990ff48eb18..02015c699c9 100644 --- a/litellm/proxy/db/proxy_worker_heartbeat.py +++ b/litellm/proxy/db/proxy_worker_heartbeat.py @@ -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) diff --git a/litellm/proxy/db/shadow_eval_funnel.py b/litellm/proxy/db/shadow_eval_funnel.py index 345a85d9fdd..0986a502c74 100644 --- a/litellm/proxy/db/shadow_eval_funnel.py +++ b/litellm/proxy/db/shadow_eval_funnel.py @@ -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", diff --git a/litellm/proxy/db/spend_log_tool_index.py b/litellm/proxy/db/spend_log_tool_index.py index 6d012c64b95..b0d8bb9aba1 100644 --- a/litellm/proxy/db/spend_log_tool_index.py +++ b/litellm/proxy/db/spend_log_tool_index.py @@ -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) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index dc2723c4267..f937b439042 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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 diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 35ff9186f5e..33fb069afbd 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -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: diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py index 73e26aa827a..f9c09128f66 100644 --- a/litellm/proxy/management_endpoints/credential_migration.py +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -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 diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 8943f5a7416..fe976c861e5 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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"], diff --git a/litellm/proxy/management_helpers/access_group_team_sync.py b/litellm/proxy/management_helpers/access_group_team_sync.py index 664e36c9f10..555481d06a4 100644 --- a/litellm/proxy/management_helpers/access_group_team_sync.py +++ b/litellm/proxy/management_helpers/access_group_team_sync.py @@ -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) diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 560363ca7d7..225e96179ff 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -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", diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index eed2c257e87..4ebff86ac06 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 diff --git a/litellm/repositories/daily_activity_repository.py b/litellm/repositories/daily_activity_repository.py index 2e34582091a..2230a9e7fa8 100644 --- a/litellm/repositories/daily_activity_repository.py +++ b/litellm/repositories/daily_activity_repository.py @@ -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) diff --git a/tests/e2e/logging/test_otel_trace_e2e.py b/tests/e2e/logging/test_otel_trace_e2e.py index 8d8595cf221..4902b0703c3 100644 --- a/tests/e2e/logging/test_otel_trace_e2e.py +++ b/tests/e2e/logging/test_otel_trace_e2e.py @@ -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 ". 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 ". 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", ( diff --git a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_endpoints.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_endpoints.py index 7b32d9e8c44..43f13e0ebd7 100644 --- a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_endpoints.py +++ b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_endpoints.py @@ -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"}, + ) diff --git a/tests/unit/integrations/otel/test_db_endpoint.py b/tests/unit/integrations/otel/test_db_endpoint.py index 5ab0a927b52..6c55b7c4dec 100644 --- a/tests/unit/integrations/otel/test_db_endpoint.py +++ b/tests/unit/integrations/otel/test_db_endpoint.py @@ -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", diff --git a/tests/unit/integrations/otel/test_otel_v2_logger.py b/tests/unit/integrations/otel/test_otel_v2_logger.py index b7bf678d1fe..c0fa890e8a2 100644 --- a/tests/unit/integrations/otel/test_otel_v2_logger.py +++ b/tests/unit/integrations/otel/test_otel_v2_logger.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py b/tests/unit/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py index b6c946b95fa..0dc52b13950 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py @@ -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"}, + ) diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index a7219ac059b..fc8bc289735 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -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"}, + ) diff --git a/tests/unit/proxy/common_utils/test_reset_budget_job.py b/tests/unit/proxy/common_utils/test_reset_budget_job.py index 131db55ee01..8308d3a7664 100644 --- a/tests/unit/proxy/common_utils/test_reset_budget_job.py +++ b/tests/unit/proxy/common_utils/test_reset_budget_job.py @@ -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,) diff --git a/tests/unit/proxy/conftest.py b/tests/unit/proxy/conftest.py index 6dec588b763..cb7e9969bca 100644 --- a/tests/unit/proxy/conftest.py +++ b/tests/unit/proxy/conftest.py @@ -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 diff --git a/tests/unit/proxy/db/db_transaction_queue/test_spend_logs_partition_manager.py b/tests/unit/proxy/db/db_transaction_queue/test_spend_logs_partition_manager.py index 609dd13afc2..0d4751346d9 100644 --- a/tests/unit/proxy/db/db_transaction_queue/test_spend_logs_partition_manager.py +++ b/tests/unit/proxy/db/db_transaction_queue/test_spend_logs_partition_manager.py @@ -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",) diff --git a/tests/unit/proxy/db/fake_prisma_engine.py b/tests/unit/proxy/db/fake_prisma_engine.py new file mode 100644 index 00000000000..4221eeced5a --- /dev/null +++ b/tests/unit/proxy/db/fake_prisma_engine.py @@ -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) diff --git a/tests/unit/proxy/db/test_autorouter_session_rollup.py b/tests/unit/proxy/db/test_autorouter_session_rollup.py index 835399c568f..659d29cda16 100644 --- a/tests/unit/proxy/db/test_autorouter_session_rollup.py +++ b/tests/unit/proxy/db/test_autorouter_session_rollup.py @@ -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,) diff --git a/tests/unit/proxy/db/test_budget_window_spend_writer.py b/tests/unit/proxy/db/test_budget_window_spend_writer.py index 130f0c56ccf..6fc438feee8 100644 --- a/tests/unit/proxy/db/test_budget_window_spend_writer.py +++ b/tests/unit/proxy/db/test_budget_window_spend_writer.py @@ -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",) diff --git a/tests/unit/proxy/db/test_db_span.py b/tests/unit/proxy/db/test_db_span.py new file mode 100644 index 00000000000..b707eb2710d --- /dev/null +++ b/tests/unit/proxy/db/test_db_span.py @@ -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) diff --git a/tests/unit/proxy/db/test_db_spend_update_writer.py b/tests/unit/proxy/db/test_db_spend_update_writer.py index 4de90d5d7f7..fe5e31d00b2 100644 --- a/tests/unit/proxy/db/test_db_spend_update_writer.py +++ b/tests/unit/proxy/db/test_db_spend_update_writer.py @@ -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), diff --git a/tests/unit/proxy/db/test_gateway_request_tracking.py b/tests/unit/proxy/db/test_gateway_request_tracking.py index 045261e2d53..a6689b38039 100644 --- a/tests/unit/proxy/db/test_gateway_request_tracking.py +++ b/tests/unit/proxy/db/test_gateway_request_tracking.py @@ -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",) diff --git a/tests/unit/proxy/db/test_health_check_latest.py b/tests/unit/proxy/db/test_health_check_latest.py index 6322891ae9e..29063d7b060 100644 --- a/tests/unit/proxy/db/test_health_check_latest.py +++ b/tests/unit/proxy/db/test_health_check_latest.py @@ -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",) diff --git a/tests/unit/proxy/db/test_log_db_metrics.py b/tests/unit/proxy/db/test_log_db_metrics.py index 1ec1afe6106..8aa8d0dabc8 100644 --- a/tests/unit/proxy/db/test_log_db_metrics.py +++ b/tests/unit/proxy/db/test_log_db_metrics.py @@ -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") diff --git a/tests/unit/proxy/db/test_prisma_query_span.py b/tests/unit/proxy/db/test_prisma_query_span.py new file mode 100644 index 00000000000..674679f8a18 --- /dev/null +++ b/tests/unit/proxy/db/test_prisma_query_span.py @@ -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 diff --git a/tests/unit/proxy/db/test_proxy_worker_heartbeat.py b/tests/unit/proxy/db/test_proxy_worker_heartbeat.py index 33ae6190411..967e4d3471f 100644 --- a/tests/unit/proxy/db/test_proxy_worker_heartbeat.py +++ b/tests/unit/proxy/db/test_proxy_worker_heartbeat.py @@ -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", + ) diff --git a/tests/unit/proxy/db/test_shadow_eval_funnel.py b/tests/unit/proxy/db/test_shadow_eval_funnel.py index 065d4e6ca1a..59d099b89fd 100644 --- a/tests/unit/proxy/db/test_shadow_eval_funnel.py +++ b/tests/unit/proxy/db/test_shadow_eval_funnel.py @@ -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 diff --git a/tests/unit/proxy/db/test_spend_log_tool_index.py b/tests/unit/proxy/db/test_spend_log_tool_index.py index 282c2a7cfaa..610faebe17c 100644 --- a/tests/unit/proxy/db/test_spend_log_tool_index.py +++ b/tests/unit/proxy/db/test_spend_log_tool_index.py @@ -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", + ) diff --git a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py index 28376be64b6..7afa275c801 100644 --- a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py @@ -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] diff --git a/tests/unit/proxy/management_endpoints/test_credential_migration.py b/tests/unit/proxy/management_endpoints/test_credential_migration.py index c638f1b30b2..ac5a45499d3 100644 --- a/tests/unit/proxy/management_endpoints/test_credential_migration.py +++ b/tests/unit/proxy/management_endpoints/test_credential_migration.py @@ -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"}, + ) diff --git a/tests/unit/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/unit/proxy/spend_tracking/test_key_metadata_recovery.py index 7c7a0b31348..5d02b289360 100644 --- a/tests/unit/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/unit/proxy/spend_tracking/test_key_metadata_recovery.py @@ -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",) diff --git a/tests/unit/proxy/test_spend_log_cleanup.py b/tests/unit/proxy/test_spend_log_cleanup.py index 399c76d97c1..05bf9fff9a0 100644 --- a/tests/unit/proxy/test_spend_log_cleanup.py +++ b/tests/unit/proxy/test_spend_log_cleanup.py @@ -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",) diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_prisma_client_writes.py b/tests/unit/proxy/utils/prisma_and_spend/test_prisma_client_writes.py index 6e69444a1b5..747e8da8043 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_prisma_client_writes.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_prisma_client_writes.py @@ -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})] diff --git a/tests/unit/proxy/utils/proxy_logging/test_alerting.py b/tests/unit/proxy/utils/proxy_logging/test_alerting.py index 77c0f71dbf9..43ee1094330 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_alerting.py +++ b/tests/unit/proxy/utils/proxy_logging/test_alerting.py @@ -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] diff --git a/tests/unit/test_service_logger.py b/tests/unit/test_service_logger.py index 2ee04cc3b96..3d74642a03e 100644 --- a/tests/unit/test_service_logger.py +++ b/tests/unit/test_service_logger.py @@ -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