diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index a94a75fdfa3..9aaedfc8d42 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -18,7 +18,7 @@ import os import re import secrets import traceback -from collections.abc import Mapping +from collections.abc import Awaitable, Mapping from datetime import datetime, timedelta, timezone from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, cast @@ -466,6 +466,9 @@ def common_key_access_checks( router = APIRouter() +CustomKeyGenerateHook = Callable[[GenerateKeyRequest], Awaitable[Mapping[str, object]]] +CustomKeyUpdateHook = Callable[[UpdateKeyRequest], Awaitable[Mapping[str, object]]] + def handle_key_type(data: GenerateKeyRequest, data_json: dict) -> dict: """ @@ -490,7 +493,7 @@ _NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS = frozenset({"llm_api_routes", "info_rout def _validate_caller_can_change_key_ownership( data: Optional[BaseModel], - existing_key_row: Any, + existing_key_row: LiteLLM_VerificationToken, user_api_key_dict: UserAPIKeyAuth, ) -> None: """ @@ -670,7 +673,7 @@ async def validate_team_id_used_in_service_account_request( ) # check if team_id exists in the database - team = await TeamRepository(prisma_client).table.find_unique( + team = await TeamRepository(prisma_client).table.find_unique( # any-ok: existence check only, avoid model_validate where={"team_id": team_id}, ) if team is None: @@ -889,7 +892,7 @@ async def _common_key_generation_helper( ) new_budget = prisma_client.jsonify_object(budget_row.json(exclude_none=True)) - _budget = await BudgetRepository(prisma_client).table.create( + _budget = await BudgetRepository(prisma_client).table.create( # any-ok: avoid model_validate on mocked rows data={ **new_budget, # type: ignore "created_by": user_api_key_dict.user_id or litellm_proxy_admin_name, @@ -1261,7 +1264,7 @@ async def _check_team_key_limits( # calculate allocated tpm/rpm limit # check if specified tpm/rpm limit is greater than allocated tpm/rpm limit - keys = await VerificationTokenRepository(prisma_client).table.find_many( + keys = await VerificationTokenRepository(prisma_client).table.find_many( # any-ok: avoid model_validate on mocks where={"team_id": team_table.team_id}, ) # Exclude the key being updated to avoid double-counting its limits. @@ -1408,7 +1411,7 @@ async def _validate_caller_can_assign_key_org( user_row = await UserRepository(prisma_client).table.find_unique( where={"user_id": user_api_key_dict.user_id}, include={"organization_memberships": True}, - ) + ) # any-ok: raw prisma row needed for the organization_memberships include; no repository helper supports it memberships = getattr(user_row, "organization_memberships", None) if user_row else None member_org_ids = { membership.organization_id for membership in (memberships or []) if membership.organization_id is not None @@ -1443,7 +1446,7 @@ async def _check_org_key_limits( # get all organization keys # calculate allocated tpm/rpm limit # check if specified tpm/rpm limit is greater than allocated tpm/rpm limit - keys = await VerificationTokenRepository(prisma_client).table.find_many( + keys = await VerificationTokenRepository(prisma_client).table.find_many( # any-ok: avoid model_validate on mocks where={"organization_id": org_table.organization_id}, ) # Exclude the key being updated to avoid double-counting its limits. @@ -1562,6 +1565,8 @@ async def generate_key_fn( user_custom_key_generate, ) + custom_key_generate_hook: CustomKeyGenerateHook | None = user_custom_key_generate + if prisma_client is None: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, @@ -1585,9 +1590,9 @@ async def generate_key_fn( detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"}, ) - if user_custom_key_generate is not None: - if inspect.iscoroutinefunction(user_custom_key_generate): - result = await user_custom_key_generate(data) # type: ignore + if custom_key_generate_hook is not None: + if inspect.iscoroutinefunction(custom_key_generate_hook): + result = await custom_key_generate_hook(data) else: raise ValueError("user_custom_key_generate must be a coroutine") decision = result.get("decision", True) @@ -1759,6 +1764,8 @@ async def generate_service_account_key_fn( user_custom_key_generate, ) + custom_key_generate_hook: CustomKeyGenerateHook | None = user_custom_key_generate + if prisma_client is None: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, @@ -1784,9 +1791,9 @@ async def generate_service_account_key_fn( verbose_proxy_logger.debug("entered /key/generate") - if user_custom_key_generate is not None: - if inspect.iscoroutinefunction(user_custom_key_generate): - result = await user_custom_key_generate(data) # type: ignore + if custom_key_generate_hook is not None: + if inspect.iscoroutinefunction(custom_key_generate_hook): + result = await custom_key_generate_hook(data) else: raise ValueError("user_custom_key_generate must be a coroutine") decision = result.get("decision", True) @@ -2050,7 +2057,9 @@ async def _get_and_validate_existing_key( existing_key_row: LiteLLM_VerificationToken | None = await VerificationTokenRepository( prisma_client - ).table.find_unique(where={"token": hashed_token}) + ).table.find_unique( # any-ok: avoid model_validate on mocked rows + where={"token": hashed_token} + ) if existing_key_row is None: raise ProxyException( @@ -2070,7 +2079,9 @@ async def _get_and_validate_existing_key( code=status.HTTP_400_BAD_REQUEST, ) - rows: list[LiteLLM_VerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many( + rows: list[LiteLLM_VerificationToken] = await VerificationTokenRepository( + prisma_client + ).table.find_many( # any-ok: avoid model_validate on mocked rows where={"key_alias": key_alias}, take=2 ) @@ -2112,11 +2123,11 @@ async def _process_single_key_update( litellm_changed_by: Optional[str], prisma_client: Optional[PrismaClient], user_api_key_cache: UserApiKeyCache, - proxy_logging_obj: Any, + proxy_logging_obj: ProxyLogging, llm_router: Optional[Router], - user_custom_key_update: Optional[Callable] = None, + user_custom_key_update: CustomKeyUpdateHook | None = None, existing_key_row: Optional[LiteLLM_VerificationToken] = None, -) -> Dict[str, Any]: +) -> Dict[str, Any]: # any-ok: consumers SuccessfulKeyUpdate/FailedKeyUpdate.key_info are Dict[str, Any] """ Process a single key update with all validations and checks. @@ -2265,9 +2276,9 @@ async def _process_single_key_update( async def _validate_mcp_servers_for_key_update( data: "UpdateKeyRequest", team_obj: Optional["LiteLLM_TeamTableCachedObj"], - existing_key_row: Any, - prisma_client: Any, - user_api_key_cache: Any, + existing_key_row: LiteLLM_VerificationToken, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, is_proxy_admin: bool, ) -> Optional[ObjectPermissionDict]: """Validate MCP servers in object_permission against the effective team.""" @@ -2302,14 +2313,17 @@ async def _validate_mcp_servers_for_key_update( async def _validate_update_key_data( data: UpdateKeyRequest, - existing_key_row: Any, + existing_key_row: LiteLLM_VerificationToken, user_api_key_dict: UserAPIKeyAuth, - llm_router: Any, + llm_router: Router | None, premium_user: bool, - prisma_client: Any, - user_api_key_cache: Any, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, ) -> None: """Validate permissions and constraints for key update.""" + if prisma_client is None: + raise HTTPException(status_code=500, detail={"error": "Database not connected"}) + # Reject NaN/±inf spend before it can reach the DB / spend counter. validate_finite_spend(data.spend) @@ -2414,8 +2428,10 @@ async def _validate_update_key_data( # _check_key_admin_access that would otherwise require team/org admin status. _key_is_team_key = getattr(existing_key_row, "team_id", None) is not None can_skip_admin_check = (caller_is_creator or _key_is_team_key) and not _is_budget_change - if (not _is_proxy_admin) and prisma_client is not None and not can_skip_admin_check: + if (not _is_proxy_admin) and not can_skip_admin_check: hashed_key = existing_key_row.token + if hashed_key is None: + raise HTTPException(status_code=500, detail={"error": "Key token missing on existing key row"}) await _check_key_admin_access( user_api_key_dict=user_api_key_dict, hashed_token=hashed_key, @@ -2631,6 +2647,8 @@ async def update_key_fn( user_custom_key_update, ) + custom_key_update_hook: CustomKeyUpdateHook | None = user_custom_key_update + try: # Validate budget values are not negative and are finite numbers if data.max_budget is not None and (not math.isfinite(data.max_budget) or data.max_budget < 0): @@ -2659,9 +2677,9 @@ async def update_key_fn( ) # Custom key update hook - if user_custom_key_update is not None: - if inspect.iscoroutinefunction(user_custom_key_update): - result = await user_custom_key_update(data) + if custom_key_update_hook is not None: + if inspect.iscoroutinefunction(custom_key_update_hook): + result = await custom_key_update_hook(data) else: raise ValueError("user_custom_key_update must be a coroutine") decision = result.get("decision", True) @@ -2821,6 +2839,8 @@ async def bulk_update_keys( user_custom_key_update, ) + custom_key_update_hook: CustomKeyUpdateHook | None = user_custom_key_update + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: raise HTTPException( status_code=403, @@ -2866,7 +2886,7 @@ async def bulk_update_keys( user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, llm_router=llm_router, - user_custom_key_update=user_custom_key_update, + user_custom_key_update=custom_key_update_hook, ) successful_updates.append( @@ -2939,7 +2959,7 @@ def _build_failed_team_key_update( else: error_message = str(exception) - key_info: Optional[Dict[str, Any]] = None + key_info: Dict[str, Any] | None = None # any-ok: FailedKeyUpdate.key_info is Dict[str, Any] if existing_key_row is not None: if hasattr(existing_key_row, "model_dump"): key_info = existing_key_row.model_dump() @@ -2983,6 +3003,8 @@ async def bulk_update_team_keys( user_custom_key_update, ) + custom_key_update_hook: CustomKeyUpdateHook | None = user_custom_key_update + if prisma_client is None: raise HTTPException( status_code=500, @@ -3009,7 +3031,9 @@ async def bulk_update_team_keys( # `blocked` is Boolean? with no default; `/key/generate` writes NULL. Prisma's `NOT` # excludes NULLs, so explicitly OR `false` with `null` to include them. now = datetime.now(timezone.utc) - existing_keys = await VerificationTokenRepository(prisma_client).table.find_many( + existing_keys = await VerificationTokenRepository( + prisma_client + ).table.find_many( # any-ok: avoid model_validate on mocked rows where={ "team_id": data.team_id, "AND": [ @@ -3027,7 +3051,7 @@ async def bulk_update_team_keys( "error": f"Team {data.team_id} has more than {MAX_BATCH_SIZE} keys. Use `key_ids` to update in batches of {MAX_BATCH_SIZE}." }, ) - requested_tokens = [row.token for row in existing_keys] + requested_tokens = [row.token for row in existing_keys if row.token is not None] else: if data.key_ids is None or len(data.key_ids) == 0: raise HTTPException( @@ -3045,7 +3069,9 @@ async def bulk_update_team_keys( seen_hashes.add(h) requested_tokens.append(k) hashed_key_ids.append(h) - existing_keys = await VerificationTokenRepository(prisma_client).table.find_many( + existing_keys = await VerificationTokenRepository( + prisma_client + ).table.find_many( # any-ok: avoid model_validate on mocked rows where={"team_id": data.team_id, "token": {"in": hashed_key_ids}} ) @@ -3109,7 +3135,7 @@ async def bulk_update_team_keys( user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, llm_router=llm_router, - user_custom_key_update=user_custom_key_update, + user_custom_key_update=custom_key_update_hook, existing_key_row=existing_by_token[db_token], ) @@ -3419,7 +3445,7 @@ async def info_key_fn_v2( alias_rows = await VerificationTokenRepository(prisma_client).table.find_many( where={"key_alias": {"in": data.key_aliases}}, include={"litellm_budget_table": True}, - ) + ) # any-ok: raw prisma rows needed for the litellm_budget_table include; no repository helper supports it alias_tokens = [row.token for row in alias_rows if row.token] tokens_to_query.extend(alias_tokens) @@ -3501,7 +3527,9 @@ async def info_key_fn( hashed_key: Optional[str] = key if key is not None: hashed_key = _hash_token_if_needed(token=key) - key_info = await VerificationTokenRepository(prisma_client).table.find_unique( + key_info: LiteLLM_VerificationToken | None = await VerificationTokenRepository( + prisma_client + ).table.find_unique( # any-ok: raw prisma row needed for the litellm_budget_table include where={"token": hashed_key}, # type: ignore include={"litellm_budget_table": True}, ) @@ -3529,27 +3557,27 @@ async def info_key_fn( ) ## REMOVE HASHED TOKEN INFO BEFORE RETURNING ## try: - key_info = key_info.model_dump() + key_info_dict = key_info.model_dump() except Exception: # if using pydantic v1 - key_info = key_info.dict() - key_token_hash = key_info.pop("token") + key_info_dict = key_info.dict() # pyright: ignore[reportDeprecated] # pydantic v1 fallback path + key_token_hash = key_info_dict.pop("token") - model_max_budget = key_info.get("model_max_budget") or {} - budget_table = key_info.get("litellm_budget_table") or {} + model_max_budget = key_info_dict.get("model_max_budget") or {} + budget_table = key_info_dict.get("litellm_budget_table") or {} if not model_max_budget and isinstance(budget_table, dict): model_max_budget = budget_table.get("model_max_budget") or {} if model_max_budget and key_token_hash: - key_info["model_max_budget_usage"] = await _build_model_max_budget_usage( + key_info_dict["model_max_budget_usage"] = await _build_model_max_budget_usage( api_key_hash=key_token_hash, model_max_budget=model_max_budget, user_api_key_cache=user_api_key_cache, ) # Attach object_permission if object_permission_id is set - key_info = await attach_object_permission_to_dict(key_info, prisma_client) + key_info_dict = await attach_object_permission_to_dict(key_info_dict, prisma_client) - return {"key": key, "info": key_info} + return {"key": key, "info": key_info_dict} except Exception as e: raise handle_exception_on_proxy(e) @@ -4018,9 +4046,9 @@ async def delete_verification_tokens( try: if prisma_client: tokens = [_hash_token_if_needed(token=key) for key in tokens] - _keys_being_deleted: List[LiteLLM_VerificationToken] = await VerificationTokenRepository( - prisma_client - ).table.find_many(where={"token": {"in": tokens}}) + _keys_being_deleted = await VerificationTokenRepository(prisma_client).find_many( + where={"token": {"in": tokens}} + ) if len(_keys_being_deleted) == 0: raise HTTPException( @@ -4088,7 +4116,7 @@ def _transform_verification_tokens_to_deleted_records( keys: List[LiteLLM_VerificationToken], user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str] = None, -) -> List[Dict[str, Any]]: +) -> List[Dict[str, Any]]: # any-ok: per-field JSON-serialized Prisma create_many payload; values are heterogeneous """Transform verification tokens into deleted token records ready for persistence.""" if not keys: return [] @@ -4141,7 +4169,7 @@ def _transform_verification_tokens_to_deleted_records( async def _save_deleted_verification_token_records( - records: List[Dict[str, Any]], + records: List[Dict[str, Any]], # any-ok: heterogeneous JSON-serialized Prisma create_many payload prisma_client: PrismaClient, ) -> None: """Save deleted verification token records to the database.""" @@ -4175,7 +4203,7 @@ async def delete_key_aliases( user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str] = None, ) -> Tuple[Optional[Dict], List[LiteLLM_VerificationToken]]: - _keys_being_deleted = await VerificationTokenRepository(prisma_client).table.find_many( + _keys_being_deleted = await VerificationTokenRepository(prisma_client).find_many( where={"key_alias": {"in": key_aliases}} ) @@ -4212,7 +4240,9 @@ async def _rotate_master_key( from litellm.proxy.proxy_server import proxy_config try: - models: Optional[List] = await ModelRepository(prisma_client).table.find_many() + models: List | None = await ModelRepository( + prisma_client + ).table.find_many() # any-ok: avoid model_validate on mocked rows except Exception: models = None # 2. process model table @@ -4234,7 +4264,7 @@ async def _rotate_master_key( _dumped["model_info"] = prisma.Json(_dumped["model_info"]) # type: ignore[attr-defined] new_models.append(_dumped) verbose_proxy_logger.debug("Resetting proxy model table") - async with prisma_client.db.tx() as tx: + async with prisma_client.db.tx() as tx: # any-ok: Prisma transaction handle is an untyped runtime wrapper await tx.litellm_proxymodeltable.delete_many() verbose_proxy_logger.debug("Creating %s models", len(new_models)) await tx.litellm_proxymodeltable.create_many( @@ -4307,7 +4337,9 @@ async def _rotate_master_key( # 5. process credentials table try: - credentials = await CredentialsRepository(prisma_client).table.find_many() + credentials = await CredentialsRepository( + prisma_client + ).table.find_many() # any-ok: CredentialsRepository exposes no typed list helper for raw rows except Exception: credentials = None if credentials: @@ -4330,8 +4362,8 @@ async def _rotate_master_key( _cred_data["credential_info"] = prisma.Json( # type: ignore[attr-defined] _cred_data["credential_info"] ) - await CredentialsRepository(prisma_client).table.update( - where={"credential_name": cred.credential_name}, + await CredentialsRepository(prisma_client).update_by_name( + credential_name=cred.credential_name, data={ **_cred_data, "updated_by": user_api_key_dict.user_id, @@ -4481,7 +4513,9 @@ async def _insert_deprecated_key( try: revoke_at = datetime.now(timezone.utc) + timedelta(seconds=grace_seconds) - await DeprecatedVerificationTokenRepository(prisma_client).table.upsert( + await DeprecatedVerificationTokenRepository( + prisma_client + ).table.upsert( # any-ok: PrismaTableRepository is a thin passthrough with no typed upsert helper where={"token": old_token_hash}, data={ "create": { @@ -4574,7 +4608,7 @@ async def _execute_virtual_key_regeneration( updated_token = await VerificationTokenRepository(prisma_client).table.update( where={"token": hashed_api_key}, data=update_data, # type: ignore - ) + ) # any-ok: raw prisma row needed verbatim for GenerateKeyResponse.model_validate below updated_token_dict = dict(updated_token) if updated_token is not None else {} updated_token_dict["key"] = new_token updated_token_dict["token_id"] = updated_token_dict.pop("token") @@ -4774,7 +4808,7 @@ async def regenerate_key_fn( _key_in_db = await VerificationTokenRepository(prisma_client).table.find_unique( where={"token": hashed_api_key}, - ) + ) # any-ok: tests patch VerificationTokenRepository/table.find_unique directly if _key_in_db is None: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -4840,7 +4874,7 @@ async def regenerate_key_fn( verbose_proxy_logger.info( "Key regeneration requested: key_alias=%s", - getattr(_key_in_db, "key_alias", None), + _key_in_db.key_alias, ) verbose_proxy_logger.debug("key_in_db: %s", _key_in_db) @@ -4956,7 +4990,7 @@ async def reset_key_spend_fn( None, description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", ), -) -> Dict[str, Any]: +) -> Mapping[str, object]: try: from litellm.proxy.proxy_server import ( hash_token, @@ -4976,10 +5010,12 @@ async def reset_key_spend_fn( else: hashed_api_key = hash_token(key) - _key_in_db = await VerificationTokenRepository(prisma_client).table.find_unique( + _key_in_db: LiteLLM_VerificationToken | None = await VerificationTokenRepository( + prisma_client + ).table.find_unique( where={"token": hashed_api_key}, include={"litellm_budget_table": True}, - ) + ) # any-ok: raw prisma row needed for the litellm_budget_table include; no repository helper supports it if _key_in_db is None: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -4999,7 +5035,7 @@ async def reset_key_spend_fn( updated_key = await VerificationTokenRepository(prisma_client).table.update( where={"token": hashed_api_key}, data={"spend": reset_to}, - ) + ) # any-ok: avoid model_validate on mocked rows if updated_key is None: raise HTTPException( @@ -5067,7 +5103,9 @@ async def validate_key_list_check( param="user_id", code=status.HTTP_403_FORBIDDEN, ) - complete_user_info_db_obj: Optional[BaseModel] = await UserRepository(prisma_client).table.find_unique( + complete_user_info_db_obj: BaseModel | None = await UserRepository( + prisma_client + ).table.find_unique( # any-ok: raw prisma row needed for the organization_memberships include where={"user_id": user_api_key_dict.user_id}, include={"organization_memberships": True}, ) @@ -5114,7 +5152,9 @@ async def validate_key_list_check( if key_hash: try: - key_info = await VerificationTokenRepository(prisma_client).table.find_unique( + key_info: LiteLLM_VerificationToken = await VerificationTokenRepository( + prisma_client + ).table.find_unique( # any-ok: preserves pre-existing behavior of erroring if key_hash doesn't exist where={"token": key_hash}, ) except Exception: @@ -5147,7 +5187,9 @@ async def _fetch_user_team_objects( if complete_user_info is None or not complete_user_info.teams: return [] - teams: Optional[List[BaseModel]] = await TeamRepository(prisma_client).table.find_many( + teams: List[BaseModel] | None = await TeamRepository( + prisma_client + ).table.find_many( # any-ok: avoid model_validate on mocked rows where={"team_id": {"in": complete_user_info.teams}} ) if teams is None: @@ -5421,8 +5463,8 @@ async def list_keys( async def _apply_non_admin_alias_scope( user_api_key_dict: UserAPIKeyAuth, - prisma_client: Any, - query_params: List[Any], + prisma_client: PrismaClient, + query_params: List[str | int], # mutable-ok: built incrementally across conditional branches for SQL args where_parts: List[str], ) -> None: """Append SQL scope conditions so non-admin users only see aliases for @@ -5435,7 +5477,9 @@ async def _apply_non_admin_alias_scope( # Look up the user's teams from the user table user_teams: List[str] = [] if user_api_key_dict.user_id: - user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_api_key_dict.user_id}) + user_row = await UserRepository(prisma_client).table.find_unique( # any-ok: avoid model_validate on mocks + where={"user_id": user_api_key_dict.user_id} + ) if user_row is not None: user_teams = getattr(user_row, "teams", []) or [] @@ -5463,7 +5507,7 @@ async def key_aliases( size: int = Query(50, ge=1, le=100, description="Page size"), search: Optional[str] = Query(None, description="Search key aliases (case-insensitive partial match)"), team_id: Optional[str] = Query(None, description="Filter aliases to keys belonging to this team"), -) -> Dict[str, Any]: +) -> Mapping[str, object]: """ Lists key aliases with pagination and optional search. @@ -5493,7 +5537,9 @@ async def key_aliases( # support column-level SELECT projection on find_many. # # $1 is always UI_SESSION_TOKEN_TEAM_ID (filters out UI session tokens). - query_params: List[Any] = [UI_SESSION_TOKEN_TEAM_ID] + query_params: List[str | int] = [ # mutable-ok: built incrementally across conditional branches for SQL args + UI_SESSION_TOKEN_TEAM_ID + ] where_parts = [ "key_alias IS NOT NULL", "key_alias != ''", @@ -5601,7 +5647,9 @@ def _validate_sort_params(sort_by: Optional[str], sort_order: str) -> Optional[D return order_by -def _build_expires_where_clause(expires_filter: str, now: datetime) -> dict[str, Any]: +def _build_expires_where_clause( + expires_filter: str, now: datetime +) -> dict[str, Any]: # any-ok: heterogeneous Prisma where-clause shape if expires_filter == "expired": return {"AND": [{"expires": {"not": None}}, {"expires": {"lt": now}}]} return {"OR": [{"expires": None}, {"expires": {"gte": now}}]} @@ -5622,7 +5670,9 @@ def _build_key_filter_conditions( agent_id: Optional[str] = None, use_substring_matching: bool = False, expires_filter: str | None = None, -) -> Dict[str, Union[str, Dict[str, Any], List[Dict[str, Any]]]]: +) -> Dict[ + str, Union[str, Dict[str, Any], List[Dict[str, Any]]] +]: # any-ok: heterogeneous Prisma where-clause built incrementally across branches """Build filter conditions for key listing. Visibility rules: @@ -5634,14 +5684,16 @@ def _build_key_filter_conditions( so former members cannot see service accounts they created after leaving. """ # Prepare filter conditions - where: Dict[str, Union[str, Dict[str, Any], List[Dict[str, Any]]]] = {} + where: Dict[ + str, Union[str, Dict[str, Any], List[Dict[str, Any]]] + ] = {} # any-ok: same heterogeneous Prisma where-clause shape as the return type where.update(_get_condition_to_filter_out_ui_session_tokens()) # Build the OR conditions for user's keys and admin team keys - or_conditions: List[Dict[str, Any]] = [] + or_conditions: List[Dict[str, Any]] = [] # any-ok: heterogeneous Prisma OR-clause entries # Base conditions for user's own keys - user_condition: Dict[str, Any] = {} + user_condition: Dict[str, Any] = {} # any-ok: heterogeneous Prisma where-clause fragment if user_id and isinstance(user_id, str): if use_substring_matching: user_condition["user_id"] = { @@ -5815,7 +5867,9 @@ async def _list_key_helper( # Fetch keys with pagination if use_deleted_table: - keys = await DeletedVerificationTokenRepository(prisma_client).table.find_many( + keys = await DeletedVerificationTokenRepository( + prisma_client + ).table.find_many( # any-ok: multi-column list-based order isn't supported by the typed find_many helper where=where, # type: ignore skip=skip, # type: ignore take=size, # type: ignore @@ -5829,7 +5883,9 @@ async def _list_key_helper( ), ) else: - keys = await VerificationTokenRepository(prisma_client).table.find_many( + keys = await VerificationTokenRepository( + prisma_client + ).table.find_many( # any-ok: multi-column list-based order + include aren't supported by find_many where=where, # type: ignore skip=skip, # type: ignore take=size, # type: ignore @@ -5848,13 +5904,9 @@ async def _list_key_helper( # Get total count of keys if use_deleted_table: - total_count = await DeletedVerificationTokenRepository(prisma_client).table.count( - where=where # type: ignore - ) + total_count = await DeletedVerificationTokenRepository(prisma_client).table.count(where=where) # type: ignore else: - total_count = await VerificationTokenRepository(prisma_client).table.count( - where=where # type: ignore - ) + total_count = await VerificationTokenRepository(prisma_client).table.count(where=where) # type: ignore verbose_proxy_logger.debug(f"Total count of keys: {total_count}") @@ -5868,7 +5920,9 @@ async def _list_key_helper( created_by_ids = [key.created_by for key in keys if key.created_by] all_ids = list(set(user_ids + created_by_ids)) # Remove duplicates if all_ids: - users = await UserRepository(prisma_client).table.find_many(where={"user_id": {"in": all_ids}}) + users = await UserRepository(prisma_client).table.find_many( # any-ok: avoid model_validate on mocked rows + where={"user_id": {"in": all_ids}} + ) user_map = {user.user_id: user for user in users} # Prepare response @@ -5917,7 +5971,9 @@ async def _list_key_helper( ) -def _get_condition_to_filter_out_ui_session_tokens() -> Dict[str, Any]: +def _get_condition_to_filter_out_ui_session_tokens() -> Dict[ + str, Any +]: # any-ok: heterogeneous Prisma where-clause shape """ Condition to filter out UI session tokens """ @@ -5932,7 +5988,7 @@ def _get_condition_to_filter_out_ui_session_tokens() -> Dict[str, Any]: async def _check_key_admin_access( user_api_key_dict: UserAPIKeyAuth, hashed_token: str, - prisma_client: Any, + prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, route: str, ) -> None: @@ -5951,7 +6007,9 @@ async def _check_key_admin_access( return # Look up the target key to find its team - target_key_row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token}) + target_key_row = await VerificationTokenRepository(prisma_client).table.find_unique( + where={"token": hashed_token} + ) # any-ok: avoid model_validate on mocked rows if target_key_row is None: raise HTTPException( status_code=404, @@ -6047,7 +6105,9 @@ async def block_key( ) # Check if the key exists before trying to block it - existing_record = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token}) + existing_record = await VerificationTokenRepository(prisma_client).table.find_unique( + where={"token": hashed_token} + ) # any-ok: avoid model_validate on mocked rows if existing_record is None: raise ProxyException( message="Key not found.", @@ -6080,7 +6140,7 @@ async def block_key( record = await VerificationTokenRepository(prisma_client).table.update( where={"token": hashed_token}, data={"blocked": True}, # type: ignore - ) + ) # any-ok: avoid model_validate on mocked rows ## UPDATE KEY CACHE - invalidate so next read re-fetches from DB await _delete_cache_key_object( @@ -6158,7 +6218,9 @@ async def unblock_key( ) # Check if the key exists before trying to unblock it - existing_record = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token}) + existing_record = await VerificationTokenRepository(prisma_client).table.find_unique( + where={"token": hashed_token} + ) # any-ok: avoid model_validate on mocked rows if existing_record is None: raise ProxyException( message="Key not found.", @@ -6191,7 +6253,7 @@ async def unblock_key( record = await VerificationTokenRepository(prisma_client).table.update( where={"token": hashed_token}, data={"blocked": False}, # type: ignore - ) + ) # any-ok: avoid model_validate on mocked rows ## UPDATE KEY CACHE - invalidate so next read re-fetches from DB await _delete_cache_key_object( @@ -6322,7 +6384,7 @@ async def _can_user_query_key_info( async def test_key_logging( user_api_key_dict: UserAPIKeyAuth, request: Request, - key_logging: List[Dict[str, Any]], + key_logging: List[Dict[str, Any]], # any-ok: heterogeneous logging-callback config from decrypt_callback_vars ) -> LoggingCallbackStatus: """ Test the key-based logging @@ -6443,7 +6505,7 @@ def _validate_key_alias_format(key_alias: Optional[str]) -> None: async def _enforce_unique_key_alias( key_alias: Optional[str], - prisma_client: Any, + prisma_client: PrismaClient | None, existing_key_token: Optional[str] = None, ) -> None: """ @@ -6451,7 +6513,7 @@ async def _enforce_unique_key_alias( Args: key_alias (Optional[str]): The key alias to check - prisma_client (Any): Prisma client instance + prisma_client (Optional[PrismaClient]): Prisma client instance existing_key_token (Optional[str]): ID of existing key being updated, to exclude from uniqueness check (The Admin UI passes key_alias, in all Edit key requests. So we need to be sure that if we find a key with the same alias, it's not the same key we're updating) @@ -6459,12 +6521,16 @@ async def _enforce_unique_key_alias( ProxyException: If key alias already exists on a different key """ if key_alias is not None and prisma_client is not None: - where_clause: dict[str, Any] = {"key_alias": key_alias} + where_clause: dict[str, Any] = {"key_alias": key_alias} # any-ok: heterogeneous Prisma where-clause fragment if existing_key_token: # Exclude the current key from the uniqueness check where_clause["NOT"] = {"token": existing_key_token} - existing_key = await VerificationTokenRepository(prisma_client).table.find_first(where=where_clause) + existing_key = await VerificationTokenRepository( + prisma_client + ).table.find_first( # any-ok: no typed find_first helper on the repository + where=where_clause + ) if existing_key is not None: raise ProxyException( message=f"Key with alias '{key_alias}' already exists. Unique key aliases across all keys are required.",