mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
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:
parent
6dbd65b230
commit
f6f3920337
82 changed files with 1050 additions and 131 deletions
|
|
@ -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}
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
########################################################
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)},
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]]] = {}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]":
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]] = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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]):
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]:
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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"]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
||||
|
|
|
|||
|
|
@ -68,6 +68,7 @@ class _FakeDB:
|
|||
class _FakePrismaClient:
|
||||
def __init__(self, db: _FakeDB) -> None:
|
||||
self.db = db
|
||||
self.replica_db = self.db
|
||||
|
||||
|
||||
class _RecordingAggregate:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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: ()))
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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()
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"}]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue