refactor(db): make read-replica routing an explicit opt-in via PrismaClient.replica_db

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yuneng 2026-09-24 08:14:48 +00:00
parent 6dbd65b230
commit f6f3920337
82 changed files with 1050 additions and 131 deletions

View file

@ -90,7 +90,7 @@ class _ENTERPRISE_BlockedUserList(CustomLogger):
if end_user_cache_obj is None and self.prisma_client is not None:
# check db
end_user_obj = (
await self.prisma_client.db.litellm_endusertable.find_unique(
await self.prisma_client.replica_db.litellm_endusertable.find_unique(
where={"user_id": user}
)
)

View file

@ -813,7 +813,7 @@ class BaseEmailLogger(CustomLogger):
)
return None
user_row = await prisma_client.db.litellm_usertable.find_unique(
user_row = await prisma_client.replica_db.litellm_usertable.find_unique(
where={"user_id": user_id}
)
@ -898,7 +898,7 @@ class BaseEmailLogger(CustomLogger):
try:
# Try to get existing invitation
existing_invitations = (
await prisma_client.db.litellm_invitationlink.find_many(
await prisma_client.replica_db.litellm_invitationlink.find_many(
where={"user_id": user_id},
order={"created_at": "desc"},
)

View file

@ -25,7 +25,7 @@ async def _get_email_settings(prisma_client) -> Dict[str, bool]:
"""Helper function to get email settings from general_settings in db"""
try:
# Get general settings from db
general_settings_entry = await prisma_client.db.litellm_config.find_unique(
general_settings_entry = await prisma_client.replica_db.litellm_config.find_unique(
where={"param_name": "general_settings"}
)
@ -71,7 +71,7 @@ async def _save_email_settings(prisma_client, settings: Dict[str, bool]):
)
# Get current general settings
general_settings_entry = await prisma_client.db.litellm_config.find_unique(
general_settings_entry = await prisma_client.replica_db.litellm_config.find_unique(
where={"param_name": "general_settings"}
)

View file

@ -57,24 +57,24 @@ class _ManagedObjectRow(Protocol):
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.replica_db.litellm_managedobjecttable
return table
def _user_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_UserTable]":
table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.db.litellm_usertable
table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.replica_db.litellm_usertable
return table
def _token_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_VerificationToken]":
table: Final[TableActions[prisma_models.LiteLLM_VerificationToken]] = (
prisma_client.db.litellm_verificationtoken
prisma_client.replica_db.litellm_verificationtoken
)
return table
def _team_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_TeamTable]":
table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.db.litellm_teamtable
table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.replica_db.litellm_teamtable
return table

View file

@ -43,7 +43,7 @@ class _ManagedObjectRow(Protocol):
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.replica_db.litellm_managedobjecttable
return table

View file

@ -197,11 +197,11 @@ class _CursorPageArgs(TypedDict, total=False):
def _managed_file_table(prisma_client: PrismaClient) -> _ManagedFileTableActions:
return prisma_client.db.litellm_managedfiletable
return prisma_client.replica_db.litellm_managedfiletable
def _managed_object_table(prisma_client: PrismaClient) -> _ManagedObjectTableActions:
return prisma_client.db.litellm_managedobjecttable
return prisma_client.replica_db.litellm_managedobjecttable
def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, str]:

View file

@ -59,7 +59,7 @@ class ManagedVectorStoreTable(Protocol):
def managed_vector_store_table(prisma_client: "PrismaClient") -> ManagedVectorStoreTable:
"""The Prisma table actions for managed vector stores, behind a typed surface."""
return prisma_client.db.litellm_managedvectorstorestable
return prisma_client.replica_db.litellm_managedvectorstorestable
########################################################

View file

@ -56,7 +56,9 @@ async def fetch_user_spend_rows(
today_str: Final = today.strftime("%Y-%m-%d")
month_start_str: Final = today.replace(day=1).strftime("%Y-%m-%d")
baseline_start_str: Final = (today - datetime.timedelta(days=max(baseline_days, 1))).strftime("%Y-%m-%d")
raw: Final = await prisma_client.db.query_raw(USER_SPEND_QUERY, today_str, month_start_str, baseline_start_str)
raw: Final = await prisma_client.replica_db.query_raw(
USER_SPEND_QUERY, today_str, month_start_str, baseline_start_str
)
return USER_SPEND_ROWS_ADAPTER.validate_python(raw)

View file

@ -93,7 +93,7 @@ class LiteLLMDatabase:
query += " LIMIT $3"
try:
db_response: Final = await client.db.query_raw(query, *params)
db_response: Final = await client.replica_db.query_raw(query, *params)
from litellm.proxy.spend_tracking.key_metadata_recovery import (
fill_missing_api_key_aliases,
)

View file

@ -51,7 +51,7 @@ async def get_all_team_member_emails(team_id: str | None = None) -> list:
WHERE user_id = ANY($1::TEXT[]);
"""
_result: Final = await prisma_client.db.query_raw(sql_query, _team_member_user_ids)
_result: Final = await prisma_client.replica_db.query_raw(sql_query, _team_member_user_ids)
verbose_logger.debug("Email Alerting: Got all Emails for team, emails=%s", _result)

View file

@ -95,7 +95,7 @@ class FocusLiteLLMDatabase:
"""
try:
db_response: Final = await client.db.query_raw(query, *query_params)
db_response: Final = await client.replica_db.query_raw(query, *query_params)
from litellm.proxy.spend_tracking.key_metadata_recovery import (
fill_missing_api_key_aliases,
)
@ -123,7 +123,7 @@ class FocusLiteLLMDatabase:
ORDER BY ordinal_position;
"""
try:
columns_response: Final = await client.db.query_raw(info_query)
columns_response: Final = await client.replica_db.query_raw(info_query)
return {"columns": columns_response, "table_name": "LiteLLM_DailyUserSpend"}
except Exception as exc:
raise RuntimeError(f"Error getting table info: {exc}") from exc

View file

@ -863,14 +863,14 @@ class ShadowEvalLogger(CustomLogger):
if prisma is None:
return _EMPTY_JOBS
try:
records: Final = await prisma.db.litellm_shadowevaljob.find_many(
records: Final = await prisma.replica_db.litellm_shadowevaljob.find_many(
where={ # mutable-ok: Prisma filter
"stopped_at": None,
"ends_at": {"gt": datetime.now(timezone.utc)}, # mutable-ok: Prisma filter
},
)
grouped: Final = (
await prisma.db.litellm_shadowevalattempt.group_by(
await prisma.replica_db.litellm_shadowevalattempt.group_by(
by=["job_id"],
count=True,
sum={"judge_cost": True, "shadow_cost": True, "shadow_classifier_cost": True},

View file

@ -90,7 +90,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
self.prisma_client = prisma_client
def _resource_table(self) -> _ManagedResourceTable[ResourceObjectType]:
return getattr(self.prisma_client.db, self.table_name)
return getattr(self.prisma_client.replica_db, self.table_name)
# ============================================================================
# ABSTRACT METHODS

View file

@ -99,12 +99,12 @@ class _MCPUserCredentialsTable(Protocol):
def _mcp_server_table(prisma_client: PrismaClient) -> _MCPServerTable:
"""The MCP server table, typed so the untyped prisma client surface stops here."""
return prisma_client.db.litellm_mcpservertable
return prisma_client.replica_db.litellm_mcpservertable
def _mcp_user_credentials_table(prisma_client: PrismaClient) -> _MCPUserCredentialsTable:
"""The per-user MCP credential table, typed so the untyped prisma client surface stops here."""
return prisma_client.db.litellm_mcpusercredentials
return prisma_client.replica_db.litellm_mcpusercredentials
def _decrypted_credentials(raw_credentials: str | Mapping[str, JsonValue] | None) -> MCPCredentials | None:

View file

@ -105,7 +105,7 @@ def _is_stamped_issuer_row(row: _MCPServerRow) -> bool:
async def backfill_discovery_stamped_issuers(prisma_client: PrismaClient) -> int:
"""Clear gateway-written issuer stamps, returning the number of rows healed."""
candidate_rows: Final[list[_MCPServerRow]] = await prisma_client.db.litellm_mcpservertable.find_many(
candidate_rows: Final[list[_MCPServerRow]] = await prisma_client.replica_db.litellm_mcpservertable.find_many(
where={
"updated_by": _DISCOVERY_ACTOR,
"auth_type": {"in": list(_AUTH_TYPES_WITH_ISSUER_ANCHORING)},

View file

@ -59,12 +59,12 @@ class _MCPServerTable(Protocol):
def _assertion_table(prisma_client: PrismaClient) -> _SSOAssertionTable:
"""The SSO assertion table, typed so the untyped prisma client surface stops here."""
return prisma_client.db.litellm_ssoidentityassertion
return prisma_client.replica_db.litellm_ssoidentityassertion
def _mcp_server_table(prisma_client: PrismaClient) -> _MCPServerTable:
"""The MCP server table, typed so the untyped prisma client surface stops here."""
return prisma_client.db.litellm_mcpservertable
return prisma_client.replica_db.litellm_mcpservertable
class SSOIdentityAssertion(BaseModel):

View file

@ -121,7 +121,7 @@ async def _attach_keys_to_agents(agents: Sequence[AgentResponse], prisma_client)
)
if not agent_ids:
return
key_rows: Final = await prisma_client.db.litellm_verificationtoken.find_many(
key_rows: Final = await prisma_client.replica_db.litellm_verificationtoken.find_many(
where={"agent_id": {"in": agent_ids}},
)
keys_by_agent: Final[dict[str, list[AgentKeySummary]]] = {}

View file

@ -22,7 +22,7 @@ class _SupportsRawQueryDb(Protocol):
"""A prisma client handle, narrowed to the raw-query surface used here."""
@property
def db(self) -> _SupportsQueryRaw: ...
def replica_db(self) -> _SupportsQueryRaw: ...
class CacheActivityGroup(BaseModel):
@ -169,12 +169,14 @@ async def get_cache_activity(
key_aliases_json: Final = json.dumps(list(key_aliases))
models_json: Final = json.dumps(list(models))
group_rows, error_rows, key_alias_rows, model_rows = await asyncio.gather(
prisma_client.db.query_raw(GROUPS_SQL, start_date, end_date, key_aliases_json, models_json, INFO_ROUTES_JSON),
prisma_client.db.query_raw(
prisma_client.replica_db.query_raw(
GROUPS_SQL, start_date, end_date, key_aliases_json, models_json, INFO_ROUTES_JSON
),
prisma_client.replica_db.query_raw(
ERROR_BREAKDOWN_SQL, start_date, end_date, key_aliases_json, models_json, INFO_ROUTES_JSON
),
prisma_client.db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date, INFO_ROUTES_JSON),
prisma_client.db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date, INFO_ROUTES_JSON),
prisma_client.replica_db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date, INFO_ROUTES_JSON),
prisma_client.replica_db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date, INFO_ROUTES_JSON),
)
groups: Final = _groups_adapter.validate_python(group_rows or [])
return CacheActivityResponse(

View file

@ -248,7 +248,7 @@ 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
row: Final[object] = await prisma_client.replica_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,

View file

@ -735,7 +735,7 @@ async def _fetch_global_spend_with_event_coordination(
"""
async def _load_global_spend() -> float | None:
proxy_budget_row: Final = await prisma_client.db.litellm_usertable.find_unique(
proxy_budget_row: Final = await prisma_client.replica_db.litellm_usertable.find_unique(
where={"user_id": LITELLM_PROXY_BUDGET_NAME}
)
return float(proxy_budget_row.spend) if proxy_budget_row is not None else None

View file

@ -1524,7 +1524,7 @@ 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: self.prisma_client.replica_db.query_raw(source.page_query(), cursor, RESET_BUDGET_JOB_BATCH_SIZE),
reason=f"reset_budget_read_{source.retry_subject}_windows_failure",
)
for row in rows:

View file

@ -147,9 +147,9 @@ async def spend_logs_seed_totals(
else:
return None
rows: Final = (
await prisma_client.db.query_raw(unbounded_sql, entity_id, window_start)
await prisma_client.replica_db.query_raw(unbounded_sql, entity_id, window_start)
if batch_started_at is None
else await prisma_client.db.query_raw(
else await prisma_client.replica_db.query_raw(
bounded_sql,
entity_id,
window_start,
@ -182,7 +182,7 @@ async def _existing_primary_keys(
prisma_client: "PrismaClient",
transactions: tuple[WindowSpendTransaction, ...],
) -> frozenset[tuple[str, str, str]]:
rows: Final = await prisma_client.db.query_raw(
rows: Final = await prisma_client.replica_db.query_raw(
_SELECT_EXISTING_ROWS_SQL,
tuple(transaction["entity_type"] for transaction in transactions),
tuple(transaction["entity_id"] for transaction in transactions),

View file

@ -75,7 +75,7 @@ _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)
rows: Final = await prisma_client.replica_db.query_raw(LATEST_HEALTH_CHECKS_SQL)
return _ROWS_ADAPTER.validate_python(rows)
@ -93,7 +93,7 @@ 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))
rows: Final = await prisma_client.replica_db.query_raw(LATEST_HEALTH_CHECKS_FOR_MODELS_SQL, list(model_names))
return _ROWS_ADAPTER.validate_python(rows)
except Exception as query_err: # noqa: BLE001 # a paged model list must not fail on its health decoration
verbose_proxy_logger.error("Error getting latest health checks for models: %s", query_err)

View file

@ -80,6 +80,10 @@ class WriterPinnedClient:
def __init__(self, db: "PrismaWrapper | RoutingPrismaWrapper") -> None:
self.db: Final = db.writer if isinstance(db, RoutingPrismaWrapper) and not db.writer_unavailable else db
@property
def replica_db(self) -> "PrismaWrapper | RoutingPrismaWrapper":
return self.db
def writer_wrapper(db: "PrismaWrapper | RoutingPrismaWrapper") -> PrismaWrapper:
"""Unlike `WriterPinnedClient`, ignores `writer_unavailable`: a raw SQL write has no replica fallback."""

View file

@ -183,31 +183,31 @@ def _team_table(prisma_client: "PrismaClient") -> _TeamTable:
def _verification_tokens(prisma_client: "PrismaClient") -> _VerificationTokenTable:
return prisma_client.db.litellm_verificationtoken
return prisma_client.replica_db.litellm_verificationtoken
def _team_rows(prisma_client: "PrismaClient") -> _TeamRowsTable:
return prisma_client.db.litellm_teamtable
return prisma_client.replica_db.litellm_teamtable
def _user_rows(prisma_client: "PrismaClient") -> _UserRowsTable:
return prisma_client.db.litellm_usertable
return prisma_client.replica_db.litellm_usertable
def _shadow_eval_jobs(prisma_client: "PrismaClient") -> _ShadowEvalJobTable:
return prisma_client.db.litellm_shadowevaljob
return prisma_client.replica_db.litellm_shadowevaljob
def _shadow_eval_funnel(prisma_client: "PrismaClient") -> _ShadowEvalFunnelTable:
return prisma_client.db.litellm_shadowevalfunnel # pyright: ignore[reportAttributeAccessIssue] # generated client
return prisma_client.replica_db.litellm_shadowevalfunnel # pyright: ignore[reportAttributeAccessIssue] # generated client
def _shadow_eval_attempts(prisma_client: "PrismaClient") -> _ShadowEvalAttemptTable:
return prisma_client.db.litellm_shadowevalattempt
return prisma_client.replica_db.litellm_shadowevalattempt
async def _query_raw(prisma_client: "PrismaClient", query: str, *args: object) -> Sequence[Mapping[str, object]]:
return await prisma_client.db.query_raw(query, *args)
return await prisma_client.replica_db.query_raw(query, *args)
async def _authorize_router_dry_run(user_api_key_dict: UserAPIKeyAuth, team_id: str | None) -> LiteLLM_TeamTable | None:

View file

@ -206,7 +206,7 @@ async def _query_raw_optional(
) -> list[dict[str, object]] | None: # mutable-ok: prisma query_raw return shape
if query is None:
return None
return await prisma_client.db.query_raw(query[0], *query[1])
return await prisma_client.replica_db.query_raw(query[0], *query[1])
def _reported_flat_cost(record: DailySpendRecord | _RollupMetricsRow) -> float:
@ -832,7 +832,7 @@ def _build_aggregated_sql_query(
"""Build the GROUPING SETS query for aggregated daily activity.
Returns:
Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw().
Tuple of (sql_query, params_list) ready for prisma_client.replica_db.query_raw().
"""
pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name)
if pg_table is None:
@ -1306,7 +1306,7 @@ async def get_daily_activity(
include_current_utc_day=include_current_utc_day,
)
spend_table: Final[TableActions[DailySpendRecord]] = getattr(prisma_client.db, table_name)
spend_table: Final[TableActions[DailySpendRecord]] = getattr(prisma_client.replica_db, table_name)
# Get total count for pagination
total_count: Final[int] = await spend_table.count(where=where_conditions)
@ -1471,7 +1471,7 @@ async def get_daily_activity_aggregated(
entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None
raw_rows, raw_entity_rows = await asyncio.gather(
prisma_client.db.query_raw(sql_query, *sql_params),
prisma_client.replica_db.query_raw(sql_query, *sql_params),
_query_raw_optional(prisma_client, entity_query),
)

View file

@ -222,7 +222,7 @@ async def _migrate_config_settings_row(
dict with selected sensitive fields (vantage_settings / cloudzero_settings).
"""
report: Final = LocationReport(location=param_name)
record: Final = await prisma_client.db.litellm_config.find_unique(where={"param_name": param_name})
record: Final = await prisma_client.replica_db.litellm_config.find_unique(where={"param_name": param_name})
if record is None or record.param_value is None:
return report
@ -274,7 +274,7 @@ async def _migrate_sso_config(prisma_client: object, dry_run: bool) -> LocationR
every present string field.
"""
report: Final = LocationReport(location="sso_config")
record: Final = await prisma_client.db.litellm_ssoconfig.find_unique(where={"id": "sso_config"})
record: Final = await prisma_client.replica_db.litellm_ssoconfig.find_unique(where={"id": "sso_config"})
if record is None or record.sso_settings is None:
return report
@ -341,10 +341,10 @@ async def _migrate_callback_vars_table(
report: Final = LocationReport(location=f"{table_name}.callback_vars")
if table_name == "team":
table = prisma_client.db.litellm_teamtable
table = prisma_client.replica_db.litellm_teamtable
pk = "team_id"
else:
table = prisma_client.db.litellm_verificationtoken
table = prisma_client.replica_db.litellm_verificationtoken
pk = "token"
rows: Final = await table.find_many()
@ -510,7 +510,7 @@ async def _scan_one_table(
scalar_columns: tuple,
) -> LocationReport:
report: Final = LocationReport(location=location)
table: Final = getattr(prisma_client.db, db_attr, None)
table: Final = getattr(prisma_client.replica_db, db_attr, None)
if table is None:
return report
try:
@ -544,7 +544,9 @@ async def _scan_config_env_vars(prisma_client: object) -> LocationReport:
"""Scan the ``environment_variables`` config row (``param_value`` dict)."""
report: Final = LocationReport(location="config_environment_variables")
try:
record: Final = await prisma_client.db.litellm_config.find_unique(where={"param_name": "environment_variables"})
record: Final = await prisma_client.replica_db.litellm_config.find_unique(
where={"param_name": "environment_variables"}
)
except Exception as e: # pragma: no cover - defensive
verbose_proxy_logger.debug("scan: config env vars unavailable: %s", str(e))
return report

View file

@ -122,7 +122,7 @@ async def get_gateway_daily_activity(
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
default_start, default_end = _default_range()
raw_rows: Final = await prisma_client.db.query_raw( # pyright: ignore[reportAny] # untyped prisma client
raw_rows: Final = await prisma_client.replica_db.query_raw( # pyright: ignore[reportAny] # untyped prisma client
_AGGREGATE_SQL,
start_date or default_start,
end_date or default_end,

View file

@ -1169,7 +1169,7 @@ async def user_info_v2(
async def _fetch_admin_teams_and_keys_rows(
prisma_client: "PrismaClient", sql_query: str
) -> Sequence[Mapping[str, Sequence[Mapping[str, object]] | None]]:
return await prisma_client.db.query_raw(sql_query)
return await prisma_client.replica_db.query_raw(sql_query)
async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth):

View file

@ -6733,7 +6733,9 @@ async def key_aliases(
where_sql: Final = " AND ".join(where_parts)
count_sql: Final = f'SELECT COUNT(*) AS count FROM "LiteLLM_VerificationToken" WHERE {where_sql}'
count_rows: Final[Sequence[Mapping[str, int]]] = await prisma_client.db.query_raw(count_sql, *query_params)
count_rows: Final[Sequence[Mapping[str, int]]] = await prisma_client.replica_db.query_raw(
count_sql, *query_params
)
total_count: Final = int(count_rows[0]["count"]) if count_rows else 0
aliases_params: Final = query_params + [size, (page - 1) * size]
@ -6746,7 +6748,9 @@ async def key_aliases(
f" ORDER BY key_alias ASC"
f" LIMIT ${limit_idx} OFFSET ${offset_idx}"
)
alias_rows: Final[Sequence[Mapping[str, str]]] = await prisma_client.db.query_raw(aliases_sql, *aliases_params)
alias_rows: Final[Sequence[Mapping[str, str]]] = await prisma_client.replica_db.query_raw(
aliases_sql, *aliases_params
)
aliases: Final[list[str]] = [row["key_alias"] for row in alias_rows if row.get("key_alias")]
total_pages: Final = -(-total_count // size) if total_count > 0 else 0

View file

@ -85,7 +85,7 @@ class PrismaBudgetListExecutor:
async def count(self, where: tuple[Predicate, ...]) -> int:
clauses, params = where_sql(where)
sql: Final = f"SELECT COUNT(*) AS count FROM {BUDGET_TABLE}" + (f" WHERE {clauses}" if clauses else "")
rows: Final = await self.prisma_client.db.query_raw(sql, *params)
rows: Final = await self.prisma_client.replica_db.query_raw(sql, *params)
counted: Final = _ROW_COUNTS.validate_python(rows)
return counted[0].count if counted else 0
@ -97,7 +97,7 @@ class PrismaBudgetListExecutor:
+ f" ORDER BY {order_by_sql(plan.order)}"
+ f" LIMIT ${len(params) + 1} OFFSET ${len(params) + 2}"
)
rows: Final = await self.prisma_client.db.query_raw(sql, *params, plan.take, plan.skip)
rows: Final = await self.prisma_client.replica_db.query_raw(sql, *params, plan.take, plan.skip)
return _BUDGET_ROWS.validate_python(rows)

View file

@ -143,7 +143,7 @@ async def _list_spend_log_facet(
f" ORDER BY {column_sql} ASC"
f" LIMIT ${scan_idx + 1} OFFSET ${scan_idx + 2}"
)
rows: Final = await prisma_client.db.query_raw(facet_sql, *params)
rows: Final = await prisma_client.replica_db.query_raw(facet_sql, *params)
values: Final[list[str]] = [row[column] for row in rows if row.get(column)]
has_more: Final = len(values) > page_size

View file

@ -223,6 +223,10 @@ class _ModelTransactionClient(BaseModel):
class _TransactionClient:
db: _TxModelTables
@property
def replica_db(self) -> _TxModelTables:
return self.db
_RowT = TypeVar("_RowT")
@ -260,7 +264,7 @@ def _repo_team_table(prisma_client: PrismaClient) -> _TeamLookupTable:
def _db_team_table(prisma_client: PrismaClient) -> _TeamTable:
return prisma_client.db.litellm_teamtable
return prisma_client.replica_db.litellm_teamtable
def _model_alias_table(prisma_client: PrismaClient) -> "TableActions[prisma_models.LiteLLM_ModelTable]":

View file

@ -26,7 +26,7 @@ class _ModelDb(Protocol):
class _PrismaClient(Protocol):
@property
@abstractmethod
def db(self) -> _ModelDb:
def replica_db(self) -> _ModelDb:
pass
@ -112,7 +112,7 @@ async def validate_router_settings_weights(
return
if prisma_client is None:
raise HTTPException(status_code=503, detail="Database unavailable while validating router weights")
stored_models: Final = await prisma_client.db.litellm_proxymodeltable.find_many(
stored_models: Final = await prisma_client.replica_db.litellm_proxymodeltable.find_many(
where={"model_id": {"in": list(deployment_ids)}}
)
stored_by_id: Final = {

View file

@ -6894,7 +6894,7 @@ 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(
rows: Final[Sequence[_TeamUserSpendDbRow]] = await prisma_client.replica_db.query_raw(
_team_user_spend_sql(team_count=len(scoped_team_ids), restrict_to_user=own_user_only),
start_date,
end_date,

View file

@ -156,7 +156,7 @@ def _typed_table(repo: DailyTagSpendRepository | VerificationTokenRepository | U
async def _query_raw(prisma_client: "PrismaClient", sql_query: str, *params: object) -> object:
return await prisma_client.db.query_raw(sql_query, *params)
return await prisma_client.replica_db.query_raw(sql_query, *params)
@router.get(

View file

@ -10124,7 +10124,7 @@ class ProxyStartupEvent:
from prisma.errors import UniqueViolationError
try:
config_table: Final = prisma_client.db.litellm_config
config_table: Final = prisma_client.replica_db.litellm_config
row: Final = await config_table.find_unique(
where={"param_name": TUNING_BASELINE_PARAM_NAME} # mutable-ok: Prisma rejects mappingproxy input
)
@ -15143,7 +15143,7 @@ async def model_streaming_metrics(
"""
_all_api_bases: Final = set()
db_response: Final[Sequence[_TTFTRow] | None] = await prisma_client.db.query_raw(
db_response: Final[Sequence[_TTFTRow] | None] = await prisma_client.replica_db.query_raw(
sql_query, _selected_model_group, startTime, endTime
)
_daily_entries: dict = {} # {"Jun 23": {"model1": 0.002, "model2": 0.003}}
@ -15267,7 +15267,7 @@ async def model_metrics(
avg_latency_per_token DESC;
"""
_all_api_bases: Final = set()
db_response: Final[Sequence[_LatencyRow] | None] = await prisma_client.db.query_raw(
db_response: Final[Sequence[_LatencyRow] | None] = await prisma_client.replica_db.query_raw(
sql_query, _selected_model_group, startTime, endTime, api_key, customer
)
_daily_entries: dict = {} # {"Jun 23": {"model1": 0.002, "model2": 0.003}}
@ -15385,7 +15385,7 @@ ORDER BY
slow_count DESC;
"""
db_response: Final = await prisma_client.db.query_raw(
db_response: Final = await prisma_client.replica_db.query_raw(
sql_query,
alerting_threshold,
_selected_model_group,
@ -15458,7 +15458,7 @@ async def model_metrics_exceptions(
ORDER BY total_exceptions DESC
LIMIT 200;
"""
db_response: Final[Sequence[_ExceptionRow] | None] = await prisma_client.db.query_raw(
db_response: Final[Sequence[_ExceptionRow] | None] = await prisma_client.replica_db.query_raw(
sql_query, startTime, endTime, _selected_model_group, api_key
)
response: Final[list[dict]] = []

View file

@ -177,7 +177,7 @@ async def _advance_marker(prisma_client: "PrismaClient", days: tuple[str, ...],
async def _db_now(prisma_client: "PrismaClient") -> _NowRow:
rows: Final = await prisma_client.db.query_raw(_DB_NOW_SQL)
rows: Final = await prisma_client.replica_db.query_raw(_DB_NOW_SQL)
return _NowRow.model_validate(rows[0])
@ -189,9 +189,9 @@ async def _scan_pending(prisma_client: "PrismaClient") -> _PendingScan:
db_now: Final = await _db_now(prisma_client)
last_closed_day: Final = (date.fromisoformat(db_now.today) - timedelta(days=1)).isoformat()
rows: Final = (
await prisma_client.db.query_raw(_ALL_CLOSED_DAYS_SQL, last_closed_day)
await prisma_client.replica_db.query_raw(_ALL_CLOSED_DAYS_SQL, last_closed_day)
if marker is None or marker.scanned_at is None
else await prisma_client.db.query_raw(
else await prisma_client.replica_db.query_raw(
_PENDING_DAYS_SQL, last_closed_day, marker.reconciled_through, marker.scanned_at
)
)

View file

@ -134,7 +134,7 @@ async def _reverse_hash_key_metadata(
warning: str,
) -> Mapping[str, KeyMetadataDict]:
rows: Final = await _db_or_empty(
lambda: prisma_client.db.query_raw(sql, sorted(wanted)),
lambda: prisma_client.replica_db.query_raw(sql, sorted(wanted)),
warning,
len(wanted),
)

View file

@ -233,7 +233,7 @@ class _SpendDailySummaryRow(TypedDict):
async def _query_raw(prisma_client: PrismaClient, sql_query: str, *args: object) -> Sequence[_RowT]:
"""Run a raw read query and return its rows as the row type the caller declares."""
return await prisma_client.db.query_raw(sql_query, *args)
return await prisma_client.replica_db.query_raw(sql_query, *args)
async def _query_raw_or_none(prisma_client: PrismaClient, sql_query: str, *args: object) -> Sequence[_RowT] | None:
@ -2922,7 +2922,7 @@ async def ui_view_spend_logs(
)
sql_params.extend([page_size, skip])
data: Final = await prisma_client.db.query_raw(sql_query, *sql_params)
data: Final = await prisma_client.replica_db.query_raw(sql_query, *sql_params)
if request_id is not None and not is_v2 and not is_admin_view:
await _assert_user_owns_fetched_spend_rows(

View file

@ -899,7 +899,7 @@ async def _query_raw_rows(
sql_query: str,
*args: object,
) -> Sequence[Mapping[str, object]] | None:
return await prisma_client.db.query_raw(sql_query, *args)
return await prisma_client.replica_db.query_raw(sql_query, *args)
async def get_spend_by_team(

View file

@ -4511,6 +4511,11 @@ class PrismaClient:
return self.db.read_target
return self.db
@property
def replica_db(self) -> "PrismaWrapper | RoutingPrismaWrapper":
"""Explicit opt-in to read-replica routing; reads reached through it may go to the reader."""
return self.db
def tx(self, *, timeout: timedelta = _PRISMA_DEFAULT_TX_TIMEOUT) -> "TransactionManager":
"""Open an interactive transaction on the writer.
@ -4597,7 +4602,7 @@ class PrismaClient:
required_view: Final = "LiteLLM_VerificationTokenView"
expected_views_str: Final = ", ".join(f"'{view}'" for view in expected_views)
pg_schema: Final = os.getenv("DATABASE_SCHEMA", "public")
ret: Final[Sequence[_ViewCountRow]] = await self.db.query_raw(f"""
ret: Final[Sequence[_ViewCountRow]] = await self.replica_db.query_raw(f"""
WITH existing_views AS (
SELECT viewname
FROM pg_views
@ -4636,9 +4641,9 @@ class PrismaClient:
""",
)
else:
should_create_views: Final = await should_create_missing_views(db=self.db)
should_create_views: Final = await should_create_missing_views(db=self.replica_db)
if should_create_views:
await create_missing_views(db=self.db)
await create_missing_views(db=self.replica_db)
else:
# don't block execution if these views are missing
# Convert lists to sets for efficient difference calculation
@ -4692,7 +4697,7 @@ class PrismaClient:
)
return await config_table.find_first(where={key: value})
elif table_name == "spend":
return await self.db.l.find_first(where={key: value})
return await self.replica_db.l.find_first(where={key: value})
return None
try:
@ -4764,7 +4769,7 @@ class PrismaClient:
"""
stale_read_engine: Final = _StaleReadEngine.observe(self.read_db)
try:
return await self.db.query_first(sql_query, *args)
return await self.replica_db.query_first(sql_query, *args)
except Exception as e:
if "cached plan must not change result type" not in str(e):
raise
@ -4779,7 +4784,7 @@ class PrismaClient:
force_recreate=True,
stale_read_engine=stale_read_engine,
)
return await self.db.query_first(sql_query, *args)
return await self.replica_db.query_first(sql_query, *args)
@backoff.on_exception(
backoff.expo,
@ -4967,7 +4972,7 @@ class PrismaClient:
LIMIT $1
OFFSET $2
"""
response = await self.db.query_raw(sql_query, limit, offset)
response = await self.replica_db.query_raw(sql_query, limit, offset)
return response
elif table_name == "spend":
verbose_proxy_logger.debug("PrismaClient: get_data: table_name == 'spend'")
@ -5113,7 +5118,9 @@ class PrismaClient:
# If not found in main table, check deprecated keys (grace period)
# check_deprecated=False on the recursive call prevents unbounded chaining
if response is None and hashed_token is not None and check_deprecated:
active_token_id: Final = await _lookup_deprecated_key(db=self.db, hashed_token=hashed_token)
active_token_id: Final = await _lookup_deprecated_key(
db=self.replica_db, hashed_token=hashed_token
)
if active_token_id:
# The recursive call returns a finished
# LiteLLM_VerificationTokenView; the dict
@ -6557,7 +6564,7 @@ class PrismaClient:
async def _view_setup_gate_table_present(self) -> bool:
rows: Final = _VIEW_SETUP_GATE_PROBE_ROWS.validate_python(
await self.db.query_raw("SELECT to_regclass($1) IS NOT NULL AS present", _VIEW_SETUP_GATE_TABLE)
await self.replica_db.query_raw("SELECT to_regclass($1) IS NOT NULL AS present", _VIEW_SETUP_GATE_TABLE)
)
return rows[0]["present"]
@ -6566,7 +6573,7 @@ class PrismaClient:
try:
await asyncio.sleep(self._db_health_watchdog_interval_seconds)
await asyncio.wait_for(
self.db.query_raw("SELECT 1"),
self.replica_db.query_raw("SELECT 1"),
timeout=self._db_health_watchdog_probe_timeout_seconds,
)
if isinstance(self.db, RoutingPrismaWrapper) and self.db.writer_unavailable:
@ -6778,7 +6785,7 @@ class PrismaClient:
FROM pg_class
WHERE oid = '"LiteLLM_SpendLogs"'::regclass;
"""
result: Final[Sequence[_RelTuplesRow]] = await self.db.query_raw(query=sql_query)
result: Final[Sequence[_RelTuplesRow]] = await self.replica_db.query_raw(query=sql_query)
return result[0]["reltuples"]
try:

View file

@ -15,7 +15,7 @@ if TYPE_CHECKING:
class AutoRouterSessionRepository(BaseRepository[LiteLLM_AutoRouterSession]):
@property
def table(self) -> TableActions["prisma_models.LiteLLM_AutoRouterSession"]:
return self.prisma_client.db.litellm_autoroutersession
return self.prisma_client.replica_db.litellm_autoroutersession
@property
def model_class(self) -> type[LiteLLM_AutoRouterSession]:

View file

@ -14,7 +14,7 @@ if TYPE_CHECKING:
class _BudgetDb(Protocol):
"""The single Prisma table this repository reaches for on ``prisma_client.db``."""
"""The single Prisma table this repository reaches for on ``prisma_client.replica_db``."""
@property
def litellm_budgettable(self) -> TableActions["prisma_models.LiteLLM_BudgetTable"]: ...
@ -24,7 +24,7 @@ class _PrismaClientView(Protocol):
"""The one attribute this repository reads off the untyped Prisma client wrapper."""
@property
def db(self) -> _BudgetDb: ...
def replica_db(self) -> _BudgetDb: ...
class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]):
@ -33,7 +33,7 @@ class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]):
@property
def table(self) -> TableActions["prisma_models.LiteLLM_BudgetTable"]:
client: Final[_PrismaClientView] = self.prisma_client
return client.db.litellm_budgettable
return client.replica_db.litellm_budgettable
@property
def model_class(self) -> type[LiteLLM_BudgetTable]:

View file

@ -55,7 +55,7 @@ class ConfigRepository:
@property
def _config_table(self) -> _ConfigTable:
return cast(_ConfigTable, self.prisma_client.db.litellm_config)
return cast(_ConfigTable, self.prisma_client.replica_db.litellm_config)
@property
def table(self) -> _ConfigTable:

View file

@ -27,7 +27,7 @@ class _PrismaCredentialsDb(Protocol):
class _PrismaClientView(Protocol):
@property
def db(self) -> _PrismaCredentialsDb: ...
def replica_db(self) -> _PrismaCredentialsDb: ...
class CredentialsRepository:
@ -46,7 +46,7 @@ class CredentialsRepository:
@property
def table(self) -> "_CredentialsTable":
return wrap_table_actions_for_config_sync(
actions=self.prisma_client.db.litellm_credentialstable,
actions=self.prisma_client.replica_db.litellm_credentialstable,
table_name="litellm_credentialstable",
)

View file

@ -17,7 +17,7 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]):
@property
def table(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]:
return self.prisma_client.db.litellm_objectpermissiontable
return self.prisma_client.replica_db.litellm_objectpermissiontable
@property
def model_class(self) -> type[LiteLLM_ObjectPermissionTable]:

View file

@ -14,7 +14,7 @@ if TYPE_CHECKING:
class _OrganizationDb(Protocol):
"""The single Prisma table this repository reaches for on ``prisma_client.db``."""
"""The single Prisma table this repository reaches for on ``prisma_client.replica_db``."""
@property
def litellm_organizationtable(self) -> TableActions["prisma_models.LiteLLM_OrganizationTable"]: ...
@ -24,7 +24,7 @@ class _PrismaClientView(Protocol):
"""The one attribute this repository reads off the untyped Prisma client wrapper."""
@property
def db(self) -> _OrganizationDb: ...
def replica_db(self) -> _OrganizationDb: ...
class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]):
@ -33,7 +33,7 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]):
@property
def table(self) -> TableActions["prisma_models.LiteLLM_OrganizationTable"]:
client: Final[_PrismaClientView] = self.prisma_client
return client.db.litellm_organizationtable
return client.replica_db.litellm_organizationtable
@property
def model_class(self) -> type[LiteLLM_OrganizationTable]:

View file

@ -1,7 +1,7 @@
"""
Typed Protocol seams over prisma-client-py surfaces.
Modules that reach Prisma through an untyped handle (``prisma_client.db`` or a
Modules that reach Prisma through an untyped handle (``prisma_client.replica_db`` or a
repository ``.table``) annotate against these Protocols instead of hand-rolling
private ones per file.
"""
@ -14,7 +14,7 @@ RowT_co = TypeVar("RowT_co", covariant=True)
class DatabaseClient(Protocol):
@property
def db(self) -> object: ...
def replica_db(self) -> object: ...
class TableActions(Protocol[RowT_co]):

View file

@ -18,7 +18,7 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]):
@property
def table(self) -> TableActions["prisma_models.LiteLLM_ProjectTable"]:
return self.prisma_client.db.litellm_projecttable
return self.prisma_client.replica_db.litellm_projecttable
@property
def model_class(self) -> type[LiteLLM_ProjectTable]:

View file

@ -32,7 +32,7 @@ class PrismaTableRepository(Generic[RowT_co]):
@property
def table(self) -> TableActions[RowT_co]:
actions: Final[TableActions[RowT_co]] = getattr(self.prisma_client.db, self.table_name)
actions: Final[TableActions[RowT_co]] = getattr(self.prisma_client.replica_db, self.table_name)
return wrap_table_actions_for_config_sync(actions=actions, table_name=self.table_name)

View file

@ -68,7 +68,7 @@ class _PrismaTeamDb(_TeamTables, Protocol):
class _PrismaClientView(Protocol):
@property
def db(self) -> _PrismaTeamDb: ...
def replica_db(self) -> _PrismaTeamDb: ...
_MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member])
@ -88,7 +88,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
@property
def _db(self) -> _PrismaTeamDb:
client: Final[_PrismaClientView] = self.prisma_client
return client.db
return client.replica_db
@property
def table(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]:

View file

@ -40,7 +40,7 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]):
@property
def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]:
return self.prisma_client.db.litellm_usertable
return self.prisma_client.replica_db.litellm_usertable
@property
def model_class(self) -> type[LiteLLM_UserTable]:

View file

@ -54,7 +54,7 @@ class _PrismaVerificationTokenDb(_VerificationTokenTables, Protocol):
class _PrismaClientView(Protocol):
@property
def db(self) -> _PrismaVerificationTokenDb: ...
def replica_db(self) -> _PrismaVerificationTokenDb: ...
_JSON_ENCODED_TOKEN_FIELDS: Final = (
@ -76,7 +76,7 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
@property
def _db(self) -> _PrismaVerificationTokenDb:
client: Final[_PrismaClientView] = self.prisma_client
return client.db
return client.replica_db
@property
def table(self) -> TableActions["PrismaVerificationToken"]:

View file

@ -308,7 +308,7 @@ class ResponsesSessionHandler:
for attempt in range(max_attempts):
if attempt:
await asyncio.sleep(RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL)
if spend_logs := await prisma_client.db.query_raw(query, response_id):
if spend_logs := await prisma_client.replica_db.query_raw(query, response_id):
verbose_proxy_logger.debug(
"Found the following spend logs for previous response id %s: %s",
response_id,

View file

@ -152,6 +152,7 @@ async def test_send_user_spend_alerts_sends_and_dedupes():
alerting_args={"daily_spend_per_user_threshold": 50.0, "spend_anomaly_min_spend": 1000.0},
)
mock_prisma: Final = AsyncMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{
@ -189,5 +190,6 @@ async def test_send_user_spend_alerts_noop_when_alert_types_disabled():
alerting_args={"daily_spend_per_user_threshold": 50.0},
)
mock_prisma: Final = AsyncMock()
mock_prisma.replica_db = mock_prisma.db
await slack_alerting.send_user_spend_alerts(prisma_client=mock_prisma)
mock_prisma.db.query_raw.assert_not_called()

View file

@ -62,6 +62,7 @@ def _job(**overrides) -> ActiveShadowEvalJob:
def _prisma(jobs=(), attempt_counts=(), attempt_costs=()) -> MagicMock:
costs = {job_id: {"judge_cost": judge, "shadow_cost": shadow} for job_id, judge, shadow in attempt_costs}
prisma = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=list(jobs))
prisma.db.litellm_shadowevalattempt.group_by = AsyncMock(
return_value=[

View file

@ -60,6 +60,7 @@ def _make_prisma(stored: dict, db_has_id_jag_server: bool = False):
``db_has_id_jag_server`` drives the retention gate's authoritative DB fallback;
it is wired explicitly so the gate never reads a truthy bare MagicMock."""
prisma = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=MagicMock() if db_has_id_jag_server else None)
async def _upsert(where, data):
@ -417,6 +418,7 @@ async def test_retain_none_assertion_never_consults_gate_or_store():
@pytest.mark.asyncio
async def test_retain_swallows_store_failure():
prisma = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_ssoidentityassertion.upsert = AsyncMock(side_effect=RuntimeError("db down"))
with (
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
@ -465,6 +467,7 @@ async def test_db_store_converts_a_driver_failure_into_assertion_store_unavailab
"""The live store must not let a raw driver error escape: the resolver distinguishes an outage
from an absent assertion, and only a typed failure lets it do that."""
prisma = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_ssoidentityassertion.find_unique = AsyncMock(side_effect=RuntimeError("connection refused"))
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
with pytest.raises(AssertionStoreUnavailable):

View file

@ -133,6 +133,7 @@ def _byok_key_row(server_id):
def _mock_prisma(null_rows, token_rows):
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=null_rows)
mock_prisma.db.litellm_mcpservertable.update_many = AsyncMock(return_value=MagicMock())
mock_prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=token_rows)

View file

@ -28,6 +28,7 @@ def _row(**overrides):
def _prisma(rows):
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=rows)
prisma_client.db.litellm_mcpservertable.update = AsyncMock()
return prisma_client

View file

@ -76,6 +76,7 @@ client = TestClient(app)
@pytest.fixture
def mock_prisma_client():
with patch("litellm.proxy.proxy_server.prisma_client") as mock:
mock.replica_db = mock.db
yield mock
@ -208,6 +209,7 @@ def test_agent_error_schema_consistency(
@pytest.mark.asyncio
async def test_get_agent_daily_activity_admin_param_passing(monkeypatch):
mock_prisma = AsyncMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
@ -246,6 +248,7 @@ async def test_get_agent_daily_activity_admin_param_passing(monkeypatch):
@pytest.mark.asyncio
async def test_get_agent_daily_activity_with_agent_names(monkeypatch):
mock_prisma = AsyncMock()
mock_prisma.replica_db = mock_prisma.db
mock_agent1 = MagicMock()
mock_agent1.agent_id = "agent-1"
mock_agent1.agent_name = "First Agent"
@ -303,6 +306,7 @@ async def test_attach_keys_to_agents_groups_by_agent_and_omits_secret():
agent_without_keys = _sample_agent_response(agent_id="agent-2")
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[
_Row("hash-aaa", "agent-1", "primary", "sk-...aaa"),
@ -348,6 +352,7 @@ class TestAgentByIdKeyRedaction:
test_client = _make_app_with_role(role)
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
return_value=None
)
@ -410,6 +415,7 @@ class TestAgentRBACInternalUser:
return_value=_sample_agent_response()
)
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
return_value=None
)
@ -532,6 +538,7 @@ class TestAgentRBACProxyAdminViewOnly:
key_row.key_name = "sk-...aaa"
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[key_row]
@ -578,6 +585,7 @@ class TestAgentRBACProxyAdmin:
def test_should_allow_admin_to_create_agent(self, monkeypatch):
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
self.mock_registry.get_agent_by_name = MagicMock(return_value=None)
self.mock_registry.add_agent_to_db = AsyncMock(
return_value=_sample_agent_response()
@ -660,6 +668,7 @@ class TestAgentRBACProxyAdmin:
def test_update_agent_response_never_echoes_secret(self):
"""LIT-6736: PUT /v1/agents/{id} must not echo the stored secret back."""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # test-quality-ok: proxy_server module global is the endpoint's only injection point
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
return_value={
"agent_id": "agent-123",
@ -695,6 +704,7 @@ class TestAgentRBACProxyAdmin:
def test_patch_agent_response_never_echoes_secret(self):
"""LIT-6736: PATCH /v1/agents/{id} must not echo the stored secret back."""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # test-quality-ok: proxy_server module global is the endpoint's only injection point
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
return_value={
"agent_id": "agent-123",
@ -730,6 +740,7 @@ class TestAgentRBACProxyAdmin:
"agent_card_params": _sample_agent_card_params(),
}
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
return_value=existing
)

View file

@ -117,6 +117,7 @@ class CountingRedis(RedisCache):
def _prisma(rows: dict[str, object] | None = ALL_ROWS) -> MagicMock:
prisma = MagicMock(name="prisma_client")
prisma.replica_db = prisma.db
prisma.db.query_first = AsyncMock(return_value=rows)
return prisma

View file

@ -2049,6 +2049,7 @@ async def test_auto_register_binds_api_key_to_token_hash():
)
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
@ -2105,6 +2106,7 @@ async def test_auto_register_first_request_propagates_user_email(active: bool) -
general_settings = {"enable_jwt_auth": True}
user_api_key_cache = DualCache()
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
jwt_handler = MagicMock()
jwt_handler.is_jwt.return_value = True
jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "user1"})
@ -2223,6 +2225,7 @@ async def test_auto_register_stamps_new_key_with_jwt_agent_id():
credential_ref=CredentialRef(token_id=token_hash),
)
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
@ -2280,6 +2283,7 @@ async def test_auto_register_race_loser_keeps_winners_agent_id(losing_agent_id:
credential_ref=CredentialRef(token_id=winner_hash),
)
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock(
side_effect=Exception("Unique constraint failed on the fields: (`jwt_claim_name`,`jwt_claim_value`)")
)
@ -2661,6 +2665,7 @@ class TestJWTOAuth2Coexistence:
general_settings = {"enable_jwt_auth": True}
user_api_key_cache = DualCache()
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
jwt_handler = MagicMock()
jwt_handler.is_jwt.return_value = True
jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "user1"})
@ -2759,6 +2764,7 @@ class TestJWTOAuth2Coexistence:
general_settings = {"enable_jwt_auth": True}
user_api_key_cache = DualCache()
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
jwt_handler = MagicMock()
jwt_handler.is_jwt.return_value = True
jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "mapped-user"})
@ -2837,6 +2843,7 @@ class TestJWTOAuth2Coexistence:
general_settings = {"enable_jwt_auth": True}
user_api_key_cache = DualCache()
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
jwt_handler = MagicMock()
jwt_handler.is_jwt.return_value = True
jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "jwt-principal"})
@ -2914,6 +2921,7 @@ class TestJWTOAuth2Coexistence:
general_settings = {"enable_jwt_auth": True}
user_api_key_cache = DualCache()
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
jwt_handler = MagicMock()
jwt_handler.is_jwt.return_value = True
jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "mapped-user"})
@ -4421,6 +4429,7 @@ async def _run_centralized_checks_with_key_end_user_budget(
return _end_user_budget_row(budget_id, budgets[budget_id]) if budget_id in budgets else None
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
prisma_client.db.litellm_endusertable.find_unique = AsyncMock(return_value=end_user_row)
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(side_effect=_find_budget)
@ -4687,6 +4696,7 @@ def _unrestricted_end_user_prisma(spend: float):
end_user_row.dict = lambda: {"user_id": "customer-1", "blocked": False, "spend": spend}
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=end_user_row)
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
@ -6661,8 +6671,10 @@ def _proxy_attrs_for_db_lookup():
``_user_api_key_auth_builder`` down to the DB key lookup."""
proxy_logging_obj = MagicMock()
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
return {
"prisma_client": MagicMock(),
"prisma_client": prisma_client,
"user_api_key_cache": DualCache(),
"proxy_logging_obj": proxy_logging_obj,
"master_key": "sk-test-master",
@ -7735,6 +7747,7 @@ async def test_global_proxy_spend_reads_resettable_proxy_budget_row():
proxy_budget_row = MagicMock()
proxy_budget_row.spend = 42.5
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=proxy_budget_row)
prisma_client.db.query_raw = AsyncMock(
side_effect=AssertionError("global spend must not be loaded from the fixed-30d MonthlyGlobalSpend view")
@ -7762,6 +7775,7 @@ async def test_global_proxy_spend_none_when_proxy_budget_row_missing():
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
result = await _fetch_global_spend_with_event_coordination(
@ -8516,7 +8530,7 @@ def _per_issuer_virtual_key_jwt_handler(
def _fake_prisma_with_jwt_key_mapping(hashed_token: str | None) -> tuple[SimpleNamespace, AsyncMock]:
"""Every ``find_first`` call (issuer-scoped or global fallback) resolves the same way."""
find_first = AsyncMock(return_value=None if hashed_token is None else SimpleNamespace(token=hashed_token))
prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first)))
prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first)), replica_db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first)))
return prisma_client, find_first
@ -8535,7 +8549,7 @@ def _fake_prisma_jwt_key_mapping_table(rows: list[dict[str, object]]) -> tuple[S
return None
find_first = AsyncMock(side_effect=_find_first)
prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first)))
prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first)), replica_db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first)))
return prisma_client, find_first

View file

@ -156,6 +156,7 @@ class MockPrismaClient:
}
self.get_data_calls: List[Dict[str, Any]] = []
self.db = MockDB()
self.replica_db = self.db
async def get_data(self, table_name, query_type, **kwargs):
self.get_data_calls.append({"table_name": table_name, "query_type": query_type, **kwargs})
@ -933,6 +934,7 @@ def _make_reset_budget_windows_job(
Returns (job, prisma_client_mock, spend_counter_cache_mock).
"""
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
async def fake_query_raw(query: str, *args, **kwargs):
# Dispatch by table name in the SQL so a single stub covers both calls.
@ -1239,6 +1241,7 @@ def test_reset_budget_windows_query_error_does_not_break_team_path(monkeypatch):
expired = (now - timedelta(minutes=1)).isoformat() + "Z"
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
async def fake_query_raw(query: str, *args, **kwargs):
if '"LiteLLM_VerificationToken"' in query:
@ -1468,6 +1471,7 @@ def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch):
now = datetime.now(timezone.utc)
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
matching_key = type(
"Key",
@ -1940,6 +1944,7 @@ def _job_with_expired_budget(db, proxy_logging=None):
something to invalidate and its absence is a real signal."""
prisma_client = MockPrismaClient()
prisma_client.db = db
prisma_client.replica_db = prisma_client.db
prisma_client.data["budget"] = [_budget_row(budget_id="budget-1", budget_duration="7d")]
db.litellm_tagtable.set_find_many_results([type("Tag", (), {"tag_name": "tenant-42"})])
job = ResetBudgetJob(
@ -2344,6 +2349,7 @@ def test_budget_table_reset_stops_when_the_cascade_fails(monkeypatch):
monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2)
client, job = _chunked_job({"budget": [[_budget_row("b1"), _budget_row("b2")]]})
client.db = FailingCommitDB()
client.replica_db = client.db
asyncio.run(job.reset_budget_for_litellm_budget_table())
@ -2620,6 +2626,7 @@ def _paginating_window_job(monkeypatch, pages_by_table: Dict[str, List[List[Dict
Returns (job, calls) where calls is a list of (sql, cursor, limit).
"""
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
remaining = {table: list(pages) for table, pages in pages_by_table.items()}
calls: List[Dict[str, Any]] = []
@ -2696,6 +2703,7 @@ def test_reset_budget_windows_pages_to_the_end_of_a_large_table(monkeypatch):
def test_reset_budget_windows_survives_one_table_failing(monkeypatch):
"""A broken key scan must not cost the team scan its sweep."""
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
async def fake_query_raw(query: str, *args, **kwargs):
if '"LiteLLM_VerificationToken"' in query:
@ -2759,6 +2767,7 @@ def _cursor_paginating_window_job(monkeypatch, key_rows: List[Dict[str, Any]]):
genuinely re-reads the same prefix.
"""
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
ordered = sorted(key_rows, key=lambda r: r["token"])
visited: List[str] = []

View file

@ -68,6 +68,7 @@ class _FakeDB:
class _FakePrismaClient:
def __init__(self, db: _FakeDB) -> None:
self.db = db
self.replica_db = self.db
class _RecordingAggregate:

View file

@ -14,6 +14,7 @@ from litellm.proxy.db.health_check_latest import (
def _prisma(rows):
prisma = MagicMock()
prisma.replica_db = prisma.db
prisma.db.query_raw = AsyncMock(return_value=rows)
return prisma

View file

@ -76,6 +76,7 @@ def query_raw(monkeypatch):
"""Mocks the one call the executor makes. `count` reads the first result, `find_many` the second."""
mock = AsyncMock(side_effect=[[{"count": 0}], []])
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.query_raw = mock
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
return mock

View file

@ -51,6 +51,7 @@ WINDOW = "filter[startTime][gte]=2026-07-23T00:00:00Z&filter[startTime][lte]=202
@pytest.fixture
def mock_prisma_client(monkeypatch):
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.query_raw = AsyncMock(return_value=[])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
return prisma_client

View file

@ -643,7 +643,8 @@ class TestAutoRouterBenchmarks:
async def query_raw(self, sql: str, *params: object):
return rows
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})())
db = _DB()
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": db, "replica_db": db})())
monkeypatch.setattr(proxy_server, "llm_router", type("R", (), {"model_list": model_list})())
return await get_auto_router_benchmarks(
user_api_key_dict=ADMIN,
@ -817,7 +818,7 @@ class TestAutoRouterBenchmarks:
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks
query: Final = AsyncMock(return_value=[])
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(query_raw=query)))
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(query_raw=query), replica_db=SimpleNamespace(query_raw=query)))
app: Final = FastAPI()
app.get("/auto_router/benchmarks")(get_auto_router_benchmarks)
app.dependency_overrides[user_api_key_auth] = lambda: ADMIN
@ -858,7 +859,8 @@ class TestAutoRouterBenchmarks:
captured["params"] = params
return [TestAutoRouterBenchmarks.ROW.model_dump()]
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})())
db = _DB()
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": db, "replica_db": db})())
response = await get_auto_router_benchmarks(
user_api_key_dict=UserAPIKeyAuth(user_role=role, api_key="sk-admin", user_id="viewer"),
@ -921,7 +923,8 @@ class TestAutoRouterBenchmarks:
async def query_raw(self, sql: str, *params: object):
return [{**TestAutoRouterBenchmarks.ROW.model_dump(), "tier_turns": wire_value}]
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})())
db = _DB()
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": db, "replica_db": db})())
response = await get_auto_router_benchmarks(
user_api_key_dict=ADMIN,
@ -1091,10 +1094,11 @@ class TestAutoRouterSession:
]
return max(matching, key=lambda r: r["last_turn_at"], default=None)
db = type("D", (), {"litellm_autoroutersession": _Table()})()
monkeypatch.setattr(
proxy_server,
"prisma_client",
type("P", (), {"db": type("D", (), {"litellm_autoroutersession": _Table()})()})(),
type("P", (), {"db": db, "replica_db": db})(),
)
return lookups
@ -1354,6 +1358,7 @@ def _shadow_prisma(
direction sees the opposite-direction legs a key may hold at the same time, and a
group read that matched on a leg id would come back empty."""
prisma = MagicMock()
prisma.replica_db = prisma.db
teams: Final = key_teams or {}
team_aliases: Final = known_teams or {}
user_emails: Final = known_users or {}
@ -3085,6 +3090,7 @@ async def test_routing_test_never_confirms_models_the_caller_cannot_use(monkeypa
team_row.model_dump.return_value = row_data
team_row.dict.return_value = row_data
prisma = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
return prisma
@ -3146,6 +3152,7 @@ async def test_validate_config_gates_like_the_write_it_rehearses(monkeypatch: py
"members_with_roles": [{"role": "admin", "user_id": "team-admin"}],
}
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(proxy_server, "premium_user", True)
@ -3178,6 +3185,7 @@ def _configure_member_preview(monkeypatch: pytest.MonkeyPatch, *, allowed: bool
team_member_permissions=["/auto_router/manage"] if allowed else [],
)
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=None)
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
@ -3615,7 +3623,7 @@ async def test_availability_counts_db_and_yaml_without_disclosing_router_names(m
monkeypatch.setattr(
proxy_server,
"prisma_client",
SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many))),
SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many)), replica_db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many))),
)
monkeypatch.setattr(proxy_server.proxy_config, "auto_router_db_catalog", build_auto_router_catalog((row,)))
monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(config_deployments=lambda: (yaml_row,)))
@ -3659,7 +3667,7 @@ async def test_availability_denies_another_teams_edit_exemption(monkeypatch):
monkeypatch.setattr(
proxy_server,
"prisma_client",
SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many))),
SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many)), replica_db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many))),
)
monkeypatch.setattr(proxy_server.proxy_config, "auto_router_db_catalog", build_auto_router_catalog((row,)))
monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(config_deployments=lambda: ()))

View file

@ -42,6 +42,7 @@ async def test_get_daily_activity_empty_entity_id_list():
# Mock PrismaClient
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# Mock the table methods
mock_table = MagicMock()
@ -94,6 +95,7 @@ async def test_get_daily_activity_order_has_id_tiebreaker():
"""
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=0)
mock_table.find_many = AsyncMock(return_value=[])
@ -147,6 +149,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
# Mock PrismaClient
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# query_raw now returns rollup rows produced by GROUPING SETS, each
# tagged with its grouping level via GROUPING_ID(). The dispatcher
@ -312,6 +315,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
async def test_get_api_key_metadata_returns_active_key_metadata():
"""Test that get_api_key_metadata should return metadata for active keys."""
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# Mock active key record
mock_active_key = MagicMock()
@ -335,6 +339,7 @@ async def test_get_api_key_metadata_returns_active_key_metadata():
async def test_get_api_key_metadata_falls_back_to_deleted_keys():
"""Test that get_api_key_metadata should fall back to deleted keys table for missing keys."""
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# No active keys found
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
@ -367,6 +372,7 @@ async def test_get_api_key_metadata_falls_back_to_deleted_keys():
async def test_get_api_key_metadata_mixed_active_and_deleted_keys():
"""Test that get_api_key_metadata should return metadata for both active and deleted keys."""
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# One active key found
mock_active_key = MagicMock()
@ -401,6 +407,7 @@ async def test_get_api_key_metadata_mixed_active_and_deleted_keys():
async def test_get_api_key_metadata_deleted_table_not_queried_when_all_keys_found():
"""Test that get_api_key_metadata should not query deleted table when all keys are active."""
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_active_key = MagicMock()
mock_active_key.token = "key-hash-1"
@ -426,6 +433,7 @@ async def test_get_api_key_metadata_deleted_table_not_queried_when_all_keys_foun
async def test_get_api_key_metadata_deleted_table_error_handled_gracefully():
"""Test that get_api_key_metadata should handle errors from deleted table gracefully."""
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# No active keys found
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
@ -446,6 +454,7 @@ async def test_get_api_key_metadata_deleted_table_error_handled_gracefully():
async def test_get_api_key_metadata_regenerated_key_uses_most_recent_deleted_record():
"""Test that get_api_key_metadata should use the most recent deleted record for regenerated keys."""
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# No active keys found (old hash no longer in active table after regeneration)
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
@ -486,6 +495,7 @@ async def test_get_api_key_metadata_recovers_double_hashed_key_via_reverse_hash(
double_hashed = hash_token("a" * 64)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable.find_many = AsyncMock(
@ -515,6 +525,7 @@ async def test_get_api_key_metadata_permanent_miss_never_pages_tokens_or_reads_s
double_hashed = hash_token("b" * 64)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[])
@ -563,6 +574,7 @@ async def test_get_api_key_metadata_permanent_miss_with_a_window_reads_spend_log
double_hashed = hash_token("permanent-miss-with-window-6852")
window = (datetime(2024, 1, 1), datetime(2024, 1, 4))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[])
@ -586,6 +598,7 @@ async def test_get_daily_activity_recovers_a_session_key_alias_from_spend_logs_a
records = [_daily_user_spend_record(user_id="session-user", api_key=session_digest, spend=1.5)]
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=len(records))
mock_table.find_many = AsyncMock(return_value=records)
@ -738,6 +751,7 @@ async def test_tag_daily_activity_metadata_totals_not_zero():
"""
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# Create mock tag spend records (request_id is NULL for aggregated rows)
mock_record_1 = MagicMock()
@ -841,6 +855,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
"""Test that the full aggregation pipeline should preserve metadata for deleted keys."""
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# GROUPING SETS rollup rows. The api_key metadata lookup is driven
# by any non-NULL api_key in the result set, so the (date, endpoint,
@ -935,6 +950,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
async def test_aggregated_activity_flags_only_keys_that_key_info_can_still_resolve():
"""/key/info reads the active key table only, so deleted and never-stored (session) keys must not claim to exist."""
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
base = {
"date": "2024-01-01",
"endpoint": "/v1/chat/completions",
@ -1033,6 +1049,7 @@ async def test_get_daily_activity_applies_resolve_entity_metadata_to_breakdown()
"""
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
records = [
_daily_user_spend_record(user_id="user-with-email", api_key="key-1", spend=7.0),
@ -1088,6 +1105,7 @@ async def test_model_groups_breakdown_keys_by_public_name_with_model_fallback():
"""
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
records = [
_daily_user_spend_record(user_id="u1", api_key="key-1", spend=7.0, model="gpt-5.2", model_group="gpt-5.2-eu"),
@ -1398,6 +1416,7 @@ async def test_get_daily_activity_aggregated_empty_result_set():
"""
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_rows = [
{
@ -1569,6 +1588,7 @@ async def test_get_daily_activity_aggregated_bounds_api_key_rollups(
row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts)
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
@ -1638,6 +1658,7 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both
row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts)
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
@ -1666,6 +1687,7 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both
def _prisma_with_marker(marker: str | None) -> MagicMock:
prisma = MagicMock()
prisma.db = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
row = (
@ -1857,6 +1879,7 @@ async def test_get_daily_activity_aggregated_reports_exact_limit_key_count_as_co
row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts)
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
@ -1906,6 +1929,7 @@ async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_mo
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, [])
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
@ -2560,6 +2584,7 @@ class TestPtuCostAttributionDisabled:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=2)
mock_table.find_many = AsyncMock(
@ -2598,6 +2623,7 @@ class TestPtuCostAttributionDisabled:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=2)
mock_table.find_many = AsyncMock(
@ -2727,6 +2753,7 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown():
query's rollup dispatch."""
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
base = {
"model": None,
@ -2831,6 +2858,7 @@ async def test_get_api_key_metadata_resolves_session_key_via_spend_log_window():
session_digest = hash_token("cli-session-user-42")
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable.find_many = AsyncMock(

View file

@ -133,6 +133,7 @@ async def test_migrate_requires_aes_gate(salt_key, monkeypatch):
def _config_prisma(record):
"""Build an AsyncMock prisma client whose litellm_config returns `record`."""
client = MagicMock()
client.replica_db = client.db
client.db.litellm_config.find_unique = AsyncMock(return_value=record)
client.db.litellm_config.update = AsyncMock()
return client
@ -218,6 +219,7 @@ async def test_sso_walker_real_run_migrates_and_clears_residual(salt_key, monkey
_enable_aes(monkeypatch)
record = SimpleNamespace(sso_settings={"client_secret": legacy, "client_id": "id"})
client = MagicMock()
client.replica_db = client.db
client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=record)
client.db.litellm_ssoconfig.update = AsyncMock()
@ -235,6 +237,7 @@ async def test_sso_walker_dry_run_reports_residual_not_migrated(salt_key, monkey
_enable_aes(monkeypatch)
record = SimpleNamespace(sso_settings={"client_secret": legacy})
client = MagicMock()
client.replica_db = client.db
client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=record)
client.db.litellm_ssoconfig.update = AsyncMock()
@ -254,6 +257,7 @@ async def test_check_reports_residual_legacy(salt_key, monkeypatch):
_enable_aes(monkeypatch)
client = MagicMock()
client.replica_db = client.db
# Net-new walker tables: empty team / token / sso, one legacy vantage field.
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
@ -278,6 +282,7 @@ async def test_check_reports_residual_legacy(salt_key, monkeypatch):
async def test_check_reports_zero_after_migration(salt_key, monkeypatch):
_enable_aes(monkeypatch)
client = MagicMock()
client.replica_db = client.db
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None)
@ -314,6 +319,7 @@ async def test_callback_vars_walker_migrates_team_metadata(salt_key, monkeypatch
team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta)
client = MagicMock()
client.replica_db = client.db
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
client.db.litellm_teamtable.update = AsyncMock()
@ -342,6 +348,7 @@ async def test_callback_vars_walker_dry_run_reports_legacy(salt_key, monkeypatch
team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta)
client = MagicMock()
client.replica_db = client.db
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
client.db.litellm_teamtable.update = AsyncMock()
@ -379,6 +386,7 @@ async def test_callback_vars_walker_migrates_callback_settings_shape(
team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta)
client = MagicMock()
client.replica_db = client.db
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
client.db.litellm_teamtable.update = AsyncMock()
@ -414,6 +422,7 @@ async def test_check_reports_callback_var_legacy_with_gate_off(salt_key, monkeyp
team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta)
client = MagicMock()
client.replica_db = client.db
_empty_covered_tables(client)
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
@ -439,6 +448,7 @@ async def test_scan_covered_tables_classifies_legacy_and_v2(salt_key, monkeypatc
v2 = encrypt_value_helper("cred-secret")
client = MagicMock()
client.replica_db = client.db
_empty_covered_tables(client)
client.db.litellm_proxymodeltable.find_many = AsyncMock(
return_value=[
@ -489,6 +499,7 @@ async def test_check_classifies_mcp_secret_maps(
value: Final = json.dumps(cases[case]) if as_json and case != "null" else cases[case]
row: Final = SimpleNamespace(**{column: value})
client: Final = MagicMock()
client.replica_db = client.db
_empty_covered_tables(client)
client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
client.db.litellm_mcpservertable.update = AsyncMock()
@ -526,6 +537,7 @@ async def test_check_counts_covered_table_residual(salt_key, monkeypatch):
_enable_aes(monkeypatch)
client = MagicMock()
client.replica_db = client.db
_empty_covered_tables(client)
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
@ -554,6 +566,7 @@ async def test_migrate_covered_tables_reports_real_counts(salt_key, monkeypatch)
row = SimpleNamespace(litellm_params={"api_key": legacy})
client = MagicMock()
client.replica_db = client.db
_empty_covered_tables(client)
client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[row])
client.db.litellm_config.find_unique = AsyncMock(return_value=None)

View file

@ -77,6 +77,7 @@ def _admin() -> UserAPIKeyAuth:
def _prisma_returning(rows: list) -> MagicMock:
client = MagicMock()
client.db = MagicMock()
client.replica_db = client.db
client.db.query_raw = AsyncMock(return_value=rows)
return client

View file

@ -55,6 +55,7 @@ async def test_ui_view_users_with_null_email(mocker, caplog):
"""
# Mock the prisma client
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Create mock user data with null email
mock_user = mocker.MagicMock()
@ -97,6 +98,7 @@ async def test_ui_view_users_proxy_admin_no_org_filter(mocker):
Proxy admin: find_many is called without organization_memberships in where.
"""
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
async def mock_find_many(*args, **kwargs):
assert "organization_memberships" not in (kwargs.get("where") or {})
@ -165,6 +167,7 @@ def test_ui_view_users_search_matches_user_id_or_email(
return [user for user in users if _matches_user_where(user, where)]
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_usertable.find_many = mock_find_many
mocker.patch( # test-quality-ok: endpoint reads settings via module global; same seam as sibling tests
"litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.get_ui_settings_cached",
@ -194,6 +197,7 @@ async def test_ui_view_users_org_admin_filtered_by_org(mocker):
from litellm.proxy._types import LiteLLM_OrganizationMembershipTable
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
org_id = "org-123"
async def mock_find_many(*args, **kwargs):
@ -253,6 +257,7 @@ async def test_ui_view_users_non_org_admin_returns_403(mocker):
from fastapi import HTTPException
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Flag ON
mocker.patch(
@ -296,6 +301,7 @@ async def test_ui_view_users_flag_off_internal_user_can_search(mocker):
Flag OFF (default): any authenticated user can search all users without org filtering.
"""
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
async def mock_find_many(*args, **kwargs):
where = kwargs.get("where") or {}
@ -331,6 +337,7 @@ async def test_ui_view_users_flag_on_team_admin_org_team(mocker):
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
org_id = "org-456"
tid = "team-789"
@ -402,6 +409,7 @@ async def test_ui_view_users_flag_on_team_admin_non_org_team_403(mocker):
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
tid = "team-no-org"
# Flag ON
@ -463,6 +471,7 @@ async def test_ui_view_users_flag_on_team_admin_org_member_no_team_id(mocker):
should succeed and filter by the user's org membership.
"""
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
org_id = "org-member-org"
async def mock_find_many(*args, **kwargs):
@ -521,6 +530,7 @@ async def test_ui_view_users_flag_on_team_admin_not_in_org_resolves_via_key_team
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
org_id = "org-from-team"
tid = "key-team-id"
@ -611,6 +621,7 @@ async def test_get_users_includes_timestamps(mocker):
"""
# Mock the prisma client
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Create mock user data with timestamps
mock_user_data = {
@ -677,6 +688,7 @@ async def test_get_users_redacts_scim_enterprise_metadata(mocker):
the rest of the metadata intact, matching the user-info endpoints.
"""
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_user_row = mocker.MagicMock()
mock_user_row.user_id = "listed-user"
@ -871,6 +883,7 @@ async def test_new_user_license_over_limit(mocker):
# Mock the prisma client
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# 1000 billable users (no SCIM-deactivated rows): the filtered count used
# for "deactivated" returns 0, so billable == total == 1000
@ -954,6 +967,7 @@ async def test_new_user_license_gate_counts_only_billable_users(mocker):
def _prisma(total, deactivated):
client = mocker.MagicMock()
client.replica_db = client.db
async def _count(*args, where=None, **kwargs):
return deactivated if where is not None else total
@ -991,6 +1005,7 @@ async def test_new_user_non_admin_cannot_create_admin(mocker):
# Mock the prisma client
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Setup the mock count response (under license limit)
async def mock_count(*args, **kwargs):
@ -1063,6 +1078,7 @@ async def test_new_user_non_admin_permissions_non_empty_rejected(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
async def mock_count(*args, **kwargs):
return 5
@ -1106,6 +1122,7 @@ async def test_new_user_non_admin_permissions_explicit_empty_rejected(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
async def mock_count(*args, **kwargs):
return 5
@ -1150,6 +1167,7 @@ async def test_new_user_non_admin_omits_permissions_succeeds(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
async def mock_count(*args, **kwargs):
return 5
@ -1200,6 +1218,7 @@ async def test_new_user_admin_can_set_permissions(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
async def mock_count(*args, **kwargs):
return 5
@ -1254,6 +1273,7 @@ async def test_update_single_user_non_admin_permissions_rejected(mocker):
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin")
@ -1281,6 +1301,7 @@ async def test_update_single_user_non_admin_permissions_explicit_empty_rejected(
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin")
@ -1310,6 +1331,7 @@ async def test_user_info_url_encoding_plus_character(mocker):
# Mock the prisma client
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Create a real LiteLLM_UserTable instance (BaseModel) so isinstance check passes
mock_user = LiteLLM_UserTable(
@ -1382,6 +1404,7 @@ async def test_user_info_nonexistent_user(mocker):
# Mock the prisma client
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Mock get_data to return None (user doesn't exist)
async def mock_get_data(*args, **kwargs):
@ -1428,6 +1451,7 @@ async def test_user_info_no_user_id_view_only_admin_gets_proxy_admin_payload(moc
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.get_data = mocker.AsyncMock(return_value=None)
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
@ -1460,6 +1484,7 @@ async def test_new_user_default_teams_flow(mocker):
# Mock the prisma client
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Setup the mock count response (under license limit)
async def mock_count(*args, **kwargs):
@ -1705,6 +1730,7 @@ async def test_check_duplicate_user_email_case_insensitive(mocker):
# Mock the prisma client
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Test Case 1: Duplicate found with different case
# Mock existing user with uppercase email
@ -1770,6 +1796,7 @@ async def test_check_duplicate_user_id(mocker):
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Duplicate user_id should raise
mock_existing_user = mocker.MagicMock()
@ -1937,6 +1964,7 @@ async def test_get_users_user_id_partial_match(mocker):
from litellm.proxy._types import UserAPIKeyAuth
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_user_data = {
"user_id": "test-user-partial-match",
@ -2034,6 +2062,7 @@ def test_get_users_search_matches_user_id_or_email(mocker):
return 0
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_usertable.find_many = mock_find_many
mock_prisma_client.db.litellm_usertable.count = mock_count
mock_prisma_client.db.litellm_verificationtoken.count = mock_key_count
@ -2186,6 +2215,7 @@ async def test_user_model_budget_update_by_email_refreshes_cached_user(mocker: M
max_budget=50,
)
prisma_client: Final = mocker.MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user)
prisma_client.get_data = mocker.AsyncMock(return_value=[saved_user])
prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": saved_user.user_id, "data": saved_user})
@ -2225,6 +2255,7 @@ async def test_user_status_update_refreshes_cached_user(
metadata={"scim_active": False if active is None else not active, "department": "engineering"},
)
prisma_client: Final = mocker.MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user)
prisma_client.get_data = mocker.AsyncMock(return_value=[saved_user])
prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": saved_user.user_id, "data": saved_user})
@ -2263,6 +2294,7 @@ async def test_bulk_user_model_budget_clear_serializes_and_refreshes_cache(mocke
saved_user: Final = LiteLLM_UserTable(user_id="user-spruce", model_max_budget={"model-spruce": {"budget_limit": 5}})
prisma_client: Final = mocker.MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_usertable.find_many = mocker.AsyncMock(return_value=[saved_user])
prisma_client.db.litellm_usertable.update_many = mocker.AsyncMock(return_value=1)
mocker.patch("litellm.proxy.proxy_server.prisma_client", prisma_client) # test-quality-ok: substitute the database dependency
@ -2308,6 +2340,7 @@ async def test_user_max_budget_update_evicts_cached_user_on_every_worker(mocker:
saved_user: Final = LiteLLM_UserTable(user_id="user-spruce", max_budget=500.0)
prisma_client: Final = mocker.MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user)
prisma_client.db.litellm_usertable.find_many = mocker.AsyncMock(return_value=[saved_user])
prisma_client.db.litellm_usertable.update_many = mocker.AsyncMock(return_value=1)
@ -2377,6 +2410,7 @@ async def test_get_user_daily_activity_non_admin_cannot_view_other_users(monkeyp
# Mock the prisma client so the DB-not-connected check passes
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
# Non-admin caller
@ -2449,6 +2483,7 @@ async def test_get_user_daily_activity_rejects_service_account_caller(monkeypatc
)
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
# Tripwire: ensure get_daily_activity is never reached
@ -2499,6 +2534,7 @@ async def test_get_user_daily_activity_aggregated_rejects_service_account_caller
)
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_get_daily_agg = AsyncMock()
@ -2544,6 +2580,7 @@ async def test_get_user_daily_activity_aggregated_admin_global_view(monkeypatch,
# Mock the prisma client
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
# Mock the downstream helper so we don't need a real DB
@ -2610,6 +2647,7 @@ async def test_get_user_daily_activity_aggregated_non_admin_cannot_view_other_us
)
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
non_admin_key_dict = UserAPIKeyAuth(
@ -2671,6 +2709,7 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Mock user lookup
mock_user_row = mocker.MagicMock()
@ -2784,6 +2823,7 @@ async def test_delete_user_evicts_jwt_key_mapping_cache_of_its_keys(mocker):
user_row.model_dump.return_value = {"user_id": "jwt-user", "user_email": "jwt-user@example.com", "teams": []}
mock_prisma_client: Final = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(return_value=user_row)
mock_prisma_client.db.litellm_teamtable.find_many = mocker.AsyncMock(return_value=[])
mock_prisma_client.db.litellm_jwtkeymapping = jwt_table
@ -2835,6 +2875,7 @@ async def test_delete_user_rejects_org_admin_deleting_outside_scope(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Target user exists and is a member of org-B only.
mock_target_user = mocker.MagicMock()
@ -2904,6 +2945,7 @@ async def test_user_update_rejects_silent_create_for_non_proxy_admin(mocker):
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# user_email lookup yields None → would silently create pre-fix.
mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=None)
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
@ -2939,6 +2981,7 @@ async def test_user_info_v2_proxy_admin_can_query_any_user(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_user_row = mocker.MagicMock()
mock_user_row.model_dump.return_value = {
@ -3002,6 +3045,7 @@ async def test_user_info_v2_redacts_scim_enterprise_metadata(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_user_row = mocker.MagicMock()
mock_user_row.model_dump.return_value = {
@ -3083,6 +3127,7 @@ async def test_user_info_v2_internal_user_can_query_self(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_user_row = mocker.MagicMock()
mock_user_row.model_dump.return_value = {
@ -3137,6 +3182,7 @@ async def test_user_info_v2_internal_user_cannot_query_other(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Caller user has no teams (so no team admin access)
mock_caller_row = mocker.MagicMock()
@ -3177,6 +3223,7 @@ async def test_user_info_v2_no_user_id_defaults_to_self(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_user_row = mocker.MagicMock()
mock_user_row.model_dump.return_value = {
@ -3231,6 +3278,7 @@ async def test_user_info_v2_nonexistent_user_returns_404(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
async def mock_find_unique(*args, **kwargs):
return None
@ -3266,6 +3314,7 @@ async def test_user_info_v2_response_shape(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_user_row = mocker.MagicMock()
mock_user_row.model_dump.return_value = {
@ -3356,6 +3405,7 @@ async def test_user_info_v2_team_admin_can_query_team_member(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Caller (team admin)
mock_caller = mocker.MagicMock()
@ -3435,6 +3485,7 @@ async def test_user_info_v2_team_admin_cannot_query_non_team_member(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Caller (team admin of team-A)
mock_caller = mocker.MagicMock()
@ -3497,6 +3548,7 @@ async def test_user_info_v2_url_encoding_plus_character(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
expected_user_id = "machine-user+admin@example.com"
@ -3729,6 +3781,7 @@ async def test_ghsa_wvg4_non_admin_cannot_self_escalate_max_budget(mocker, budge
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": "user-1", "data": {"user_id": "user-1"}})
existing_user = mocker.MagicMock()
existing_user.model_dump.return_value = {
@ -3762,6 +3815,7 @@ async def test_ghsa_wvg4_non_admin_cannot_self_escalate_spend(mocker):
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
existing_user = mocker.MagicMock()
existing_user.model_dump.return_value = {
"user_id": "user-1",
@ -3794,6 +3848,7 @@ async def test_ghsa_wvg4_proxy_admin_can_update_user_budget(mocker):
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
existing_user = mocker.MagicMock()
existing_user.model_dump.return_value = {
"user_id": "target-user",
@ -3828,6 +3883,7 @@ async def test_admin_user_update_spend_invalidates_counter(mocker):
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
existing_user = mocker.MagicMock()
existing_user.model_dump.return_value = {"user_id": "target-user", "spend": 50.0}
existing_user.user_id = "target-user"
@ -3863,6 +3919,7 @@ async def test_user_update_rejects_non_finite_spend(mocker):
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
existing_user = mocker.MagicMock()
existing_user.model_dump.return_value = {"user_id": "target-user", "spend": 50.0}
existing_user.user_id = "target-user"
@ -3897,6 +3954,7 @@ async def test_resolve_user_email_metadata_maps_page_user_ids_to_email(mocker):
"""
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
find_many = mocker.AsyncMock(
return_value=[
SimpleNamespace(user_id="u1", user_email="alice@example.com", user_alias="Alice"),
@ -3925,6 +3983,7 @@ async def test_resolve_user_email_metadata_maps_page_user_ids_to_email(mocker):
async def test_resolve_user_email_metadata_skips_db_when_no_user_ids(mocker):
"""No user_ids on the page (e.g. all spend is unattributed) means no query."""
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
find_many = mocker.AsyncMock(return_value=[])
mock_prisma_client.db.litellm_usertable.find_many = find_many
@ -4075,6 +4134,7 @@ async def test_get_user_info_for_proxy_admin_validates_keys_and_teams():
]
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.query_raw = AsyncMock(return_value=raw_rows)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
@ -4091,6 +4151,7 @@ async def test_get_user_info_for_proxy_admin_validates_keys_and_teams():
def _object_permission_mocks(mocker, existing_object_permission_id=None):
"""Prisma double whose user row optionally already links a permission row."""
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
existing_user = mocker.MagicMock()
existing_user.model_dump.return_value = {
"user_id": "target-user",
@ -4316,6 +4377,7 @@ async def test_new_user_persists_the_requested_mcp_entitlement(mocker):
"""generate_key_helper_fn only forwards object_permission_id, so /user/new has to create the
grants row itself; otherwise the entitlement the admin sent is silently dropped."""
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_objectpermissiontable.create = mocker.AsyncMock(
return_value=SimpleNamespace(object_permission_id="perm-created")
)
@ -4459,6 +4521,7 @@ def _admin_prisma(mocker):
this file repeats per-test; consolidated here since these three share it
verbatim)."""
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mocker.patch( # test-quality-ok: same module-global mocking every test in this file already uses
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
)
@ -4697,6 +4760,7 @@ async def test_delete_user_evicts_cached_user_rows(mocker: MockerFixture) -> Non
deleted: Final = LiteLLM_UserTable(user_id="user-gone", user_email="gone@example.test", teams=[])
survivor: Final = LiteLLM_UserTable(user_id="user-stays", user_email="stays@example.test", teams=[])
prisma_client: Final = mocker.MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(return_value=deleted)
prisma_client.db.litellm_teamtable.find_many = mocker.AsyncMock(return_value=[])
prisma_client.db.litellm_jwtkeymapping.find_many = mocker.AsyncMock(return_value=[])
@ -4739,6 +4803,7 @@ _DB_OUTAGE_503_BODY: Final = {
def _user_read_raising(mocker: MockerFixture, error: Exception) -> tuple[MagicMock, MagicMock]:
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=error)
cache = MagicMock()
cache.async_get_cache = AsyncMock(return_value=None)
@ -4859,6 +4924,7 @@ async def test_delete_user_writes_deleted_audit_log_for_user_keys(mocker):
from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_user_row = mocker.MagicMock()
mock_user_row.user_id = "doomed-user"

View file

@ -57,6 +57,7 @@ class MockPrismaClient:
self.user_admin = user_admin
self.sibling_deployments = sibling_deployments or []
self.db = self
self.replica_db = self.db
async def find_unique(self, where):
if self.team_exists:
@ -385,6 +386,7 @@ class TestModelManagementAuthChecks:
)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
@ -515,6 +517,7 @@ class TestModelManagementAuthChecks:
)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
@ -608,6 +611,7 @@ class TestModelManagementAuthChecks:
existing_row = MagicMock()
existing_row.model_dump.return_value = existing.model_dump()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row)
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock()
with (
@ -702,6 +706,7 @@ class TestDeleteTeamModelAlias:
# Create mock prisma client
mock_prisma = MockPrismaClient(team_exists=True)
mock_prisma.db = MockPrismaWrapper(model_aliases_list)
mock_prisma.replica_db = mock_prisma.db
# Call the function
await delete_team_model_alias(
@ -760,6 +765,7 @@ class TestDeleteTeamModelAlias:
# Create mock prisma client
mock_prisma = MockPrismaClient(team_exists=True)
mock_prisma.db = MockPrismaWrapper(model_aliases_list)
mock_prisma.replica_db = mock_prisma.db
# Call the function with non-existent model
await delete_team_model_alias(
@ -790,6 +796,7 @@ class TestClearCache:
)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_logging = MagicMock()
with (
@ -848,6 +855,7 @@ class TestClearCache:
)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_logging = MagicMock()
with (
@ -1125,6 +1133,7 @@ class TestDeleteModelClearsRouterRegistry:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
@ -1187,6 +1196,7 @@ class TestDeleteModelClearsRouterRegistry:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
@ -1274,9 +1284,8 @@ class TestDeletedAutoRouterAvailability:
find_unique=AsyncMock(return_value=row),
delete=AsyncMock(return_value=row, side_effect=None if delete_succeeds else RuntimeError("delete failed")),
)
prisma = SimpleNamespace(
db=SimpleNamespace(litellm_proxymodeltable=table, query_raw=AsyncMock(return_value=[]))
)
db = SimpleNamespace(litellm_proxymodeltable=table, query_raw=AsyncMock(return_value=[]))
prisma = SimpleNamespace(db=db, replica_db=db)
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
@ -1374,6 +1383,7 @@ class TestUpdateModel:
updated_row.model_dump_json.return_value = "{}"
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
return_value=existing_row
)
@ -1434,6 +1444,7 @@ class TestUpdateModel:
updated_row = MagicMock()
updated_row.model_dump_json.return_value = "{}"
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row)
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row)
mock_router = MagicMock()
@ -2529,6 +2540,7 @@ class TestAddAndDeleteModelLifecycle:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=db_row)
@ -2642,6 +2654,7 @@ class TestDeleteTeamBYOKModelGhost:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
@ -2725,6 +2738,7 @@ class TestDeleteTeamBYOKModelGhost:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
@ -2802,6 +2816,7 @@ class TestDeleteTeamBYOKModelGhost:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
@ -2888,6 +2903,7 @@ class TestDeleteTeamBYOKModelGhost:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
@ -2970,6 +2986,7 @@ class TestDeleteTeamBYOKModelGhost:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
@ -3040,6 +3057,7 @@ class TestDeleteModelTeamAuth:
)
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
@ -3159,6 +3177,7 @@ class TestDeleteModelTeamAuth:
)
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
@ -3329,6 +3348,7 @@ class _TxPrismaClient:
return False
self.db = MagicMock()
self.replica_db = self.db
self.db.tx = MagicMock(return_value=_TxCM())
@ -4135,6 +4155,7 @@ class TestModelInfoServerDerivedPricingFilter:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=db_row)
@ -4980,6 +5001,7 @@ class TestPatchModelBlockedAuthGate:
existing_row.model_dump_json.return_value = "{}"
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
return_value=existing_row
)
@ -5023,6 +5045,7 @@ class TestPatchModelBlockedAuthGate:
updated_row.model_dump_json.return_value = "{}"
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
return_value=existing_row
)
@ -5078,6 +5101,7 @@ class TestPatchModelRowDeletedBeforeWrite:
existing_row.model_dump_json.return_value = "{}"
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
return_value=existing_row
)
@ -5462,6 +5486,7 @@ class TestDeleteEvictionsHoldTheReconcileLock:
table.delete = AsyncMock(return_value=row)
prisma = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_proxymodeltable = table
prisma.db.query_raw = AsyncMock(return_value=[])
@ -5517,6 +5542,7 @@ class TestDeleteEvictionsHoldTheReconcileLock:
tx_ctx.__aexit__ = AsyncMock(return_value=False)
prisma = MagicMock()
prisma.replica_db = prisma.db
prisma.db.tx = MagicMock(return_value=tx_ctx)
monkeypatch.setattr(
@ -5898,6 +5924,7 @@ class TestStrategyRouterWriteValidation:
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
@ -6043,6 +6070,7 @@ class TestStrategyRouterWriteValidation:
self, db_models: list[str], existing_row: object = None, tuning_rows: list[dict[str, object]] | None = None
) -> None:
self.db = self
self.replica_db = self.db
self.tx_obj = TestStrategyRouterWriteValidation._FakeTx(db_models, tuning_rows=tuning_rows)
self.litellm_proxymodeltable = MagicMock(
create=AsyncMock(), update=AsyncMock(), find_unique=AsyncMock(return_value=existing_row)
@ -6681,6 +6709,7 @@ class TestStrategyRouterWriteValidation:
}
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row)
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock()
@ -7097,6 +7126,7 @@ class TestBlockModelResponseSerialization:
updated_row = prisma_models.LiteLLM_ProxyModelTable(blocked=blocked, **row_fields)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row)
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row)
@ -7180,6 +7210,7 @@ class TestAccessGroupModelSync:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(side_effect=query_raw)
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=row)
@ -7505,7 +7536,7 @@ class TestTeamMemberAutoRouterWrites:
litellm_proxymodeltable=table,
tx=MagicMock(return_value=context),
)
return MagicMock(db=db, transaction=transaction)
return MagicMock(db=db, replica_db=db, transaction=transaction)
@staticmethod
def _catalog() -> Router:
@ -7727,6 +7758,7 @@ class TestModelManagementActorEdges:
actor: Final = UserAPIKeyAuth(user_id="internal-user", user_role=LitellmUserRoles.INTERNAL_USER)
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
deployment: Final = Deployment(
model_name="internal-model",
litellm_params=LiteLLM_Params(model="openai/test-model"),
@ -7754,6 +7786,7 @@ class TestModelManagementActorEdges:
user_id="view-only-user", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY
)
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
deployment: Final = Deployment(
model_name="view-only-model",
litellm_params=LiteLLM_Params(model="openai/test-model"),
@ -7779,6 +7812,7 @@ class TestModelManagementActorEdges:
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
deployment: Final = Deployment(
model_name="database-disabled-model",
litellm_params=LiteLLM_Params(model="openai/test-model"),
@ -7813,6 +7847,7 @@ class TestModelManagementActorEdges:
updated_row: Final = MagicMock()
updated_row.model_dump_json.return_value = "{}"
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row)
prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row)
router: Final = MagicMock()
@ -7862,6 +7897,7 @@ class TestModelManagementActorEdges:
updated_row: Final = MagicMock()
updated_row.model_dump_json.return_value = "{}"
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row)
prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row)
router: Final = MagicMock()
@ -7901,6 +7937,7 @@ class TestModelManagementActorEdges:
model_id: Final = "config-model-id"
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_proxymodeltable.update = AsyncMock()
router: Final = MagicMock()
@ -7944,6 +7981,7 @@ class TestModelManagementActorEdges:
def test_post_model_new_binds_to_actor_guard(self):
actor: Final = UserAPIKeyAuth(user_id="internal-user", user_role=LitellmUserRoles.INTERNAL_USER)
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam
@ -7983,6 +8021,7 @@ class TestModelManagementActorEdges:
updated_by="admin",
)
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row)
prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row)
router: Final = MagicMock()
@ -8024,6 +8063,7 @@ class TestModelManagementActorEdges:
def test_patch_config_model_binds_to_patch_route(self):
model_id: Final = "config-route-model-id"
prisma: Final = MagicMock()
prisma.replica_db = prisma.db
prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_proxymodeltable.update = AsyncMock()
router: Final = MagicMock()

View file

@ -107,6 +107,7 @@ class _FakePrisma:
self.marker_landing_on_day: dict[str, str] = {}
self.reconciled: list[str] = []
self.db = _FakeDb(self)
self.replica_db = self.db
def write_late_row(self, day: str) -> None:
"""A per-key row for ``day`` lands now, after whatever scans already happened."""

View file

@ -73,6 +73,7 @@ def _query_raw_by_table(
async def test_recover_double_hashed_key_metadata_via_active_token_digest():
double_hashed = hash_token("a" * 64)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = _query_raw_by_table(
active_rows=[_digest_row(double_hashed, "batch-worker", "team-1", "alice")],
deleted_rows=[],
@ -91,6 +92,7 @@ async def test_recover_double_hashed_key_metadata_via_active_token_digest():
async def test_recover_double_hashed_key_metadata_falls_back_to_deleted_tokens():
double_hashed = hash_token("y" * 64)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = _query_raw_by_table(
active_rows=[],
deleted_rows=[_digest_row(double_hashed, "deleted-key", "team-del", "erin")],
@ -109,6 +111,7 @@ async def test_recover_only_asks_deleted_tokens_for_digests_active_keys_missed()
found_active = hash_token("1" * 64)
found_deleted = hash_token("2" * 64)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = _query_raw_by_table(
active_rows=[_digest_row(found_active, "active-key", None, None)],
deleted_rows=[_digest_row(found_deleted, "deleted-key", None, None)],
@ -128,6 +131,7 @@ async def test_recover_only_asks_deleted_tokens_for_digests_active_keys_missed()
async def test_recover_permanent_miss_costs_two_digest_lookups_and_no_table_walk():
double_hashed = hash_token("b" * 64)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = _query_raw_by_table(active_rows=[], deleted_rows=[])
result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed})
@ -141,6 +145,7 @@ async def test_recover_permanent_miss_costs_two_digest_lookups_and_no_table_walk
@pytest.mark.asyncio
async def test_recover_skips_keys_that_are_not_sha256_digests():
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(return_value=[])
result = await recover_double_hashed_key_metadata(mock_prisma, {"sk-plain-key", "key-hash-short"})
@ -153,6 +158,7 @@ async def test_recover_skips_keys_that_are_not_sha256_digests():
async def test_recover_returns_empty_when_digest_lookup_raises_prisma_error():
double_hashed = hash_token("c" * 64)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(side_effect=PrismaError("db down"))
result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed})
@ -164,6 +170,7 @@ async def test_recover_returns_empty_when_digest_lookup_raises_prisma_error():
async def test_fill_missing_api_key_aliases_updates_null_alias_and_email_rows():
double_hashed = hash_token("d" * 64)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = _query_raw_by_table(
active_rows=[_digest_row(double_hashed, "recovered-alias", "team-9", "bob")],
deleted_rows=[],
@ -202,6 +209,7 @@ async def test_fill_missing_api_key_aliases_updates_null_alias_and_email_rows():
@pytest.mark.asyncio
async def test_fill_missing_api_key_aliases_leaves_rows_untouched_when_nothing_is_missing():
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(return_value=[])
rows = ({"api_key": hash_token("e" * 64), "api_key_alias": "named", "user_email": "x@example.com"},)
@ -215,6 +223,7 @@ async def test_fill_missing_api_key_aliases_leaves_rows_untouched_when_nothing_i
async def test_fill_missing_api_key_aliases_keeps_spend_user_email_when_alias_is_missing():
double_hashed = hash_token("f" * 64)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = _query_raw_by_table(
active_rows=[_digest_row(double_hashed, "team-key", "team-9", "key-owner")],
deleted_rows=[],
@ -243,6 +252,7 @@ async def test_fill_missing_api_key_aliases_keeps_spend_user_email_when_alias_is
@pytest.mark.asyncio
async def test_fill_missing_api_key_aliases_skips_named_keys_that_have_no_email():
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(return_value=[])
rows = (
{
@ -264,6 +274,7 @@ async def test_recover_key_metadata_from_spend_logs_resolves_session_token_from_
session_digest = hash_token("cli-session-repro-user-6852")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(
mock_prisma,
_query_raw_spend_logs(
@ -284,6 +295,7 @@ async def test_recover_key_metadata_from_spend_logs_resolves_session_token_from_
async def test_recover_key_metadata_from_spend_logs_skips_query_when_no_missing_keys():
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(mock_prisma, AsyncMock(return_value=[]))
result = await recover_key_metadata_from_spend_logs(mock_prisma, set(), window, cache=InMemoryCache())
@ -296,6 +308,7 @@ async def test_recover_key_metadata_from_spend_logs_skips_query_when_no_missing_
async def test_recover_key_metadata_from_spend_logs_returns_empty_on_prisma_error():
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(mock_prisma, AsyncMock(side_effect=PrismaError("db down")))
result = await recover_key_metadata_from_spend_logs(
@ -312,6 +325,7 @@ async def test_recover_key_metadata_from_spend_logs_ignores_foreign_and_all_null
foreign = hash_token("cli-session-foreign")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(
mock_prisma,
_query_raw_spend_logs(
@ -334,6 +348,7 @@ async def test_recover_key_metadata_from_spend_logs_ignores_foreign_and_all_null
async def test_recover_key_metadata_from_spend_logs_skips_non_sha256_keys():
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(mock_prisma, AsyncMock(return_value=[]))
result = await recover_key_metadata_from_spend_logs(
@ -349,6 +364,7 @@ async def test_recover_key_metadata_from_spend_logs_accepts_hashed_jwt_digests()
jwt_digest = f"hashed-jwt-{hash_token('jwt-subject-1')}"
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(jwt_digest, None, "team-jwt", "jwt-user")]))
result = await recover_key_metadata_from_spend_logs(mock_prisma, {jwt_digest}, window, cache=InMemoryCache())
@ -366,6 +382,7 @@ async def test_recover_key_metadata_from_spend_logs_serves_repeat_lookups_from_t
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(found, "found-alias", None, "owner-1")]))
first = await recover_key_metadata_from_spend_logs(mock_prisma, {found, unknown}, window, cache=cache)
@ -384,6 +401,7 @@ async def test_recover_key_metadata_from_spend_logs_only_queries_digests_the_cac
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(cached_digest, "cached-alias", None, None)]))
await recover_key_metadata_from_spend_logs(mock_prisma, {cached_digest}, window, cache=cache)
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(new_digest, "new-alias", None, None)]))
@ -401,6 +419,7 @@ async def test_recover_key_metadata_from_spend_logs_rescans_when_the_window_chan
digest = hash_token("cli-session-windowed")
cache = InMemoryCache()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([]))
await recover_key_metadata_from_spend_logs(
mock_prisma, {digest}, (datetime(2026, 9, 1), datetime(2026, 9, 4)), cache=cache
@ -421,6 +440,7 @@ async def test_recover_key_metadata_from_spend_logs_retries_a_failed_query_only_
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(mock_prisma, AsyncMock(side_effect=PrismaError("statement timeout")))
started = time.time()
assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {}
@ -442,6 +462,7 @@ async def test_recover_key_metadata_from_spend_logs_drops_the_owner_of_a_digest_
shared_ui_digest = hash_token("ui-token")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(
mock_prisma,
_query_raw_spend_logs(
@ -461,6 +482,7 @@ async def test_recover_key_metadata_from_spend_logs_keeps_the_owner_when_every_n
digest = hash_token("cli-session-one-owner")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(
mock_prisma,
_query_raw_spend_logs(
@ -480,6 +502,7 @@ async def test_recover_key_metadata_from_spend_logs_forgets_a_miss_long_before_a
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(found, "found-alias", None, None)]))
started = time.time()
@ -498,6 +521,7 @@ async def test_recover_key_metadata_from_spend_logs_runs_one_query_for_concurren
cache = InMemoryCache()
lock = asyncio.Lock()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
async def slow_query_raw(sql: str, *params: object) -> list[dict[str, str | None]]:
await asyncio.sleep(0.01)
@ -522,6 +546,7 @@ async def test_recover_key_metadata_from_spend_logs_keeps_a_repeated_miss_as_lon
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([]))
await recover_key_metadata_from_spend_logs(mock_prisma, {unknown}, window, cache=cache)
first_miss_key = next(key for key in cache.ttl_dict if unknown in key and not key.endswith(":missed-before"))
@ -539,6 +564,7 @@ async def test_recover_key_metadata_from_spend_logs_keeps_the_owner_older_rows_a
digest = hash_token("cli-session-owner-from-older-rows")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
_spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(digest, None, "team-x", "alice")]))
result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache())
@ -551,6 +577,7 @@ async def test_recover_key_metadata_from_spend_logs_names_nothing_for_a_field_wh
digest = hash_token("cli-session-disagreeing-rows")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
_spend_log_transaction(
mock_prisma,
_query_raw_spend_logs(
@ -576,6 +603,7 @@ async def test_recover_key_metadata_from_spend_logs_bounds_the_scan_with_a_state
digest = hash_token("cli-session-bounded-scan")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
calls: list[str] = []
transaction = MagicMock()
transaction.execute_raw = AsyncMock(side_effect=lambda sql: calls.append(sql) or 0)

View file

@ -245,6 +245,7 @@ def make_ui_spend_logs_mock_prisma(mock_spend_logs, filter_fn, team_lookup_fn=No
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
self.db.litellm_spendlogs = self.db
if team_lookup_fn is not None:
self.db.litellm_teamtable = self
@ -311,6 +312,7 @@ async def test_can_team_member_view_log_none_team_id():
def __init__(self):
self.db = self.DB()
self.replica_db = self.db
prisma = MockPrisma()
auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="user_1")
@ -334,6 +336,7 @@ async def test_can_team_member_view_log_team_not_found(monkeypatch):
def __init__(self):
self.db = self.DB()
self.replica_db = self.db
prisma = MockPrisma()
# Even if admin check would return True, no team means False
@ -375,6 +378,7 @@ async def test_can_team_member_view_log_not_admin(monkeypatch):
def __init__(self):
self.db = self.DB()
self.replica_db = self.db
prisma = MockPrisma()
monkeypatch.setattr(
@ -415,6 +419,7 @@ async def test_can_team_member_view_log_admin(monkeypatch):
def __init__(self):
self.db = self.DB()
self.replica_db = self.db
prisma = MockPrisma()
auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="user_1")
@ -476,6 +481,7 @@ def _make_owner_lookup_prisma(rows):
class MockPrisma:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
return MockPrisma()
@ -995,6 +1001,7 @@ async def test_ui_view_spend_logs_sort_by_and_sort_order(
class MockPrismaClient:
def __init__(self):
self.db = MagicMock()
self.replica_db = self.db
self.db.litellm_spendlogs = MagicMock()
self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count)
self.db.query_raw = AsyncMock(side_effect=mock_query_raw)
@ -1059,6 +1066,7 @@ async def test_ui_view_spend_logs_sort_validation_errors(
class MockPrismaClient:
def __init__(self):
self.db = MagicMock()
self.replica_db = self.db
self.db.litellm_spendlogs = MagicMock()
self.db.litellm_spendlogs.find_many = AsyncMock(return_value=[])
self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count)
@ -1137,6 +1145,7 @@ async def test_ui_view_spend_logs_sort_by_request_duration_ms(client, monkeypatc
class MockPrismaClient:
def __init__(self):
self.db = MagicMock()
self.replica_db = self.db
self.db.litellm_spendlogs = MagicMock()
self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count)
self.db.query_raw = AsyncMock(side_effect=mock_query_raw)
@ -1237,6 +1246,7 @@ async def test_ui_view_spend_logs_sort_by_model(
class MockPrismaClient:
def __init__(self):
self.db = MagicMock()
self.replica_db = self.db
self.db.litellm_spendlogs = MagicMock()
self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count)
self.db.query_raw = AsyncMock(side_effect=mock_query_raw)
@ -1352,6 +1362,7 @@ async def test_ui_view_spend_logs_sort_by_ttft_ms(client, monkeypatch):
class MockPrismaClient:
def __init__(self):
self.db = MagicMock()
self.replica_db = self.db
self.db.litellm_spendlogs = MagicMock()
self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count)
self.db.query_raw = AsyncMock(side_effect=mock_query_raw)
@ -2093,6 +2104,7 @@ async def test_ui_view_session_spend_logs_pagination(client, monkeypatch):
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
self.db.litellm_spendlogs = self.db
mock_prisma_client = MockPrismaClient()
@ -2145,6 +2157,7 @@ async def test_ui_view_session_spend_logs_rehydrates_metadata_jsonb_text(client,
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
self.db.litellm_spendlogs = self.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient())
@ -2186,6 +2199,7 @@ async def test_ui_view_session_spend_logs_scopes_non_admin_to_own_logs(client, m
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
self.db.litellm_spendlogs = self.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient())
@ -2248,6 +2262,7 @@ async def test_ui_view_session_spend_logs_includes_permitted_team_logs(client, m
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
self.db.litellm_spendlogs = self.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient())
@ -2736,6 +2751,7 @@ def _make_payload_lookup_prisma(rows):
class MockPrisma:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
return MockPrisma()
@ -2807,6 +2823,7 @@ async def test_ui_view_request_response_rejects_foreign_row_inserted_after_owner
class MockPrisma:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrisma())
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
@ -2838,6 +2855,7 @@ async def test_ui_view_request_response_custom_logger_denies_foreign_payload_own
class MockPrisma:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
class ColdStorageLogger:
async def get_request_response_payload(self, request_id, start_time_utc, end_time_utc):
@ -3963,10 +3981,12 @@ async def test_global_spend_keys_endpoint_limit_validation(client, monkeypatch):
"""
# Create a simple mock for prisma client with empty response
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_db = MagicMock()
mock_query_raw = AsyncMock(return_value=[])
mock_db.query_raw = mock_query_raw
mock_prisma_client.db = mock_db
mock_prisma_client.replica_db = mock_prisma_client.db
# Apply the mock to the prisma_client module
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
@ -4095,6 +4115,7 @@ async def test_view_spend_logs_summarize_parameter(client, monkeypatch):
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
# Apply the monkeypatch
mock_prisma_client = MockPrismaClient()
@ -4194,6 +4215,7 @@ async def test_view_spend_logs_bounds_row_count(client, monkeypatch):
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
def hash_token(self, token):
return f"hashed-{token}"
@ -4267,6 +4289,7 @@ async def test_view_spend_tags(client, monkeypatch):
# Mock the prisma client and get_spend_by_tags function
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
# Mock response data
@ -4447,6 +4470,7 @@ async def test_view_spend_logs_with_date_range_summarized(client, monkeypatch):
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient())
@ -4511,6 +4535,7 @@ async def test_view_spend_logs_summarize_groups_by_day_in_sql(client, monkeypatc
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
def hash_token(self, token):
return "hashed::" + token
@ -4578,6 +4603,7 @@ async def test_view_spend_logs_summarize_empty_rows(client, monkeypatch):
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient())
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
@ -4619,6 +4645,7 @@ async def test_view_spend_logs_summarize_unhashed_api_key_without_padding(client
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.replica_db = self.db
mock_prisma_client = MockPrismaClient()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
@ -4921,6 +4948,7 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts():
]
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_spendlogs.group_by = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(
return_value=[
@ -4987,6 +5015,7 @@ async def test_build_ui_spend_logs_response_caps_session_models():
over_limit_models = [f"model-{i:02d}" for i in range(_SESSION_MODELS_LIMIT + 1)]
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{
@ -5040,6 +5069,7 @@ async def test_build_ui_spend_logs_response_key_split_session_gets_per_key_aggre
]
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{
@ -5104,6 +5134,7 @@ async def test_build_ui_spend_logs_response_empty_api_key_keeps_session_aggregat
]
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{
@ -5161,6 +5192,7 @@ async def test_build_ui_spend_logs_response_sums_multi_round_session_spend():
]
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# The raw aggregate query returns the full session spend (0.01 + 0.02 + 0.03).
mock_prisma.db.query_raw = AsyncMock(
return_value=[
@ -5241,6 +5273,7 @@ async def test_build_ui_spend_logs_response_sums_multi_round_session_tokens():
]
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{
@ -5323,6 +5356,7 @@ async def test_build_ui_spend_logs_response_sums_multi_round_session_duration():
]
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{
@ -5382,6 +5416,7 @@ async def test_build_ui_spend_logs_response_session_cache_hit_count():
]
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{
@ -5451,6 +5486,7 @@ async def test_can_team_member_view_log_with_spend_logs_permission(monkeypatch):
def __init__(self):
self.db = self.DB()
self.replica_db = self.db
prisma = MockPrisma()
auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="member_1")
@ -5489,6 +5525,7 @@ async def test_can_team_member_view_log_without_spend_logs_permission(monkeypatc
def __init__(self):
self.db = self.DB()
self.replica_db = self.db
prisma = MockPrisma()
auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="member_1")
@ -5659,6 +5696,7 @@ class _CaptureFilterDB:
class _CapturePrismaClient:
def __init__(self):
self.db = _CaptureFilterDB()
self.replica_db = self.db
def hash_token(self, token):
return "hashed::" + token
@ -5839,6 +5877,7 @@ class _SpendScopeMockPrismaClient:
self.litellm_verificationtoken = _VerificationTokenTable()
self.db = _DB()
self.replica_db = self.db
async def get_data(self, table_name=None, query_type=None, **kwargs):
self.get_data_calls.append(
@ -6227,6 +6266,7 @@ async def test_ui_view_spend_logs_rehydrates_metadata_jsonb_text(client, monkeyp
class MockPrismaClient:
def __init__(self):
self.db = MagicMock()
self.replica_db = self.db
self.db.litellm_spendlogs = MagicMock()
self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count)
self.db.query_raw = AsyncMock(side_effect=mock_query_raw)
@ -6313,6 +6353,7 @@ async def test_ui_view_spend_logs_metadata_invalid_json_falls_back_to_empty_dict
class MockPrismaClient:
def __init__(self):
self.db = MagicMock()
self.replica_db = self.db
self.db.litellm_spendlogs = MagicMock()
self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count)
self.db.query_raw = AsyncMock(side_effect=mock_query_raw)
@ -6624,7 +6665,7 @@ def test_ui_view_request_response_reads_from_cold_storage(client, monkeypatch):
async def _query_raw(_sql, *_args):
return [placeholder_row]
fake_prisma = SimpleNamespace(db=SimpleNamespace(query_raw=_query_raw))
fake_prisma = SimpleNamespace(db=SimpleNamespace(query_raw=_query_raw), replica_db=SimpleNamespace(query_raw=_query_raw))
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", fake_prisma)
cold_logger = _FakeColdStorageLogger(
@ -6680,6 +6721,7 @@ _SCOPED_SPEND_REPORT_PATHS = [
def _spend_report_mock_prisma(query_raw_returns=None, team_rows=None, user_row=None):
pc = MagicMock()
pc.replica_db = pc.db
pc.db.query_raw = AsyncMock(
return_value=query_raw_returns if query_raw_returns is not None else []
)
@ -7237,6 +7279,7 @@ def _session_grouped_mock_prisma(session_page_rows, session_total, representativ
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(side_effect=mock_query_raw)
return mock_prisma
@ -7282,6 +7325,7 @@ def _session_grouped_paginating_prisma(sessions, counted_total=None):
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(side_effect=mock_query_raw)
return mock_prisma
@ -7658,6 +7702,7 @@ async def test_ui_view_spend_logs_search_returns_flat_rows_when_grouping_by_sess
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.query_raw = AsyncMock(side_effect=mock_query_raw)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
@ -7704,6 +7749,7 @@ def _fake_prisma_with_owned_spend_log(owner_user_id, messages_json, response_jso
class _Prisma:
def __init__(self):
self.db = _DB()
self.replica_db = self.db
return _Prisma()
@ -7810,7 +7856,7 @@ def test_ui_view_request_response_internal_user_missing_row_forbidden(client, mo
from types import SimpleNamespace
fake_prisma = SimpleNamespace(db=_DB())
fake_prisma = SimpleNamespace(db=_DB(), replica_db=_DB())
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", fake_prisma)
custom_logger = _RecordingAdditionalLoggingUtils(

View file

@ -106,6 +106,7 @@ def test_cors_exposes_cache_key_header_to_browser_js():
def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch):
mock_login_result = {"user_id": "test-user"}
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_authenticate_user = AsyncMock(return_value=mock_login_result)
mock_create_ui_token_object = MagicMock(return_value={"user_id": "test-user"})
mock_jwt_encode = MagicMock(return_value="signed-token")
@ -234,6 +235,7 @@ def test_login_v2_returns_json_on_proxy_exception(monkeypatch):
from litellm.proxy._types import ProxyErrorTypes, ProxyException
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_authenticate_user = AsyncMock(
side_effect=ProxyException(
message="Invalid credentials",
@ -269,6 +271,7 @@ def test_login_v2_returns_json_on_http_exception(monkeypatch):
from fastapi import HTTPException
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_authenticate_user = AsyncMock(side_effect=HTTPException(status_code=401, detail="Unauthorized"))
monkeypatch.setattr(
@ -294,6 +297,7 @@ def test_login_v2_returns_json_on_http_exception(monkeypatch):
def test_login_v2_returns_json_on_unexpected_exception(monkeypatch):
"""Test that /v2/login returns JSON error when unexpected exception occurs"""
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_authenticate_user = AsyncMock(side_effect=ValueError("Unexpected error"))
monkeypatch.setattr(
@ -338,6 +342,7 @@ def test_login_v2_returns_json_on_invalid_json_body(monkeypatch):
def test_login_v3_rejected_without_control_plane_url(monkeypatch):
"""v3/login returns 404 when control_plane_url is not configured."""
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
@ -355,6 +360,7 @@ def test_login_v3_rejected_without_control_plane_url(monkeypatch):
def test_login_v3_returns_code(monkeypatch):
"""v3/login returns an opaque code, not the JWT directly."""
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.authenticate_user",
AsyncMock(return_value={"user_id": "test-user"}),
@ -393,6 +399,7 @@ def test_login_v3_returns_code(monkeypatch):
def test_login_v3_exchange_happy_path(monkeypatch):
"""Full flow: v3/login returns code, v3/login/exchange redeems it for JWT."""
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.authenticate_user",
AsyncMock(return_value={"user_id": "test-user"}),
@ -441,6 +448,7 @@ def test_login_v3_exchange_sets_secure_cookie_behind_trusted_tls_terminating_pro
"""Regression: /v3/login/exchange's token cookie must be Secure behind a trusted
TLS-terminating reverse proxy even though litellm only sees a plain-HTTP hop."""
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.authenticate_user",
AsyncMock(return_value={"user_id": "test-user"}),
@ -485,6 +493,7 @@ def test_login_v3_exchange_sets_secure_cookie_behind_trusted_tls_terminating_pro
def test_login_v3_exchange_single_use(monkeypatch):
"""Code can only be redeemed once."""
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.authenticate_user",
AsyncMock(return_value={"user_id": "test-user"}),
@ -557,6 +566,7 @@ def test_login_v3_returns_json_on_proxy_exception(monkeypatch):
from litellm.proxy._types import ProxyErrorTypes, ProxyException
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_authenticate_user = AsyncMock(
side_effect=ProxyException(
message="Invalid credentials",
@ -847,6 +857,7 @@ async def test_initialize_scheduled_jobs_loads_credentials_only_through_add_depl
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.slack_alerting_instance = MagicMock()
mock_proxy_logging.db_spend_update_writer = MagicMock()
@ -906,6 +917,7 @@ async def test_periodic_reload_job_scheduled_without_store_model_in_db(monkeypat
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.slack_alerting_instance = MagicMock()
@ -948,6 +960,7 @@ async def test_initialize_scheduled_jobs_uses_configured_config_reload_interval(
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.slack_alerting_instance = MagicMock()
mock_proxy_logging.db_spend_update_writer = MagicMock()
@ -997,6 +1010,7 @@ async def test_initialize_scheduled_jobs_rejects_non_positive_config_reload_inte
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.slack_alerting_instance = MagicMock()
mock_proxy_logging.db_spend_update_writer = MagicMock()
@ -1044,6 +1058,7 @@ async def test_initialize_scheduled_jobs_hydrates_mcp_when_store_model_in_db_fal
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.slack_alerting_instance = MagicMock()
mock_proxy_logging.db_spend_update_writer = MagicMock()
@ -1907,10 +1922,12 @@ async def test_get_all_team_models():
# Mock prisma client
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_db = MagicMock()
mock_litellm_teamtable = MagicMock()
mock_prisma_client.db = mock_db
mock_prisma_client.replica_db = mock_prisma_client.db
mock_db.litellm_teamtable = mock_litellm_teamtable
# Make find_many async
@ -2214,6 +2231,7 @@ async def test_non_admin_all_models_returns_user_models_when_user_row_missing():
user_added_model = {"model_name": "my-model", "model_info": {"id": "user-model-1"}}
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
prisma_client.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=MagicMock(created_by="ghost-user"))
@ -2349,6 +2367,7 @@ async def test_apply_search_filter_scopes_byok_to_caller_teams():
}
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=2)
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[db_caller_row, db_other_row])
caller_user_row = MagicMock()
@ -2432,6 +2451,7 @@ async def test_apply_search_filter_bounds_db_fetch_by_page_and_cap():
)
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=10_000)
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
@ -2477,6 +2497,7 @@ async def test_apply_search_filter_honours_exact_model_name_in_db_query():
from litellm.proxy.proxy_server import _apply_search_filter_to_models
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=0)
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
proxy_config = MagicMock()
@ -2558,6 +2579,7 @@ async def test_filter_models_by_team_id_excludes_viewer_direct_access():
}
prisma = MagicMock()
prisma.replica_db = prisma.db
team_db = MagicMock()
team_db.model_dump.return_value = {
"team_id": "team-111",
@ -2607,6 +2629,7 @@ async def test_filter_models_by_team_id_rejects_non_member():
}
prisma = MagicMock()
prisma.replica_db = prisma.db
# Caller is in team-222 only
user_row = MagicMock()
user_row.teams = ["team-222"]
@ -2645,6 +2668,7 @@ async def test_filter_models_by_team_id_allows_team_member():
}
prisma = MagicMock()
prisma.replica_db = prisma.db
user_row = MagicMock()
user_row.teams = ["team-111", "team-999"]
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
@ -2767,6 +2791,7 @@ async def test_add_access_group_models_to_team_models():
mock_ag_row.access_model_names = ["claude-3", "gemini"]
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[mock_ag_row])
result = await _add_access_group_models_to_team_models(
@ -2843,6 +2868,7 @@ async def test_add_access_group_models_multiple_teams_shared_group():
mock_extra_row.access_model_names = ["gemini"]
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[mock_shared_row, mock_extra_row])
result = await _add_access_group_models_to_team_models(
@ -2883,6 +2909,7 @@ async def test_add_access_group_models_no_eligible_teams():
team.access_group_ids = None
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock()
result = await _add_access_group_models_to_team_models(
@ -2924,9 +2951,11 @@ async def test_get_all_team_models_with_access_groups():
mock_ag_row.access_model_names = ["claude-3"]
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_db = MagicMock()
mock_litellm_teamtable = MagicMock()
mock_prisma_client.db = mock_db
mock_prisma_client.replica_db = mock_prisma_client.db
mock_db.litellm_teamtable = mock_litellm_teamtable
mock_litellm_teamtable.find_many = AsyncMock(return_value=[mock_team1])
mock_db.litellm_accessgrouptable = MagicMock()
@ -3281,6 +3310,7 @@ async def test_add_proxy_budget_to_db_backfills_budget_reset_at(monkeypatch: pyt
litellm_proxy_budget_name = "litellm-proxy-budget"
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_usertable.update_many = AsyncMock(return_value={"count": 1})
mock_generate_key_helper = AsyncMock(
@ -4087,6 +4117,7 @@ async def test_write_config_to_file(monkeypatch):
# Mock prisma_client to not be None (so DB path is taken)
mock_prisma_client = AsyncMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.insert_data = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
@ -4456,6 +4487,7 @@ class TestPriceDataReloadAPI:
):
# Mock the database connection
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.upsert = AsyncMock(
return_value=_reload_schedule_row({}, reload_revision=1)
)
@ -4507,6 +4539,7 @@ class TestPriceDataReloadAPI:
):
# Mock the database connection
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_config.upsert = AsyncMock(
return_value=_reload_schedule_row({}, reload_revision=1)
@ -4521,6 +4554,7 @@ class TestPriceDataReloadAPI:
def test_schedule_model_cost_map_reload_admin_access(self, client_with_auth):
"""Admin schedule write owns param_value only, so it can't clobber the job-owned run columns"""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
# Mock database upsert
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1))
@ -4567,6 +4601,7 @@ class TestPriceDataReloadAPI:
def test_cancel_model_cost_map_reload_admin_access(self, client_with_auth):
"""Test that admin users can cancel periodic reload"""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=1)
mock_prisma.db.litellm_config.delete = AsyncMock(return_value=None)
@ -4604,6 +4639,7 @@ class TestPriceDataReloadAPI:
proxy_server_module.proxy_config.model_cost_map_loaded_at = datetime(2030, 6, 1, tzinfo=timezone.utc)
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.find_unique = AsyncMock(
return_value=_reload_schedule_row(
{"interval_hours": 6},
@ -4637,6 +4673,7 @@ class TestPriceDataReloadAPI:
def test_get_model_cost_map_reload_status_no_config(self, client_with_auth):
"""Test that status returns not scheduled when no config exists"""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
response = client_with_auth.get("/schedule/model_cost_map_reload/status")
@ -4651,6 +4688,7 @@ class TestPriceDataReloadAPI:
def test_get_model_cost_map_reload_status_no_interval(self, client_with_auth):
"""A row left behind by a manual reload (no interval) must not read as scheduled"""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.find_unique = AsyncMock(
return_value=_reload_schedule_row(
{"interval_hours": None},
@ -4670,6 +4708,7 @@ class TestPriceDataReloadAPI:
def test_get_model_cost_map_reload_status_before_first_run(self, client_with_auth):
"""Scheduled but never executed: no last_run_at means no next_run can be computed"""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.find_unique = AsyncMock(
return_value=_reload_schedule_row({"interval_hours": 6})
)
@ -4726,6 +4765,7 @@ class TestPriceDataReloadIntegration:
):
# Mock the database connection
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.upsert = AsyncMock(
return_value=_reload_schedule_row({}, reload_revision=1)
)
@ -4764,6 +4804,7 @@ class TestPriceDataReloadIntegration:
proxy_config = ProxyConfig()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1))
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
@ -4822,6 +4863,7 @@ class TestPriceDataReloadIntegration:
proxy_config = ProxyConfig()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
mock_prisma.db.litellm_config.find_unique = AsyncMock(
return_value=_reload_schedule_row(
@ -4862,6 +4904,7 @@ class TestPriceDataReloadIntegration:
proxy_config = ProxyConfig()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.find_unique = AsyncMock(
return_value=_reload_schedule_row(
{"interval_hours": 6},
@ -4911,6 +4954,7 @@ class TestPriceDataReloadIntegration:
pods = [ProxyConfig(), ProxyConfig(), ProxyConfig()]
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1))
for pod in pods:
@ -4956,6 +5000,7 @@ class TestPriceDataReloadIntegration:
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
proxy_config.model_cost_map_loaded_at = frozen_now
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
mock_prisma.db.litellm_config.find_unique = AsyncMock(
return_value=_reload_schedule_row({}, reload_revision=published_revision)
@ -4992,6 +5037,7 @@ class TestPriceDataReloadIntegration:
proxy_config = ProxyConfig()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=_reload_schedule_row({"interval_hours": 24}))
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1))
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
@ -5033,6 +5079,7 @@ class TestPriceDataReloadIntegration:
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
proxy_config.model_cost_map_loaded_at = frozen_now - timedelta(hours=9)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.find_unique = AsyncMock(
return_value=_reload_schedule_row({"interval_hours": 6}, reload_revision=7)
)
@ -5078,6 +5125,7 @@ class TestPriceDataReloadIntegration:
pod_data_loaded_at = frozen_now - timedelta(hours=9)
proxy_config.model_cost_map_loaded_at = pod_data_loaded_at
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.find_unique = AsyncMock(
return_value=_reload_schedule_row({"interval_hours": 6}, reload_revision=7)
)
@ -5120,6 +5168,7 @@ class TestPriceDataReloadIntegration:
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
proxy_config.model_cost_map_loaded_at = frozen_now - timedelta(hours=9)
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_config.find_unique = AsyncMock(
return_value=_reload_schedule_row({"interval_hours": 6}, reload_revision=7)
)
@ -5211,6 +5260,7 @@ class TestPriceDataReloadIntegration:
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
):
mock_prisma.replica_db = mock_prisma.db
mock_get_map.return_value = ModelCostMapReloaded(
model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}}
)
@ -5252,6 +5302,7 @@ class TestPriceDataReloadIntegration:
litellm_config_cache.flush_cache()
proxy_config = ProxyConfig()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
# Set up config with interval_hours=12 and force_reload=True to trigger reload
mock_config = MagicMock()
@ -5300,6 +5351,7 @@ class TestPriceDataReloadIntegration:
mock_reload.return_value = {"anthropic": {"beta_header": "test-value"}}
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
# Simulate existing config with a schedule
mock_existing = MagicMock()
mock_existing.param_value = {"interval_hours": 8, "force_reload": False}
@ -5394,6 +5446,7 @@ async def test_add_router_settings_from_db_config_merge_logic():
# Mock prisma client
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
# Call the method under test
@ -5450,6 +5503,7 @@ async def test_invalid_db_routing_groups_do_not_abort_other_router_settings():
],
}
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
await ProxyConfig()._add_router_settings_from_db_config(llm_router=router, prisma_client=mock_prisma_client)
@ -5472,6 +5526,7 @@ async def test_valid_db_routing_groups_still_replace_router_groups():
"routing_groups": [{"group_name": "g2", "models": ["m2"], "routing_strategy": "least-busy"}],
}
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
await ProxyConfig()._add_router_settings_from_db_config(llm_router=router, prisma_client=mock_prisma_client)
@ -5515,6 +5570,7 @@ async def test_add_router_settings_from_db_config_empty_db_lists_do_not_clobber_
}
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
proxy_config.router_settings.load_yaml(config_data["router_settings"])
@ -5552,6 +5608,7 @@ async def test_add_router_settings_from_db_config_empty_db_list_still_clears_unc
mock_db_config.param_value = {"fallbacks": [], "model_group_alias": {}}
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
proxy_config.router_settings.load_yaml(config_data["router_settings"])
@ -5598,6 +5655,7 @@ async def test_add_router_settings_from_db_config_edge_cases():
# Test Case 3: DB returns None (no router_settings in DB)
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
config_data = {"router_settings": {"routing_strategy": "usage-based"}}
@ -5691,6 +5749,7 @@ async def test_add_router_settings_shallow_merge_behavior():
}
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
proxy_config.router_settings.load_yaml(config_data["router_settings"])
@ -5728,6 +5787,7 @@ async def test_router_settings_reload_keeps_db_values_writable(tmp_path, monkeyp
return db_row if param_name == "router_settings" else None
mock_prisma_client: Final = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=db_row)
mock_router: Final = MagicMock()
monkeypatch.setattr(proxy_server_module, "get_config_param", read_config_row)
@ -6353,6 +6413,7 @@ async def test_init_sso_settings_in_db():
}
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config)
# Mock _decrypt_and_set_db_env_variables
@ -6397,6 +6458,7 @@ async def test_init_sso_settings_in_db_no_settings():
# Mock prisma client to return None (no SSO settings)
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None)
# Mock _decrypt_and_set_db_env_variables
@ -6423,6 +6485,7 @@ async def test_init_sso_settings_in_db_error_handling():
# Mock prisma client to raise an exception
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=Exception("Database connection error"))
# The method should not raise an exception, it should log it instead
@ -6451,6 +6514,7 @@ async def test_init_sso_settings_in_db_empty_settings():
mock_sso_config.sso_settings = {}
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config)
# Mock _decrypt_and_set_db_env_variables
@ -6491,6 +6555,7 @@ async def test_init_sso_settings_in_db_retries_on_transport_error():
return mock_sso_config
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=_flaky_find_unique)
mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True)
mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0
@ -6517,6 +6582,7 @@ async def test_init_sso_settings_in_db_propagates_when_reconnect_fails():
proxy_config = ProxyConfig()
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=prisma.errors.ClientNotConnectedError())
mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=False)
mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0
@ -6548,6 +6614,7 @@ async def test_init_hashicorp_vault_config_override_retries_on_transport_error()
return None # No config in DB → function returns early after retry.
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_configoverrides.find_unique = AsyncMock(side_effect=_flaky_find_unique)
mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True)
mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0
@ -7131,6 +7198,7 @@ class TestInvitationEndpoints:
def test_invitation_endpoints_proxy_admin_success(self, client_with_auth, endpoint, payload, mock_return):
"""Proxy admin can successfully create and delete invitations."""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_invitationlink = MagicMock()
if endpoint == "/invitation/new":
mock_create = AsyncMock(return_value=mock_return)
@ -7169,6 +7237,7 @@ class TestInvitationEndpoints:
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_invitationlink = MagicMock()
# Avoid triggering async DB calls in _user_has_admin_privileges
with patch(
@ -8344,6 +8413,7 @@ async def test_batch_cost_poller_is_confirmed_before_serving(monkeypatch):
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=None)
mock_proxy_logging = MagicMock(spec=ProxyLogging)
@ -8384,6 +8454,7 @@ async def test_store_model_in_db_db_override_when_config_false():
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Mock DB returning store_model_in_db=True in general_settings
mock_db_record = MagicMock()
@ -8429,6 +8500,7 @@ async def test_store_model_in_db_db_check_skipped_when_already_true(monkeypatch)
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
mock_proxy_logging = MagicMock(spec=ProxyLogging)
@ -8471,6 +8543,7 @@ async def test_store_model_in_db_db_failure_graceful(monkeypatch):
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
# Simulate DB failure
mock_prisma_client.db.litellm_config.find_first = AsyncMock(side_effect=Exception("DB connection error"))
@ -8710,6 +8783,7 @@ async def test_prepare_spend_counter_increment_reseeds_from_db_on_counter_miss()
db_row = MagicMock()
db_row.spend = 42.0
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row)
stale_cache = DualCache()
@ -8811,6 +8885,7 @@ async def test_primary_spend_counter_redis_concurrent_seed_does_not_double_seed(
return row
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(side_effect=slow_find_unique)
pod_a = DualCache()
@ -8872,6 +8947,7 @@ async def test_reseed_spend_from_db_user_and_org_prefixes():
org_row.spend = 305.0
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
fake_prisma.db.litellm_endusertable.find_unique = AsyncMock()
fake_prisma.db.litellm_tagtable.find_unique = AsyncMock()
@ -8904,6 +8980,7 @@ async def test_reseed_spend_from_db_skips_window_variant_keys():
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_verificationtoken.find_unique = AsyncMock()
fake_prisma.db.litellm_teamtable.find_unique = AsyncMock()
@ -8924,6 +9001,7 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
counter_cache = DualCache()
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_budgetwindowspend.find_unique = AsyncMock(return_value=None)
fake_prisma.db.litellm_spendlogs.group_by = AsyncMock(
return_value=[{"api_key": "key-window", "_sum": {"spend": 2.25}}]
@ -8998,6 +9076,7 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
db_row = MagicMock()
db_row.spend = 42.0
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row)
import litellm.proxy.proxy_server as ps
@ -9069,6 +9148,7 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_budgetwindowspend.find_unique = AsyncMock(return_value=None)
fake_prisma.db.litellm_spendlogs.group_by = AsyncMock(
return_value=[{"api_key": "key-window-stale-local", "_sum": {"spend": 2.25}}]
@ -9146,6 +9226,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_budgetwindowspend.find_unique = AsyncMock(return_value=None)
fake_prisma.db.litellm_spendlogs.group_by = AsyncMock(
return_value=[{"api_key": "key-window-concurrent-seed", "_sum": {"spend": 2.25}}]
@ -9444,6 +9525,7 @@ async def test_get_current_spend_reseeds_from_db_when_counter_missing():
db_row = MagicMock()
db_row.spend = 362.0
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=db_row)
import litellm.proxy.proxy_server as ps
@ -9548,6 +9630,7 @@ async def test_get_current_spend_coalesces_concurrent_reseeds():
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=slow_find_unique)
import litellm.proxy.proxy_server as ps
@ -9585,6 +9668,7 @@ async def test_get_current_spend_uses_db_zero_over_stale_fallback():
db_row = MagicMock()
db_row.spend = 0.0
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=db_row)
import litellm.proxy.proxy_server as ps
@ -9655,6 +9739,7 @@ async def test_concurrent_read_and_write_paths_share_one_db_query():
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=slow_find_unique)
import litellm.proxy.proxy_server as ps
@ -9770,6 +9855,7 @@ async def test_reseed_warms_cache_even_on_zero_db_spend():
return row
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=find_unique)
import litellm.proxy.proxy_server as ps
@ -9833,6 +9919,7 @@ class _FakeLitellmConfig:
class _FakePrismaClient:
def __init__(self, initial_rows=None):
self.db = mock.MagicMock()
self.replica_db = self.db
self.db.litellm_config = _FakeLitellmConfig(initial_rows=initial_rows)
self.jsonify_object = lambda obj: obj
@ -10698,6 +10785,7 @@ async def test_get_current_spend_redis_clean_miss_skips_stale_in_memory():
db_row = MagicMock()
db_row.spend = 500.0
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=db_row)
import litellm.proxy.proxy_server as ps
@ -10734,6 +10822,7 @@ async def test_get_current_spend_redis_error_falls_back_to_in_memory():
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
fake_prisma.replica_db = fake_prisma.db
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=MagicMock(spend=999.0))
import litellm.proxy.proxy_server as ps
@ -11292,6 +11381,7 @@ class TestDeleteDeploymentSync:
proxy_config = ProxyConfig()
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(side_effect=Exception("DB connection lost"))
result = await proxy_config._get_models_from_db(prisma_client=mock_prisma)
@ -11383,9 +11473,11 @@ def test_get_config_list_includes_cancel_on_disconnect(monkeypatch):
from litellm.proxy.proxy_server import app
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_config_table = MagicMock()
mock_config_table.find_first = AsyncMock(return_value=None)
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
mock_prisma.replica_db = mock_prisma.db
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
@ -11415,9 +11507,11 @@ def test_get_config_list_includes_apply_user_budget_to_team_keys(monkeypatch):
from litellm.proxy.proxy_server import app
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_config_table = MagicMock()
mock_config_table.find_first = AsyncMock(return_value=None)
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
mock_prisma.replica_db = mock_prisma.db
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
@ -11437,9 +11531,11 @@ def test_get_config_list_includes_user_api_key_cache_max_size(monkeypatch):
"""The Admin UI General Settings table renders whatever /config/list returns,
so the cache capacity has to be exposed there as an Integer to be editable."""
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_config_table = MagicMock()
mock_config_table.find_first = AsyncMock(return_value=None)
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
mock_prisma.replica_db = mock_prisma.db
monkeypatch.setattr(proxy_server_module, "prisma_client", mock_prisma)
app.dependency_overrides[proxy_server_module.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
@ -11468,9 +11564,11 @@ def test_get_config_list_includes_budget_exceeded_throttle_percentage(monkeypatc
from litellm.proxy.proxy_server import app
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_config_table = MagicMock()
mock_config_table.find_first = AsyncMock(return_value=None)
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
mock_prisma.replica_db = mock_prisma.db
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", 0.15)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
@ -11545,9 +11643,11 @@ def test_get_config_list_includes_anthropic_prompt_caching_fields(monkeypatch):
from litellm.proxy.proxy_server import app
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_config_table = MagicMock()
mock_config_table.find_first = AsyncMock(return_value=None)
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
mock_prisma.replica_db = mock_prisma.db
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
monkeypatch.setattr(litellm, "anthropic_prompt_caching_ttl", "1h")
@ -11853,9 +11953,11 @@ def test_get_config_list_marks_untouched_prompt_caching_flag_as_not_set(monkeypa
from litellm.proxy.proxy_server import app
mock_prisma = MagicMock()
mock_prisma.replica_db = mock_prisma.db
mock_config_table = MagicMock()
mock_config_table.find_first = AsyncMock(return_value=None)
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
mock_prisma.replica_db = mock_prisma.db
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", False)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
@ -12131,6 +12233,7 @@ def _config_field_info_client(monkeypatch, user_role):
mock_config_table.find_first = AsyncMock(return_value=db_record)
mock_prisma = MagicMock()
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
mock_prisma.replica_db = mock_prisma.db
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
settings = SettingsStore("general_settings")
@ -12188,6 +12291,7 @@ def _fake_prisma_with_config(existing_param_value):
"""MagicMock prisma whose litellm_config row returns existing_param_value and
whose litellm_auditlog.create records the written audit row."""
fake = MagicMock()
fake.replica_db = fake.db
config_row = MagicMock()
config_row.param_value = existing_param_value
fake.db.litellm_config.find_first = AsyncMock(return_value=config_row)
@ -13709,6 +13813,7 @@ async def test_team_window_spend_row_carries_the_request_start_time():
def _mock_startup_prisma_client(health_check_error=None, connect_error=None):
client = MagicMock()
client.replica_db = client.db
client.connect = AsyncMock(side_effect=connect_error)
client.db.start_token_refresh_task = AsyncMock()
client.check_view_exists = AsyncMock()
@ -13849,6 +13954,7 @@ async def _run_scheduled_background_jobs():
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_prisma_client.replica_db = mock_prisma_client.db
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
mock_proxy_logging = MagicMock(spec=ProxyLogging)
@ -14242,6 +14348,7 @@ async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeyp
return served_callback().prompt_manager.get_prompt("greeting_sync").content
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
try:
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[db_row("Begin every reply with AHOY")])
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
@ -14286,6 +14393,7 @@ async def test_init_prompts_in_db_syncs_remaining_rows_when_one_row_fails(monkey
return row
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
try:
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
return_value=[db_row("broken_sync", "does_not_exist"), db_row("healthy_sync", "dotprompt")]
@ -14338,6 +14446,7 @@ async def test_init_prompts_in_db_syncs_every_environment_sharing_a_versioned_id
return callback.prompt_manager.get_prompt("greeting_env").content
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
try:
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
return_value=[
@ -14398,6 +14507,7 @@ async def test_init_prompts_in_db_unloads_rows_deleted_on_another_worker(monkeyp
monkeypatch.setattr(litellm, "callbacks", [])
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
try:
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
return_value=[_prompt_db_row("greeting_del", _dotprompt_params("greeting_del"))]
@ -14435,6 +14545,7 @@ async def test_init_prompts_in_db_keeps_config_prompts_when_their_id_has_no_db_r
)
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
try:
IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(prompt=config_prompt)
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[])
@ -14457,6 +14568,7 @@ async def test_init_prompts_in_db_keeps_the_in_memory_copy_when_a_row_fails_to_p
monkeypatch.setattr(litellm, "callbacks", [])
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
try:
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
return_value=[_prompt_db_row("greeting_broken", _dotprompt_params("greeting_broken"))]
@ -14489,6 +14601,7 @@ async def test_init_prompts_in_db_keeps_a_prompt_created_while_the_sync_was_read
monkeypatch.setattr(litellm, "callbacks", [])
prisma_client = MagicMock()
prisma_client.replica_db = prisma_client.db
try:
async def create_prompt_behind_the_select() -> list:

View file

@ -2433,3 +2433,107 @@ def test_handle_exception_on_proxy_logs_bug_report_only_for_unmapped_500(caplog)
assert provider_result.code == internal_result.code == "500"
assert ISSUE_URL_BASE in caplog.text
assert ISSUE_URL_BASE not in internal_result.message
class _ReadOnlyReplicaError(RuntimeError):
pass
class _FakeUserTable:
def __init__(self, rows: list[dict[str, str]], *, read_only: bool) -> None:
self._rows = rows
self._read_only = read_only
async def find_many(self) -> list[dict[str, str]]:
return list(self._rows)
async def find_unique(self, where: dict[str, str]) -> dict[str, str] | None:
return next((row for row in self._rows if row["user_id"] == where["user_id"]), None)
async def create(self, data: dict[str, str]) -> dict[str, str]:
if self._read_only:
raise _ReadOnlyReplicaError("cannot execute INSERT in a read-only transaction")
self._rows.append(data) # mutable-ok: in-memory fake table
return data
class _FakePrisma:
def __init__(self, label: str, rows: list[dict[str, str]], *, read_only: bool, reachable: bool = True) -> None:
self._label = label
self._reachable = reachable
self._connected = False
self.litellm_usertable = _FakeUserTable(rows, read_only=read_only)
def is_connected(self) -> bool:
return self._connected
async def connect(self, timeout: object = None) -> None:
if not self._reachable:
raise ConnectionError(f"{self._label} unreachable")
self._connected = True
async def query_raw(self, sql: str, *params: object) -> list[dict[str, str]]:
return [{"served_by": self._label}]
def _replica_client(*, reader_reachable: bool = True) -> tuple[PrismaClient, list[dict[str, str]], list[dict[str, str]]]:
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer_rows: Final[list[dict[str, str]]] = [{"user_id": "u1", "user_email": "writer@example.com"}]
reader_rows: Final[list[dict[str, str]]] = [{"user_id": "u1", "user_email": "reader@example.com"}]
writer: Final = PrismaWrapper(original_prisma=_FakePrisma("writer", writer_rows, read_only=False))
reader: Final = PrismaWrapper(
original_prisma=_FakePrisma("reader", reader_rows, read_only=True, reachable=reader_reachable)
)
client: Final = PrismaClient.__new__(PrismaClient)
client.db = RoutingPrismaWrapper(writer=writer, reader=reader)
return client, writer_rows, reader_rows
class TestPrismaClientReplicaDb:
@pytest.mark.asyncio
async def test_replica_db_reads_come_from_the_reader(self) -> None:
client, _writer_rows, reader_rows = _replica_client()
await client.db.connect()
found: Final = await client.replica_db.litellm_usertable.find_unique(where={"user_id": "u1"})
raw: Final = await client.replica_db.query_raw("SELECT 1")
assert found == reader_rows[0]
assert raw == [{"served_by": "reader"}]
@pytest.mark.asyncio
async def test_replica_db_writes_land_on_the_writer(self) -> None:
client, writer_rows, reader_rows = _replica_client()
await client.db.connect()
new_row: Final = {"user_id": "u2", "user_email": "new@example.com"}
created: Final = await client.replica_db.litellm_usertable.create(data=new_row)
assert created == new_row
assert new_row in writer_rows
assert new_row not in reader_rows
@pytest.mark.asyncio
async def test_replica_db_reads_fall_back_to_the_writer_when_the_reader_is_unreachable(self) -> None:
client, writer_rows, _reader_rows = _replica_client(reader_reachable=False)
await client.db.connect()
found: Final = await client.replica_db.litellm_usertable.find_many()
raw: Final = await client.replica_db.query_raw("SELECT 1")
assert found == writer_rows
assert raw == [{"served_by": "writer"}]
@pytest.mark.asyncio
async def test_replica_db_is_the_plain_handle_without_a_replica(self) -> None:
from litellm.proxy.db.prisma_client import PrismaWrapper
rows: Final[list[dict[str, str]]] = [{"user_id": "u1", "user_email": "only@example.com"}]
client: Final = PrismaClient.__new__(PrismaClient)
client.db = PrismaWrapper(original_prisma=_FakePrisma("writer", rows, read_only=False))
assert client.replica_db is client.db
assert await client.replica_db.litellm_usertable.find_many() == rows
assert await client.replica_db.query_raw("SELECT 1") == [{"served_by": "writer"}]

View file

@ -466,6 +466,7 @@ class _FakePrismaDB:
class _FakePrismaClient:
def __init__(self, results):
self.db = _FakePrismaDB(results)
self.replica_db = self.db
def _spend_log(request_id: str, session_id: str, prompt: str, answer: str) -> dict: