mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(otel): name postgres service spans by operation and table (#44240)
Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
564d236985
commit
797353f13a
62 changed files with 2452 additions and 312 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
99
litellm/proxy/db/db_span.py
Normal file
99
litellm/proxy/db/db_span.py
Normal file
|
|
@ -0,0 +1,99 @@
|
|||
"""A ``ServiceTypes.DB`` event around Prisma I/O that ``@log_db_metrics`` cannot wrap.
|
||||
|
||||
The spend flush, the spend-log batch insert and the background jobs run raw
|
||||
``prisma_client.db`` statements and transactions, often inside retry loops, so the
|
||||
decorator (one event per decorated coroutine) cannot name the table each round
|
||||
trip touches. ``db_span`` emits one success or failure event per round trip,
|
||||
carrying the raw ``call_type`` for the metric labels and the Prisma model on
|
||||
``table_name`` so OTel renders ``postgres.{verb} {table}``; ``db_spanned`` is the
|
||||
same event around a thunk, for the retry helpers that take one. The outermost
|
||||
producer owns the event: a ``db_span`` nested in another ``db_span`` or in a
|
||||
decorated helper emits nothing, so one transaction stays one span, and a block
|
||||
whose Prisma client never reached the engine emits nothing at all.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime
|
||||
from typing import Final, TypeVar
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._service_logger import ServiceLogging, ServiceTypes
|
||||
from litellm.proxy.db.log_db_metrics import _is_exception_related_to_db, claim_db_io, db_io_claimed
|
||||
|
||||
_T = TypeVar("_T")
|
||||
|
||||
|
||||
def _service_logging() -> ServiceLogging | None:
|
||||
try:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
except ImportError:
|
||||
return None
|
||||
return proxy_logging_obj.service_logging_obj
|
||||
|
||||
|
||||
def _event_metadata(table: str | None, operation: str | None) -> dict[str, str]:
|
||||
pairs: Final = (("table_name", table), ("db_operation", operation))
|
||||
return {key: value for key, value in pairs if value is not None}
|
||||
|
||||
|
||||
async def _emit_failure(
|
||||
service_logging: ServiceLogging,
|
||||
call_type: str,
|
||||
event_metadata: Mapping[str, str],
|
||||
start_time: datetime,
|
||||
error: Exception,
|
||||
) -> None:
|
||||
end_time: Final = datetime.now()
|
||||
try:
|
||||
await service_logging.async_service_failure_hook(
|
||||
error=error,
|
||||
service=ServiceTypes.DB,
|
||||
call_type=call_type,
|
||||
parent_otel_span=None,
|
||||
duration=(end_time - start_time).total_seconds(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
event_metadata=dict(event_metadata),
|
||||
)
|
||||
except Exception as hook_error:
|
||||
verbose_proxy_logger.debug("db_span: failure hook raised for %s: %s", call_type, hook_error)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def db_span(call_type: str, table: str | None, operation: str | None = None) -> AsyncGenerator[None]:
|
||||
if db_io_claimed():
|
||||
yield
|
||||
return
|
||||
service_logging: Final = _service_logging()
|
||||
start_time: Final = datetime.now()
|
||||
event_metadata: Final = _event_metadata(table, operation)
|
||||
with claim_db_io() as witness:
|
||||
try:
|
||||
yield
|
||||
except Exception as e:
|
||||
if service_logging is not None and _is_exception_related_to_db(e):
|
||||
await _emit_failure(service_logging, call_type, event_metadata, start_time, e)
|
||||
raise
|
||||
if service_logging is None or not witness.touched:
|
||||
return
|
||||
end_time_ok: Final = datetime.now()
|
||||
asyncio.create_task(
|
||||
service_logging.async_service_success_hook(
|
||||
service=ServiceTypes.DB,
|
||||
call_type=call_type,
|
||||
parent_otel_span=None,
|
||||
duration=(end_time_ok - start_time).total_seconds(),
|
||||
start_time=start_time,
|
||||
end_time=end_time_ok,
|
||||
event_metadata=event_metadata,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def db_spanned(call_type: str, table: str | None, load: Callable[[], Awaitable[_T]]) -> _T:
|
||||
async with db_span(call_type, table):
|
||||
return await load()
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
136
litellm/proxy/db/prisma_query_span.py
Normal file
136
litellm/proxy/db/prisma_query_span.py
Normal file
|
|
@ -0,0 +1,136 @@
|
|||
"""Name the Prisma round trips that no producer claims.
|
||||
|
||||
``_TrackedPrismaEngine.query`` sees every statement the proxy sends to the query
|
||||
engine. When neither ``@log_db_metrics`` nor ``db_span`` encloses the call, the
|
||||
engine names the event itself from the GraphQL payload Prisma built: the root
|
||||
field (``findUniqueLiteLLM_VerificationToken``, ``createOneLiteLLM_SpendLogs``)
|
||||
carries the method and the model, and for ``queryRaw``/``executeRaw`` the leading
|
||||
SQL keyword gives the verb and the first ``schema.prisma`` relation the statement
|
||||
names gives the table. Only bounded names ever leave this module: relations
|
||||
declared in the schema, the spend views, ``pg_catalog`` for catalog probes and
|
||||
the setting a ``SET`` statement targets. No SQL text or values.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.integrations.otel.model.spans import PG_CATALOG, PRISMA_RELATIONS
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PrismaQuery:
|
||||
"""What one engine round trip is, for the ``ServiceTypes.DB`` event: the raw method,
|
||||
the SQL verb (``None`` when the statement is not one this module knows) and the relation."""
|
||||
|
||||
call_type: str
|
||||
operation: str | None
|
||||
table: str | None
|
||||
|
||||
|
||||
UNKNOWN_PRISMA_QUERY: Final = PrismaQuery("prisma_query", None, None)
|
||||
|
||||
_ROOT_FIELD: Final = re.compile(r"result:\s*(\w+)")
|
||||
_RAW_SQL: Final = re.compile(r'query:\s*"((?:[^"\\]|\\.)*)"')
|
||||
_LEADING_KEYWORD: Final = re.compile(r"(?:\\[nrt]|\s|\()*(\w+)")
|
||||
_SETTING: Final = re.compile(r"(?:\\[nrt]|\s)*SET\s+(?:LOCAL\s+|SESSION\s+)?([A-Za-z_.]+)", re.IGNORECASE)
|
||||
_CATALOG: Final = re.compile(r"\bpg_\w+|\bto_regclass\b|\binformation_schema\b|\bcurrent_setting\s*\(|^\s*SHOW\b")
|
||||
_PROBE: Final = re.compile(r"(?:\\[nrt]|\s)*SELECT\s+\d+\s*;?(?:\\[nrt]|\s)*$", re.IGNORECASE)
|
||||
_CTE_WRITE: Final = re.compile(r"\b(UPDATE|INSERT|DELETE)\s+(?:INTO\s+|FROM\s+)?(?:\\?\")", re.IGNORECASE)
|
||||
_RELATION: Final = re.compile(
|
||||
r"\b(?:" + "|".join(sorted(map(re.escape, PRISMA_RELATIONS), key=len, reverse=True)) + r")\b"
|
||||
)
|
||||
_MODEL_ACTIONS: Final[Mapping[str, tuple[str, str]]] = MappingProxyType(
|
||||
{
|
||||
"findUnique": ("find_unique", "select"),
|
||||
"findFirst": ("find_first", "select"),
|
||||
"findMany": ("find_many", "select"),
|
||||
"aggregate": ("count", "select"),
|
||||
"groupBy": ("group_by", "select"),
|
||||
"createOne": ("create", "insert"),
|
||||
"createMany": ("create_many", "insert"),
|
||||
"updateOne": ("update", "update"),
|
||||
"updateMany": ("update_many", "update"),
|
||||
"deleteOne": ("delete", "delete"),
|
||||
"deleteMany": ("delete_many", "delete"),
|
||||
"upsertOne": ("upsert", "upsert"),
|
||||
}
|
||||
)
|
||||
_RAW_ACTIONS: Final[Mapping[str, str]] = MappingProxyType({"queryRaw": "query_raw", "executeRaw": "execute_raw"})
|
||||
_VERB_BY_KEYWORD: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
"SELECT": "select",
|
||||
"WITH": "select",
|
||||
"INSERT": "insert",
|
||||
"UPDATE": "update",
|
||||
"DELETE": "delete",
|
||||
"CREATE": "ddl",
|
||||
"ALTER": "ddl",
|
||||
"DROP": "ddl",
|
||||
"REFRESH": "ddl",
|
||||
"TRUNCATE": "delete",
|
||||
"SET": "set",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def sql_relation(sql: str) -> str | None:
|
||||
"""The first schema relation (model or spend view) the statement names, ``pg_catalog``
|
||||
for a statement that only reads the system catalog, else ``None``."""
|
||||
relation: Final = _RELATION.search(sql)
|
||||
if relation is not None:
|
||||
return relation.group(0)
|
||||
return PG_CATALOG if _CATALOG.search(sql) else None
|
||||
|
||||
|
||||
def sql_operation(sql: str) -> tuple[str | None, str | None]:
|
||||
"""``(verb, target)`` for a raw statement: the SQL verb from its leading keyword and the
|
||||
relation it names, or for ``SET`` the setting it changes."""
|
||||
if _PROBE.match(sql):
|
||||
return "ping", None
|
||||
keyword: Final = _LEADING_KEYWORD.match(sql)
|
||||
leading: Final = keyword.group(1).upper() if keyword is not None else ""
|
||||
cte_write: Final = _CTE_WRITE.search(sql) if leading == "WITH" else None
|
||||
verb: Final = _VERB_BY_KEYWORD[cte_write.group(1).upper()] if cte_write else _VERB_BY_KEYWORD.get(leading)
|
||||
if verb != "set":
|
||||
return verb, sql_relation(sql)
|
||||
setting: Final = _SETTING.match(sql)
|
||||
return verb, setting.group(1).lower() if setting is not None else None
|
||||
|
||||
|
||||
def _query_text(content: str) -> str:
|
||||
try:
|
||||
payload: Final[object] = json.loads(content)
|
||||
except ValueError:
|
||||
return content
|
||||
query: Final = payload.get("query") if isinstance(payload, dict) else None
|
||||
return query if isinstance(query, str) else content
|
||||
|
||||
|
||||
def _model_query(root_field: str) -> PrismaQuery | None:
|
||||
action: Final = next((prefix for prefix in _MODEL_ACTIONS if root_field.startswith(prefix)), None)
|
||||
if action is None:
|
||||
return None
|
||||
call_type, verb = _MODEL_ACTIONS[action]
|
||||
model: Final = root_field.removeprefix(action).removesuffix("OrThrow")
|
||||
return PrismaQuery(call_type, verb, model) if model in PRISMA_RELATIONS else None
|
||||
|
||||
|
||||
def parse_prisma_query(content: str) -> PrismaQuery:
|
||||
"""The round trip behind one query-engine payload, ``UNKNOWN_PRISMA_QUERY`` when the
|
||||
payload is not a shape this module knows (which renders ``postgres prisma_query``)."""
|
||||
query: Final = _query_text(content)
|
||||
root: Final = _ROOT_FIELD.search(query)
|
||||
if root is None:
|
||||
return UNKNOWN_PRISMA_QUERY
|
||||
raw_call_type: Final = _RAW_ACTIONS.get(root.group(1))
|
||||
if raw_call_type is None:
|
||||
return _model_query(root.group(1)) or UNKNOWN_PRISMA_QUERY
|
||||
sql: Final = _RAW_SQL.search(query, root.end())
|
||||
verb, target = sql_operation(sql.group(1)) if sql is not None else (None, None)
|
||||
return PrismaQuery(raw_call_type, verb, target)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -3,8 +3,7 @@
|
|||
Covers logging.otel.success.exports_metric: a successful non-streaming call must
|
||||
land at the OTEL destination as ONE connected trace - a single root SERVER span
|
||||
with the auth phase and db lookups under it, the gen-AI CLIENT span parented
|
||||
into the same tree, and the cost write either under it or as the root of its
|
||||
own trace linked back to the request span. The regression this pins: the proxy publishing
|
||||
into the same tree. The regression this pins: the proxy publishing
|
||||
the global TracerProvider before callbacks init made server spans export through
|
||||
a different provider than the preset's gen-AI spans, so the destination received
|
||||
the gen-AI span alone, dangling (fixed in #30590; verified failing at its parent
|
||||
|
|
@ -29,14 +28,13 @@ from e2e_config import CHEAP_ANTHROPIC_MODEL, CHEAP_OPENAI_MODEL, OTEL_EXPORTER_
|
|||
from lifecycle import ResourceManager
|
||||
from logging_client import INVALID_UPSTREAM_API_KEY, LoggingClient, first_ok, readiness_details_body
|
||||
from models import LiteLLMParamsBody
|
||||
from otel_client import CallTraces, JaegerSpan, JaegerTrace, OtelReader, root_span
|
||||
from otel_client import CallTraces, JaegerSpan, JaegerTrace, OtelReader
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
MODEL = CHEAP_ANTHROPIC_MODEL
|
||||
COST_SPAN = "batch_write_to_db _PROXY_track_cost_callback"
|
||||
DB_SPAN_PREFIX = "postgres "
|
||||
DB_SPAN_PREFIX = "postgres."
|
||||
#: The active OTEL v2 logger's name in /health/readiness/details success_callbacks.
|
||||
OTEL_V2_LOGGER_NAME = "OpenTelemetryV2"
|
||||
|
||||
|
|
@ -79,12 +77,12 @@ def _chain_reaches(span_id: str, root_id: str, trace: JaegerTrace) -> bool:
|
|||
return False
|
||||
|
||||
|
||||
def _assert_complete_trace(traces: CallTraces, *, route: str, genai_span: str, require_cost_span: bool = True) -> None:
|
||||
def _assert_complete_trace(traces: CallTraces, *, route: str, genai_span: str) -> None:
|
||||
"""The enforced behavior: the destination holds exactly one call-id-tagged
|
||||
trace for the call, rooted at the SERVER span, with auth/db children and
|
||||
the gen-AI span all connected into that one tree - no dangling parent
|
||||
references - and the cost write either in that trace or as the root of
|
||||
its own trace linked FOLLOWS_FROM to the request SERVER span."""
|
||||
references. The spend enqueue after the response does no I/O, so it emits
|
||||
no span; the flush that writes spend is its own background trace."""
|
||||
hits = traces.hits
|
||||
assert hits, (
|
||||
"no trace for this call arrived at the destination within the deadline "
|
||||
|
|
@ -121,21 +119,6 @@ def _assert_complete_trace(traces: CallTraces, *, route: str, genai_span: str, r
|
|||
assert any(name.startswith(DB_SPAN_PREFIX) for name in names), (
|
||||
f"no db ('{DB_SPAN_PREFIX}*') span in the trace; spans: {names}"
|
||||
)
|
||||
if require_cost_span and COST_SPAN not in names:
|
||||
cost_traces = [t for t in traces.linked if (r := root_span(t)) is not None and r.operation_name == COST_SPAN]
|
||||
assert len(cost_traces) == 1, (
|
||||
f"cost write span {COST_SPAN!r} reached neither the request trace nor its own "
|
||||
f"trace linked to the request SERVER span; request spans: {names}; "
|
||||
f"linked traces: {[(t.trace_id, t.span_names()) for t in traces.linked]}"
|
||||
)
|
||||
cost_root = root_span(cost_traces[0])
|
||||
assert cost_root is not None, f"cost write trace has no single root; spans: {cost_traces[0].span_names()}"
|
||||
link = next(ref for ref in cost_root.references if ref.span_id == root.span_id)
|
||||
assert link.ref_type == "FOLLOWS_FROM" and link.trace_id == trace.trace_id, (
|
||||
f"the cost write trace's root must reference the request SERVER span FOLLOWS_FROM, "
|
||||
f"got refType={link.ref_type!r} traceID={link.trace_id!r} (request trace {trace.trace_id})"
|
||||
)
|
||||
|
||||
genai = next((span for span in trace.spans if span.operation_name == genai_span), None)
|
||||
assert genai is not None, f"gen-AI span {genai_span!r} missing; spans: {names}"
|
||||
assert genai.kind == "client", f"gen-AI span must have kind=client, got {genai.kind!r}"
|
||||
|
|
@ -145,14 +128,11 @@ def _assert_complete_trace(traces: CallTraces, *, route: str, genai_span: str, r
|
|||
)
|
||||
|
||||
|
||||
def _poll(
|
||||
otel_reader: OtelReader, *, call_id: str, route: str, genai_span: str, require_cost_span: bool = True
|
||||
) -> CallTraces:
|
||||
def _poll(otel_reader: OtelReader, *, call_id: str, route: str, genai_span: str) -> CallTraces:
|
||||
return otel_reader.poll_traces_for_call(
|
||||
call_id=call_id,
|
||||
settled_names={f"POST {route}", f"auth {route}", genai_span},
|
||||
settled_prefixes={DB_SPAN_PREFIX},
|
||||
linked_names=frozenset({COST_SPAN}) if require_cost_span else frozenset(),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -308,8 +288,7 @@ class TestOtelTraceCompleteness:
|
|||
The trace should have a single server root span for the incoming request, with
|
||||
the authentication and database work beneath it. The span for the actual model
|
||||
call must also belong to that same trace, rather than being exported separately
|
||||
with a missing parent, and the cost-recording work must land either in that
|
||||
trace or in its own trace linked to it.
|
||||
with a missing parent.
|
||||
|
||||
This matters because a split trace is easy to miss: all of the spans may still
|
||||
arrive, but the model call appears without the surrounding request context.
|
||||
|
|
@ -367,8 +346,6 @@ class TestOtelTraceCompleteness:
|
|||
The trace must have a single root span named "POST /v1/messages". The
|
||||
authentication, database, and model-call spans must all belong to
|
||||
the same trace and have valid parent relationships leading back to that root.
|
||||
The cost-writing span must land in the request trace or in its own trace
|
||||
linked to it.
|
||||
|
||||
The model-call span is expected to be named "chat <model>". The test fails if
|
||||
the request is split across multiple traces, if any span references a missing
|
||||
|
|
@ -397,13 +374,10 @@ class TestOtelTraceCompleteness:
|
|||
|
||||
The trace must have a single root span named "POST /v1/responses". The
|
||||
authentication, database, and model-call spans must all belong to the same
|
||||
trace and have valid parent relationships leading back to that root. The cost
|
||||
write finishes after the response, so it lands as the root of its own trace
|
||||
linked FOLLOWS_FROM to the request SERVER span.
|
||||
trace and have valid parent relationships leading back to that root.
|
||||
|
||||
The model-call span is expected to be named "chat <model>". The test fails on
|
||||
a split request trace, a dangling parent, a disconnected model-call span, or
|
||||
a cost write that is neither in the request trace nor linked to it."""
|
||||
a split request trace, a dangling parent or a disconnected model-call span."""
|
||||
route = "/v1/responses"
|
||||
_assert_otel_destination_configured(client)
|
||||
|
||||
|
|
@ -428,8 +402,7 @@ class TestOtelTraceCompleteness:
|
|||
"""A successful streamed `/chat/completions` request should export one
|
||||
complete OTEL trace. The trace must contain a single root `SERVER`
|
||||
span, with the auth, database, and gen-AI `CLIENT` spans all
|
||||
connected back to that root, and the cost write in that trace or in
|
||||
its own trace linked to it.
|
||||
connected back to that root.
|
||||
|
||||
Streaming has an additional lifecycle risk because the gen-AI span is
|
||||
closed by the stream-consumption path after the final chunk has
|
||||
|
|
@ -477,8 +450,7 @@ class TestOtelTraceCompleteness:
|
|||
"""A successful streamed `/v1/messages` request should export one
|
||||
complete OTEL trace. The trace must contain a single root `SERVER`
|
||||
span, with the auth, database, and gen-AI `CLIENT` spans all
|
||||
connected back to that root, and the cost write in that trace or in
|
||||
its own trace linked to it.
|
||||
connected back to that root.
|
||||
|
||||
This endpoint has the same streaming lifecycle risk as
|
||||
`/chat/completions`: the gen-AI span is closed by the
|
||||
|
|
@ -558,18 +530,14 @@ class TestOtelTraceCompleteness:
|
|||
)
|
||||
|
||||
genai_span = f"chat {CHEAP_OPENAI_MODEL}"
|
||||
traces = _poll(
|
||||
otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False
|
||||
)
|
||||
_assert_complete_trace(traces, route=route, genai_span=genai_span, require_cost_span=False)
|
||||
traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span)
|
||||
_assert_complete_trace(traces, route=route, genai_span=genai_span)
|
||||
|
||||
one_served_genai_span(traces.hits[0], genai_span)
|
||||
|
||||
spend_row = client.poll_proxy_spend_for_key(key)
|
||||
assert spend_row is not None and spend_row.spend is not None and spend_row.spend > 0, (
|
||||
"a successful streamed responses call must record a positive-spend row in /spend/logs "
|
||||
"(the cost-write SPAN is knowingly absent on this surface, LIT-4428, but the spend "
|
||||
f"itself must land); got {spend_row!r}"
|
||||
f"a successful streamed responses call must record a positive-spend row in /spend/logs; got {spend_row!r}"
|
||||
)
|
||||
assert spend_row.call_type == "aresponses", (
|
||||
f"the spend row must be attributed to the responses call type, got {spend_row.call_type!r}"
|
||||
|
|
@ -686,9 +654,7 @@ class TestOtelTraceCompleteness:
|
|||
)
|
||||
|
||||
genai_span = f"chat {CHEAP_OPENAI_MODEL}"
|
||||
traces = _poll(
|
||||
otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False
|
||||
)
|
||||
traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span)
|
||||
_assert_real_ttft(traces.hits, genai_span=genai_span)
|
||||
|
||||
@pytest.mark.covers("logging.otel.failure.exports_metric", exercised_on=["chat_completions"])
|
||||
|
|
@ -704,9 +670,7 @@ class TestOtelTraceCompleteness:
|
|||
|
||||
The test uses a deployment with an invalid upstream API key. This
|
||||
allows the request to pass LiteLLM’s proxy authentication and fail at
|
||||
the provider, which is necessary to generate a model-call error span.
|
||||
There should be no cost-write span because failed requests are not
|
||||
billed."""
|
||||
the provider, which is necessary to generate a model-call error span."""
|
||||
route = "/chat/completions"
|
||||
_assert_otel_destination_configured(client)
|
||||
|
||||
|
|
@ -736,10 +700,8 @@ class TestOtelTraceCompleteness:
|
|||
assert outcome.call_id is not None, "failed responses must still carry x-litellm-call-id"
|
||||
|
||||
genai_span = f"chat {model_name}"
|
||||
traces = _poll(
|
||||
otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False
|
||||
)
|
||||
_assert_complete_trace(traces, route=route, genai_span=genai_span, require_cost_span=False)
|
||||
traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span)
|
||||
_assert_complete_trace(traces, route=route, genai_span=genai_span)
|
||||
|
||||
root = next(span for span in traces.hits[0].spans if not span.references)
|
||||
assert str(_tag(root, "http.status_code")) == "401", (
|
||||
|
|
@ -760,8 +722,7 @@ class TestOtelTraceCompleteness:
|
|||
litellm.provider.error.llm_provider attribute.
|
||||
|
||||
Same setup as the chat sibling: a deployment with an invalid upstream
|
||||
API key passes proxy auth and fails at the provider with a real 401,
|
||||
and failed requests are not billed, so no cost-write span."""
|
||||
API key passes proxy auth and fails at the provider with a real 401."""
|
||||
route = "/v1/messages"
|
||||
_assert_otel_destination_configured(client)
|
||||
|
||||
|
|
@ -792,10 +753,8 @@ class TestOtelTraceCompleteness:
|
|||
assert outcome.call_id is not None, "failed responses must still carry x-litellm-call-id"
|
||||
|
||||
genai_span = f"chat {model_name}"
|
||||
traces = _poll(
|
||||
otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span, require_cost_span=False
|
||||
)
|
||||
_assert_complete_trace(traces, route=route, genai_span=genai_span, require_cost_span=False)
|
||||
traces = _poll(otel_reader, call_id=outcome.call_id, route=route, genai_span=genai_span)
|
||||
_assert_complete_trace(traces, route=route, genai_span=genai_span)
|
||||
|
||||
root = next(span for span in traces.hits[0].spans if not span.references)
|
||||
assert str(_tag(root, "http.status_code")) == "401", (
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",)
|
||||
|
|
|
|||
18
tests/unit/proxy/db/fake_prisma_engine.py
Normal file
18
tests/unit/proxy/db/fake_prisma_engine.py
Normal file
|
|
@ -0,0 +1,18 @@
|
|||
"""An ``AsyncMock`` standing in for a ``prisma_client.db`` method that reached the engine,
|
||||
marking the DB I/O witness the way ``_TrackedPrismaEngine`` does, so the producer under test
|
||||
emits its service event."""
|
||||
|
||||
from typing import TypeVar
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from litellm.proxy.db.log_db_metrics import record_db_io
|
||||
|
||||
_T = TypeVar("_T")
|
||||
|
||||
|
||||
def engine_call(return_value: _T | None = None) -> AsyncMock:
|
||||
async def run(*args: object, **kwargs: object) -> _T | None:
|
||||
record_db_io()
|
||||
return return_value
|
||||
|
||||
return AsyncMock(side_effect=run)
|
||||
|
|
@ -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,)
|
||||
|
|
|
|||
|
|
@ -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",)
|
||||
|
|
|
|||
124
tests/unit/proxy/db/test_db_span.py
Normal file
124
tests/unit/proxy/db/test_db_span.py
Normal file
|
|
@ -0,0 +1,124 @@
|
|||
import asyncio
|
||||
from collections.abc import Iterator
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from prisma.errors import PrismaError
|
||||
|
||||
from litellm._service_logger import ServiceTypes
|
||||
from litellm.proxy.db.db_span import db_span
|
||||
from litellm.proxy.db.log_db_metrics import record_db_io
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def service_hooks() -> Iterator[tuple[AsyncMock, AsyncMock]]:
|
||||
success: Final = AsyncMock()
|
||||
failure: Final = AsyncMock()
|
||||
service_logging: Final = MagicMock(async_service_success_hook=success, async_service_failure_hook=failure)
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=service_logging)):
|
||||
yield success, failure
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_completed_write_emits_one_db_event_named_for_the_call_and_table(
|
||||
service_hooks: tuple[AsyncMock, AsyncMock],
|
||||
) -> None:
|
||||
success, failure = service_hooks
|
||||
|
||||
async with db_span("commit_spend_updates", "LiteLLM_UserTable"):
|
||||
record_db_io()
|
||||
await asyncio.sleep(0)
|
||||
|
||||
event: Final = success.await_args.kwargs
|
||||
assert (event["service"], event["call_type"], event["event_metadata"]) == (
|
||||
ServiceTypes.DB,
|
||||
"commit_spend_updates",
|
||||
{"table_name": "LiteLLM_UserTable"},
|
||||
)
|
||||
assert event["duration"] == pytest.approx((event["end_time"] - event["start_time"]).total_seconds())
|
||||
assert failure.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_prisma_error_inside_the_write_emits_a_db_failure_event_and_propagates(
|
||||
service_hooks: tuple[AsyncMock, AsyncMock],
|
||||
) -> None:
|
||||
success, failure = service_hooks
|
||||
|
||||
with pytest.raises(PrismaError):
|
||||
async with db_span("insert_spend_logs", "LiteLLM_SpendLogs"):
|
||||
raise PrismaError("connection reset")
|
||||
await asyncio.sleep(0)
|
||||
|
||||
event: Final = failure.await_args.kwargs
|
||||
assert (event["service"], event["call_type"], event["event_metadata"], str(event["error"])) == (
|
||||
ServiceTypes.DB,
|
||||
"insert_spend_logs",
|
||||
{"table_name": "LiteLLM_SpendLogs"},
|
||||
"connection reset",
|
||||
)
|
||||
assert success.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_dropped_query_engine_connection_emits_a_db_failure_event(
|
||||
service_hooks: tuple[AsyncMock, AsyncMock],
|
||||
) -> None:
|
||||
success, failure = service_hooks
|
||||
|
||||
with pytest.raises(httpx.ReadError):
|
||||
async with db_span("write_tool_spend", "LiteLLM_DailyToolSpend"):
|
||||
raise httpx.ReadError("peer closed connection")
|
||||
await asyncio.sleep(0)
|
||||
|
||||
event: Final = failure.await_args.kwargs
|
||||
assert (event["call_type"], event["event_metadata"], str(event["error"])) == (
|
||||
"write_tool_spend",
|
||||
{"table_name": "LiteLLM_DailyToolSpend"},
|
||||
"peer closed connection",
|
||||
)
|
||||
assert success.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_non_database_error_inside_the_write_emits_no_db_event(
|
||||
service_hooks: tuple[AsyncMock, AsyncMock],
|
||||
) -> None:
|
||||
success, failure = service_hooks
|
||||
|
||||
with pytest.raises(ValueError, match="bad row"):
|
||||
async with db_span("insert_spend_logs", "LiteLLM_SpendLogs"):
|
||||
raise ValueError("bad row")
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert (success.await_count, failure.await_count) == (0, 0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_raising_failure_hook_never_replaces_the_prisma_error(
|
||||
service_hooks: tuple[AsyncMock, AsyncMock],
|
||||
) -> None:
|
||||
success, failure = service_hooks
|
||||
failure.side_effect = RuntimeError("exporter down")
|
||||
|
||||
with pytest.raises(PrismaError):
|
||||
async with db_span("commit_spend_updates", "LiteLLM_UserTable"):
|
||||
raise PrismaError("connection reset")
|
||||
|
||||
assert failure.await_count == 1
|
||||
assert success.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_block_whose_prisma_client_never_reached_the_engine_emits_no_db_event(
|
||||
service_hooks: tuple[AsyncMock, AsyncMock],
|
||||
) -> None:
|
||||
success, failure = service_hooks
|
||||
|
||||
async with db_span("team_user_spend", "LiteLLM_SpendLogs"):
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert (success.await_count, failure.await_count) == (0, 0)
|
||||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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",)
|
||||
|
|
|
|||
|
|
@ -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",)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
471
tests/unit/proxy/db/test_prisma_query_span.py
Normal file
471
tests/unit/proxy/db/test_prisma_query_span.py
Normal file
|
|
@ -0,0 +1,471 @@
|
|||
import ast
|
||||
import re
|
||||
from collections.abc import Iterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.otel.model.payloads import ServiceSpanData
|
||||
from litellm.integrations.otel.model import spans as spans_mod
|
||||
from litellm.integrations.otel.model.spans import (
|
||||
_POSTGRES_OPERATION_BY_CALL_TYPE,
|
||||
PRISMA_RELATIONS,
|
||||
service_span_name,
|
||||
)
|
||||
from litellm.proxy.db.prisma_query_span import UNKNOWN_PRISMA_QUERY, parse_prisma_query, sql_operation
|
||||
|
||||
_REPO: Final = Path(__file__).resolve().parents[4]
|
||||
_SOURCE_ROOTS: Final = ("litellm", "enterprise", "litellm-proxy-extras")
|
||||
_RAW_METHODS: Final = frozenset({"query_first", "query_raw", "execute_raw"})
|
||||
_MODEL_METHODS: Final = frozenset(
|
||||
{
|
||||
"find_unique",
|
||||
"find_unique_or_raise",
|
||||
"find_first",
|
||||
"find_first_or_raise",
|
||||
"find_many",
|
||||
"count",
|
||||
"group_by",
|
||||
"create",
|
||||
"create_many",
|
||||
"update",
|
||||
"update_many",
|
||||
"delete",
|
||||
"delete_many",
|
||||
"upsert",
|
||||
}
|
||||
)
|
||||
_MODEL_BY_ACCESSOR: Final[Mapping[str, str]] = {relation.lower(): relation for relation in PRISMA_RELATIONS}
|
||||
_GENERIC_CRUD_HELPERS: Final = frozenset({"get_data", "get_generic_data", "insert_data", "update_data", "delete_data"})
|
||||
_TRANSACTION_BODIES: Final[Mapping[str, str]] = {"litellm/proxy/db/baseline_accounting.py": "baseline_accounting"}
|
||||
_RENDERED_NAME: Final = re.compile(
|
||||
r"postgres\.(select|insert|update|delete|upsert|ddl|set|transaction) .+|postgres\.ping"
|
||||
)
|
||||
|
||||
|
||||
def _engine_payload(root_field: str, sql: str | None = None) -> str:
|
||||
selection: Final = f'queryRaw(query: "{sql}", parameters: "[]")' if sql is not None else root_field
|
||||
return (
|
||||
f'{{"query": "mutation {{ result: {selection} }}"}}'
|
||||
if sql is not None
|
||||
else f'{{"query": "query {{ result: {root_field}(where: {{token: \\"x\\"}}) {{ token }} }}"}}'
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("content", "expected"),
|
||||
[
|
||||
(
|
||||
_engine_payload("findUniqueLiteLLM_VerificationToken"),
|
||||
("find_unique", "select", "LiteLLM_VerificationToken"),
|
||||
),
|
||||
(_engine_payload("createOneLiteLLM_SpendLogs"), ("create", "insert", "LiteLLM_SpendLogs")),
|
||||
(_engine_payload("findFirstLiteLLM_UserTableOrThrow"), ("find_first", "select", "LiteLLM_UserTable")),
|
||||
(
|
||||
'{"query": "mutation { result: queryRaw(query: \\"SELECT * FROM \\\\\\"LiteLLM_UserTable\\\\\\" WHERE user_id = $1\\", parameters: \\"[]\\") }"}',
|
||||
("query_raw", "select", "LiteLLM_UserTable"),
|
||||
),
|
||||
(
|
||||
'{"query": "mutation { result: executeRaw(query: \\"SET LOCAL statement_timeout = 5000\\", parameters: \\"[]\\") }"}',
|
||||
("execute_raw", "set", "statement_timeout"),
|
||||
),
|
||||
(
|
||||
'{"query": "mutation { result: queryRaw(query: \\"SELECT to_regclass($1) IS NOT NULL AS present\\", parameters: \\"[]\\") }"}',
|
||||
("query_raw", "select", "pg_catalog"),
|
||||
),
|
||||
(
|
||||
'{"query": "mutation { result: queryRaw(query: \\"SELECT 1\\", parameters: \\"[]\\") }"}',
|
||||
("query_raw", "ping", None),
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_the_engine_names_a_round_trip_from_its_payload_without_copying_sql_text(
|
||||
content: str, expected: tuple[str, str | None, str | None]
|
||||
) -> None:
|
||||
query = parse_prisma_query(content)
|
||||
assert (query.call_type, query.operation, query.table) == expected
|
||||
assert query.table is None or " " not in query.table
|
||||
|
||||
|
||||
def test_a_payload_the_parser_does_not_know_stays_the_legacy_function_named_span() -> None:
|
||||
assert parse_prisma_query("not json at all") is UNKNOWN_PRISMA_QUERY
|
||||
assert parse_prisma_query('{"query": "mutation { result: somethingNew(x: 1) }"}') is UNKNOWN_PRISMA_QUERY
|
||||
rendered = service_span_name(ServiceSpanData(service_name="postgres", call_type=UNKNOWN_PRISMA_QUERY.call_type))
|
||||
assert rendered == "postgres prisma_query"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("sql", "expected"),
|
||||
[
|
||||
('SELECT 1 FROM "LiteLLM_VerificationTokenView" LIMIT 1', ("select", "LiteLLM_VerificationTokenView")),
|
||||
(
|
||||
'\n WITH keys AS (SELECT * FROM "LiteLLM_VerificationToken") SELECT 1',
|
||||
("select", "LiteLLM_VerificationToken"),
|
||||
),
|
||||
('INSERT INTO "LiteLLM_DailyUserSpend" (id) VALUES ($1)', ("insert", "LiteLLM_DailyUserSpend")),
|
||||
("SET LOCAL lock_timeout = 1000", ("set", "lock_timeout")),
|
||||
("SELECT COUNT(*) FROM pg_stat_activity", ("select", "pg_catalog")),
|
||||
("SELECT 1", ("ping", None)),
|
||||
("SELECT current_setting('transaction_read_only') AS transaction_read_only", ("select", "pg_catalog")),
|
||||
('REFRESH MATERIALIZED VIEW "MonthlyGlobalSpend"', ("ddl", "MonthlyGlobalSpend")),
|
||||
(
|
||||
'WITH team_rows AS (UPDATE "LiteLLM_TeamTable" SET models = $1 RETURNING team_id) SELECT team_id FROM team_rows',
|
||||
("update", "LiteLLM_TeamTable"),
|
||||
),
|
||||
("BEGIN", (None, None)),
|
||||
],
|
||||
)
|
||||
def test_sql_operation_is_the_leading_verb_and_the_first_schema_relation(
|
||||
sql: str, expected: tuple[str | None, str | None]
|
||||
) -> None:
|
||||
assert sql_operation(sql) == expected
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PrismaCallSite:
|
||||
location: str
|
||||
method: str
|
||||
owner: str
|
||||
rendered: str | None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Module:
|
||||
path: Path
|
||||
tree: ast.Module
|
||||
constants: Mapping[str, ast.expr]
|
||||
|
||||
def ancestors(self, node: ast.AST) -> tuple[ast.AST, ...]:
|
||||
parent_of: Final = _parent_map(self.tree)
|
||||
chain: Final = [node]
|
||||
while (parent := parent_of.get(id(chain[-1]))) is not None:
|
||||
chain.append(parent)
|
||||
return tuple(chain[1:])
|
||||
|
||||
|
||||
_PARENTS: Final[dict[int, Mapping[int, ast.AST]]] = {} # mutable-ok: per-tree parent map memo
|
||||
|
||||
|
||||
def _parent_map(tree: ast.Module) -> Mapping[int, ast.AST]:
|
||||
if id(tree) not in _PARENTS:
|
||||
_PARENTS[id(tree)] = {
|
||||
id(child): node for node in ast.walk(tree) for child in ast.iter_child_nodes(node)
|
||||
} # comprehension-ok: parent links
|
||||
return _PARENTS[id(tree)]
|
||||
|
||||
|
||||
def _modules() -> Iterator[_Module]:
|
||||
for root in _SOURCE_ROOTS:
|
||||
for path in sorted((_REPO / root).rglob("*.py")):
|
||||
if "tests" in path.parts or "node_modules" in path.parts:
|
||||
continue
|
||||
tree: Final = ast.parse(path.read_text(encoding="utf-8"))
|
||||
yield _Module(path, tree, _assignments(tree.body))
|
||||
|
||||
|
||||
def _imported_module(module: _Module, name: str) -> Path | None:
|
||||
for node in module.tree.body:
|
||||
if isinstance(node, ast.ImportFrom) and node.module and any(alias.name == name for alias in node.names):
|
||||
return _REPO / (node.module.replace(".", "/") + ".py")
|
||||
return None
|
||||
|
||||
|
||||
def _assignments(body: list[ast.stmt]) -> Mapping[str, ast.expr]:
|
||||
return {
|
||||
target.id: node.value
|
||||
for node in ast.walk(ast.Module(body=body, type_ignores=[]))
|
||||
if isinstance(node, (ast.Assign, ast.AnnAssign)) and node.value is not None
|
||||
for target in (node.targets if isinstance(node, ast.Assign) else (node.target,))
|
||||
if isinstance(target, ast.Name)
|
||||
} # comprehension-ok: constants by name
|
||||
|
||||
|
||||
def _mapping_values(expr: ast.expr | None) -> ast.expr | None:
|
||||
"""The dict a ``Mapping`` constant was built from, through ``MappingProxyType(...)``."""
|
||||
if (
|
||||
isinstance(expr, ast.Call)
|
||||
and isinstance(expr.func, ast.Name)
|
||||
and expr.func.id == "MappingProxyType"
|
||||
and expr.args
|
||||
):
|
||||
return expr.args[0]
|
||||
return expr if isinstance(expr, (ast.Dict, ast.DictComp)) else None
|
||||
|
||||
|
||||
def _returned_text(function_name: str, module: _Module) -> ast.expr | None:
|
||||
"""What a module-level SQL builder returns, when its body is one ``return`` of a string expression."""
|
||||
for node in module.tree.body:
|
||||
if isinstance(node, ast.FunctionDef) and node.name == function_name:
|
||||
returns: Final = [stmt for stmt in ast.walk(node) if isinstance(stmt, ast.Return)]
|
||||
return returns[0].value if len(returns) == 1 else None
|
||||
return None
|
||||
|
||||
|
||||
_DYNAMIC: Final = " ? "
|
||||
|
||||
|
||||
def _fragment(value: ast.expr, module: _Module, depth: int) -> str:
|
||||
"""One f-string piece: literal text, a module constant spliced in, or a runtime placeholder."""
|
||||
spliced: Final = (
|
||||
_sql_text(value.value, module, depth + 1)
|
||||
if isinstance(value, ast.FormattedValue)
|
||||
else _sql_text(value, module, depth)
|
||||
)
|
||||
return _DYNAMIC if spliced is None or _ALTERNATIVE in spliced else spliced
|
||||
|
||||
|
||||
def _sql_text(expr: ast.expr | None, module: _Module, depth: int = 0) -> str | None:
|
||||
if expr is None or depth > 3:
|
||||
return None
|
||||
if isinstance(expr, ast.Constant) and isinstance(expr.value, str):
|
||||
return expr.value
|
||||
if isinstance(expr, ast.JoinedStr):
|
||||
return "".join(_fragment(value, module, depth) for value in expr.values)
|
||||
if isinstance(expr, ast.BinOp) and isinstance(expr.op, ast.Add):
|
||||
left: Final = _sql_text(expr.left, module, depth)
|
||||
return left if left is not None else _sql_text(expr.right, module, depth)
|
||||
if (
|
||||
isinstance(expr, ast.Call)
|
||||
and isinstance(expr.func, ast.Attribute)
|
||||
and expr.func.attr in {"format", "strip", "lstrip"}
|
||||
):
|
||||
return _sql_text(expr.func.value, module, depth)
|
||||
if isinstance(expr, ast.Call) and isinstance(expr.func, ast.Attribute) and expr.func.attr == "dedent":
|
||||
return _sql_text(expr.args[0], module, depth) if expr.args else None
|
||||
if isinstance(expr, ast.IfExp):
|
||||
branches: Final = (_sql_text(expr.body, module, depth), _sql_text(expr.orelse, module, depth))
|
||||
return branches[0] if branches[0] == branches[1] or None in branches else _multi(branches)
|
||||
if isinstance(expr, ast.Subscript) and isinstance(expr.value, ast.Name):
|
||||
return _sql_text(_mapping_values(module.constants.get(expr.value.id)), module, depth + 1)
|
||||
if (
|
||||
isinstance(expr, ast.Call)
|
||||
and isinstance(expr.func, ast.Name)
|
||||
and expr.func.id == "MappingProxyType"
|
||||
and expr.args
|
||||
):
|
||||
return _sql_text(expr.args[0], module, depth)
|
||||
if isinstance(expr, ast.Dict):
|
||||
values: Final = tuple(_sql_text(value, module, depth) for value in expr.values)
|
||||
return _multi(values) if values and None not in values else None
|
||||
if isinstance(expr, ast.DictComp):
|
||||
return _sql_text(expr.value, module, depth)
|
||||
if isinstance(expr, ast.Call) and isinstance(expr.func, ast.Name):
|
||||
returned: Final = _returned_text(expr.func.id, module)
|
||||
return _sql_text(returned, module, depth + 1) if returned is not None else None
|
||||
if isinstance(expr, ast.Name):
|
||||
if expr.id in module.constants:
|
||||
return _sql_text(module.constants[expr.id], module, depth + 1)
|
||||
source: Final = _imported_module(module, expr.id)
|
||||
if source is None or not source.exists():
|
||||
return None
|
||||
imported: Final = ast.parse(source.read_text(encoding="utf-8"))
|
||||
imported_module: Final = _Module(source, imported, _assignments(imported.body))
|
||||
return _sql_text(imported_module.constants.get(expr.id), imported_module, depth + 1)
|
||||
return None
|
||||
|
||||
|
||||
_ALTERNATIVE: Final = "\x1f"
|
||||
|
||||
|
||||
def _multi(texts: tuple[str | None, ...]) -> str:
|
||||
return _ALTERNATIVE.join(text for text in texts if text is not None)
|
||||
|
||||
|
||||
def _parameters(parents: tuple[ast.AST, ...]) -> frozenset[str]:
|
||||
function: Final = next((p for p in parents if isinstance(p, (ast.AsyncFunctionDef, ast.FunctionDef))), None)
|
||||
if function is None:
|
||||
return frozenset()
|
||||
return frozenset(arg.arg for arg in (*function.args.args, *function.args.kwonlyargs))
|
||||
|
||||
|
||||
def _argument_for(call: ast.Call, function: ast.AsyncFunctionDef | ast.FunctionDef, parameter: str) -> ast.expr | None:
|
||||
positional: Final = tuple(arg.arg for arg in function.args.args)
|
||||
by_keyword: Final = next((k.value for k in call.keywords if k.arg == parameter), None)
|
||||
if by_keyword is not None or parameter not in positional:
|
||||
return by_keyword
|
||||
index: Final = positional.index(parameter)
|
||||
return call.args[index] if index < len(call.args) else None
|
||||
|
||||
|
||||
def _parameter_site(
|
||||
parameter: str, module: _Module, parents: tuple[ast.AST, ...], method: str, location: str
|
||||
) -> _PrismaCallSite:
|
||||
"""A statement that arrives as a parameter: a ``query_raw`` forwarder adds no round trip of its
|
||||
own, any other helper is named by what its callers in the module hand it."""
|
||||
function: Final = next(p for p in parents if isinstance(p, (ast.AsyncFunctionDef, ast.FunctionDef)))
|
||||
if function.name in _RAW_METHODS:
|
||||
return _PrismaCallSite(location, method, "forwarder", f"(callers of {function.name})")
|
||||
callers: Final = tuple(
|
||||
node
|
||||
for node in ast.walk(module.tree)
|
||||
if isinstance(node, ast.Call) and ast.unparse(node.func).endswith(function.name)
|
||||
)
|
||||
sites: Final = tuple(_caller_site(call, function, parameter, module, method, location) for call in callers)
|
||||
names: Final = tuple(site.rendered for site in sites)
|
||||
owners: Final = ", ".join(sorted({site.owner for site in sites}))
|
||||
rendered: Final = " | ".join(sorted(set(names))) if names and None not in names else None # pyright: ignore[reportArgumentType] # None filtered above
|
||||
return _PrismaCallSite(location, method, f"{owners} via {function.name} callers", rendered)
|
||||
|
||||
|
||||
def _caller_site(
|
||||
call: ast.Call,
|
||||
function: ast.AsyncFunctionDef | ast.FunctionDef,
|
||||
parameter: str,
|
||||
module: _Module,
|
||||
method: str,
|
||||
location: str,
|
||||
) -> _PrismaCallSite:
|
||||
parents: Final = module.ancestors(call)
|
||||
wrapped: Final = _wrapper_site(call, parents, method, location, module)
|
||||
if wrapped is not None:
|
||||
return wrapped
|
||||
text: Final = _sql_text(_argument_for(call, function, parameter), _scope(module, parents))
|
||||
return _PrismaCallSite(location, method, "engine", _render_statements(method, text) if text is not None else None)
|
||||
|
||||
|
||||
def _render_statements(method: str, text: str) -> str | None:
|
||||
names: Final = tuple(_render_statement(method, alternative) for alternative in text.split(_ALTERNATIVE))
|
||||
return " | ".join(sorted(set(names))) if None not in names else None # pyright: ignore[reportArgumentType] # None filtered above
|
||||
|
||||
|
||||
def _render_statement(method: str, text: str) -> str | None:
|
||||
verb, target = sql_operation(text)
|
||||
if verb is not None and target is None and verb != "ping" and _DYNAMIC in text:
|
||||
return f"postgres.{verb} {{relation built at runtime}}"
|
||||
return _render(method, target, verb) if verb is not None else None
|
||||
|
||||
|
||||
def _render(call_type: str, table: str | None, operation: str | None = None) -> str:
|
||||
metadata: Final = {
|
||||
key: value for key, value in (("table_name", table), ("db_operation", operation)) if value is not None
|
||||
}
|
||||
return service_span_name(ServiceSpanData(service_name="postgres", call_type=call_type, event_metadata=metadata))
|
||||
|
||||
|
||||
def _wrapper_site(
|
||||
call: ast.Call, parents: tuple[ast.AST, ...], method: str, location: str, module: _Module
|
||||
) -> _PrismaCallSite | None:
|
||||
for parent in parents:
|
||||
items: Final = parent.items if isinstance(parent, (ast.AsyncWith, ast.With)) else ()
|
||||
for item in items:
|
||||
context: Final = item.context_expr
|
||||
if isinstance(context, ast.Call) and isinstance(context.func, ast.Name) and context.func.id == "db_span":
|
||||
return _wrapped_by(context, method, location, "db_span", _scope(module, parents))
|
||||
if (
|
||||
isinstance(context, ast.Call)
|
||||
and isinstance(context.func, ast.Name)
|
||||
and context.func.id == "_spend_update_tx"
|
||||
):
|
||||
call_type: Final = context.args[2] if len(context.args) > 2 else ast.Constant("commit_spend_updates")
|
||||
spend_tx: Final = ast.Call(func=ast.Name("db_span"), args=[call_type, context.args[1]], keywords=[])
|
||||
return _wrapped_by(spend_tx, method, location, "_spend_update_tx", _scope(module, parents))
|
||||
if isinstance(parent, ast.Call) and isinstance(parent.func, ast.Name) and parent.func.id == "db_spanned":
|
||||
return _wrapped_by(parent, method, location, "db_spanned", _scope(module, parents))
|
||||
if isinstance(parent, (ast.AsyncFunctionDef, ast.FunctionDef)):
|
||||
decorators: Final = tuple(
|
||||
decorator.id for decorator in parent.decorator_list if isinstance(decorator, ast.Name)
|
||||
)
|
||||
if "log_db_metrics" in decorators:
|
||||
return _decorated_site(parent.name, method, location)
|
||||
return None
|
||||
|
||||
|
||||
def _wrapped_by(wrapper: ast.Call, method: str, location: str, owner: str, scope: _Module) -> _PrismaCallSite:
|
||||
call_type: Final = _sql_text(wrapper.args[0], scope)
|
||||
table_expr: Final = wrapper.args[1] if len(wrapper.args) > 1 else None
|
||||
if call_type is None:
|
||||
return _PrismaCallSite(location, method, owner, None)
|
||||
if isinstance(table_expr, ast.Constant) and table_expr.value is None:
|
||||
return _PrismaCallSite(location, method, owner, _render(call_type, None))
|
||||
table: Final = _sql_text(table_expr, scope)
|
||||
if table is not None:
|
||||
return _PrismaCallSite(location, method, owner, _render(call_type, table))
|
||||
operation: Final = _POSTGRES_OPERATION_BY_CALL_TYPE.get(call_type)
|
||||
rendered: Final = f"postgres.{operation.verb} {{relation}}" if operation is not None else None
|
||||
return _PrismaCallSite(location, method, f"{owner}(bounded)", rendered)
|
||||
|
||||
|
||||
def _decorated_site(function: str, method: str, location: str) -> _PrismaCallSite:
|
||||
if function in _GENERIC_CRUD_HELPERS:
|
||||
return _PrismaCallSite(location, method, "log_db_metrics(crud)", "postgres.{verb} {table_name}")
|
||||
operation: Final = _POSTGRES_OPERATION_BY_CALL_TYPE.get(function)
|
||||
return _PrismaCallSite(
|
||||
location, method, "log_db_metrics", _render(function, None) if operation is not None else None
|
||||
)
|
||||
|
||||
|
||||
def _scope(module: _Module, parents: tuple[ast.AST, ...]) -> _Module:
|
||||
function: Final = next((p for p in parents if isinstance(p, (ast.AsyncFunctionDef, ast.FunctionDef))), None)
|
||||
if function is None:
|
||||
return module
|
||||
return _Module(module.path, module.tree, {**module.constants, **_assignments(function.body)})
|
||||
|
||||
|
||||
def _engine_site(
|
||||
call: ast.Call, module: _Module, parents: tuple[ast.AST, ...], method: str, accessor: str | None, location: str
|
||||
) -> _PrismaCallSite:
|
||||
if accessor is not None:
|
||||
return _PrismaCallSite(location, method, "engine", _render(method, _MODEL_BY_ACCESSOR.get(accessor)))
|
||||
scope: Final = _scope(module, parents)
|
||||
statement: Final = call.args[0] if call.args else next((k.value for k in call.keywords if k.arg == "query"), None)
|
||||
if isinstance(statement, ast.Name) and statement.id in _parameters(parents):
|
||||
return _parameter_site(statement.id, module, parents, method, location)
|
||||
rendered: Final = _render_statements(method, _sql_text(statement, scope) or "")
|
||||
relative: Final = str(module.path.relative_to(_REPO))
|
||||
if _RENDERED_NAME.fullmatch(rendered or "") is None and relative in _TRANSACTION_BODIES:
|
||||
owner: Final = _TRANSACTION_BODIES[relative]
|
||||
return _PrismaCallSite(location, method, f"transaction({owner})", _render(owner, None))
|
||||
return _PrismaCallSite(location, method, "engine", rendered)
|
||||
|
||||
|
||||
def _accessor(receiver: ast.expr) -> str | None:
|
||||
if isinstance(receiver, ast.Attribute) and receiver.attr in _MODEL_BY_ACCESSOR:
|
||||
return receiver.attr
|
||||
return None
|
||||
|
||||
|
||||
def _call_sites(module: _Module) -> Iterator[_PrismaCallSite]:
|
||||
def walk(node: ast.AST, parents: tuple[ast.AST, ...]) -> Iterator[_PrismaCallSite]:
|
||||
for child in ast.iter_child_nodes(node):
|
||||
if isinstance(child, ast.Call) and isinstance(child.func, ast.Attribute):
|
||||
method: Final = child.func.attr
|
||||
accessor: Final = _accessor(child.func.value)
|
||||
if method in _RAW_METHODS or (method in _MODEL_METHODS and accessor is not None):
|
||||
location: Final = f"{module.path.relative_to(_REPO)}:{child.lineno}"
|
||||
yield _wrapper_site(child, parents, method, location, module) or _engine_site(
|
||||
child, module, parents, method, accessor, location
|
||||
)
|
||||
yield from walk(child, (child, *parents))
|
||||
|
||||
yield from walk(module.tree, ())
|
||||
|
||||
|
||||
def prisma_call_sites() -> tuple[_PrismaCallSite, ...]:
|
||||
return tuple(site for module in _modules() for site in _call_sites(module)) # comprehension-ok: flatten
|
||||
|
||||
|
||||
def test_every_prisma_call_site_in_the_proxy_renders_a_bounded_postgres_span_name() -> None:
|
||||
"""A raw ``query_raw``/``execute_raw``/``query_first`` or a direct model call that no producer
|
||||
wraps is named by the engine from its payload; this scan replays that naming (and the wrappers')
|
||||
statically so a new statement that would ship as a bare ``postgres.select`` or an unnamed
|
||||
``postgres query_raw`` fails here rather than in a trace."""
|
||||
sites = prisma_call_sites()
|
||||
assert len(sites) >= 120, f"the scan lost the Prisma call sites: {len(sites)}"
|
||||
unresolved = [site for site in sites if site.rendered is None]
|
||||
assert unresolved == [], f"Prisma call sites whose span name cannot be resolved: {unresolved}"
|
||||
half_named = [
|
||||
site
|
||||
for site in sites
|
||||
if site.owner.startswith("engine")
|
||||
and any(_RENDERED_NAME.fullmatch(name) is None for name in (site.rendered or "").split(" | "))
|
||||
]
|
||||
assert half_named == [], f"Prisma call sites that would ship a half-named or legacy span: {half_named}"
|
||||
|
||||
|
||||
def test_every_model_in_the_prisma_schema_is_a_renderable_span_table() -> None:
|
||||
schema: Final = (_REPO / "schema.prisma").read_text()
|
||||
declared: Final = frozenset(re.findall(r"^model (\w+) \{", schema, re.MULTILINE))
|
||||
|
||||
assert declared == spans_mod._PRISMA_MODELS
|
||||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",)
|
||||
|
|
|
|||
|
|
@ -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",)
|
||||
|
|
|
|||
|
|
@ -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})]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue