From 2f134a03e3e3097818019710340882ce13db769e Mon Sep 17 00:00:00 2001 From: Harshit Jain Date: Thu, 5 Feb 2026 14:14:20 +0530 Subject: [PATCH] Fix authorization issues, same alias; verified working --- .../proxy/hooks/key_management_event_hooks.py | 17 +- .../key_management_endpoints.py | 242 ++++++++++-------- .../secret_managers/aws_secret_manager_v2.py | 161 +++++++++--- .../test_aws_secret_manager_rotation.py | 109 ++++++++ 4 files changed, 394 insertions(+), 135 deletions(-) create mode 100644 tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index a8325d3461..c07f30f864 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -150,15 +150,26 @@ class KeyManagementEventHooks: existing_key_row.key_alias or f"virtual-key-{existing_key_row.token}" ) + new_secret_name = ( + response.key_alias + or data.key_alias + or f"virtual-key-{response.token_id}" + ) + verbose_proxy_logger.info( + "Updating secret in secret manager: secret_name=%s", + new_secret_name, + ) team_id = getattr(existing_key_row, "team_id", None) await KeyManagementEventHooks._rotate_virtual_key_in_secret_manager( current_secret_name=initial_secret_name, - new_secret_name=response.key_alias - or data.key_alias - or f"virtual-key-{response.token_id}", + new_secret_name=new_secret_name, new_secret_value=response.key, team_id=team_id, ) + verbose_proxy_logger.info( + "Secret updated in secret manager: secret_name=%s", + new_secret_name, + ) except Exception as e: verbose_proxy_logger.warning( f"Failed to rotate virtual key in secret manager: {e}" diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 9dadffca35..ddf746843c 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -518,7 +518,7 @@ async def _common_key_generation_helper( # noqa: PLR0915 ) # Handle special case where duration is "-1" (never expires) if value == "-1": - user_duration = float('inf') # Infinite duration + user_duration = float("inf") # Infinite duration else: user_duration = duration_in_seconds(duration=value) if user_duration > upperbound_duration: @@ -660,9 +660,9 @@ async def _common_key_generation_helper( # noqa: PLR0915 request_type="key", **data_json, table_name="key" ) - response["soft_budget"] = ( - data.soft_budget - ) # include the user-input soft budget in the response + response[ + "soft_budget" + ] = data.soft_budget # include the user-input soft budget in the response response = GenerateKeyResponse(**response) @@ -1083,12 +1083,16 @@ async def generate_key_fn( if data.max_budget is not None and data.max_budget < 0: raise HTTPException( status_code=400, - detail={"error": f"max_budget cannot be negative. Received: {data.max_budget}"} + detail={ + "error": f"max_budget cannot be negative. Received: {data.max_budget}" + }, ) if data.soft_budget is not None and data.soft_budget < 0: raise HTTPException( status_code=400, - detail={"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"} + detail={ + "error": f"soft_budget cannot be negative. Received: {data.soft_budget}" + }, ) if user_custom_key_generate is not None: @@ -1399,8 +1403,13 @@ async def prepare_key_update_data( validate_model_max_budget(non_default_values["model_max_budget"]) # Serialize router_settings to JSON if present - if "router_settings" in non_default_values and non_default_values["router_settings"] is not None: - non_default_values["router_settings"] = safe_dumps(non_default_values["router_settings"]) + if ( + "router_settings" in non_default_values + and non_default_values["router_settings"] is not None + ): + non_default_values["router_settings"] = safe_dumps( + non_default_values["router_settings"] + ) non_default_values = prepare_metadata_fields( data=data, non_default_values=non_default_values, existing_metadata=_metadata @@ -1448,19 +1457,17 @@ def is_different_team( def _validate_max_budget(max_budget: Optional[float]) -> None: """ Validate that max_budget is not negative. - + Args: max_budget: The max_budget value to validate - + Raises: HTTPException: If max_budget is negative """ if max_budget is not None and max_budget < 0: raise HTTPException( status_code=400, - detail={ - "error": f"max_budget cannot be negative. Received: {max_budget}" - }, + detail={"error": f"max_budget cannot be negative. Received: {max_budget}"}, ) @@ -1469,14 +1476,14 @@ async def _get_and_validate_existing_key( ) -> LiteLLM_VerificationToken: """ Get existing key from database and validate it exists. - + Args: token: The key token to look up prisma_client: Prisma client instance - + Returns: LiteLLM_VerificationToken: The existing key row - + Raises: HTTPException: If key is not found """ @@ -1485,19 +1492,19 @@ async def _get_and_validate_existing_key( status_code=500, detail={"error": "Database not connected"}, ) - + existing_key_row = await prisma_client.get_data( token=token, table_name="key", query_type="find_unique", ) - + if existing_key_row is None: raise HTTPException( status_code=404, detail={"error": f"Key not found: {token}"}, ) - + return existing_key_row @@ -1512,10 +1519,10 @@ async def _process_single_key_update( ) -> Dict[str, Any]: """ Process a single key update with all validations and checks. - + This function encapsulates all the logic for updating a single key, including validation, permission checks, team checks, and database updates. - + Args: key_update_item: The key update request item user_api_key_dict: The authenticated user's API key info @@ -1524,22 +1531,22 @@ async def _process_single_key_update( user_api_key_cache: User API key cache proxy_logging_obj: Proxy logging object llm_router: LLM router instance - + Returns: Dict containing the updated key information - + Raises: HTTPException: For various validation and permission errors """ # Validate max_budget _validate_max_budget(key_update_item.max_budget) - + # Get and validate existing key existing_key_row = await _get_and_validate_existing_key( token=key_update_item.key, prisma_client=prisma_client, ) - + # Check team member permissions if prisma_client is not None: await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( @@ -1549,7 +1556,7 @@ async def _process_single_key_update( existing_key_row=existing_key_row, user_api_key_cache=user_api_key_cache, ) - + # Create UpdateKeyRequest from BulkUpdateKeyRequestItem update_key_request = UpdateKeyRequest( key=key_update_item.key, @@ -1558,7 +1565,7 @@ async def _process_single_key_update( team_id=key_update_item.team_id, tags=key_update_item.tags, ) - + # Get team object and check team limits if team_id is provided team_obj: Optional[LiteLLM_TeamTableCachedObj] = None if update_key_request.team_id is not None: @@ -1568,18 +1575,16 @@ async def _process_single_key_update( user_api_key_cache=user_api_key_cache, check_db_only=True, ) - + if team_obj is not None and prisma_client is not None: await _check_team_key_limits( team_table=team_obj, data=update_key_request, prisma_client=prisma_client, ) - + # Validate team change if team is being changed - if is_different_team( - data=update_key_request, existing_key_row=existing_key_row - ): + if is_different_team(data=update_key_request, existing_key_row=existing_key_row): if llm_router is None: raise HTTPException( status_code=400, @@ -1590,9 +1595,7 @@ async def _process_single_key_update( if team_obj is None: raise HTTPException( status_code=500, - detail={ - "error": "Team object not found for team change validation" - }, + detail={"error": "Team object not found for team change validation"}, ) validate_key_team_change( key=existing_key_row, @@ -1600,31 +1603,29 @@ async def _process_single_key_update( change_initiated_by=user_api_key_dict, llm_router=llm_router, ) - + # Prepare update data non_default_values = await prepare_key_update_data( data=update_key_request, existing_key_row=existing_key_row ) - + # Update key in database if prisma_client is None: raise HTTPException( status_code=500, detail={"error": "Database not connected"}, ) - + _data = {**non_default_values, "token": key_update_item.key} - response = await prisma_client.update_data( - token=key_update_item.key, data=_data - ) - + response = await prisma_client.update_data(token=key_update_item.key, data=_data) + # Delete cache await _delete_cache_key_object( hashed_token=hash_token(key_update_item.key), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) - + # Trigger async hook asyncio.create_task( KeyManagementEventHooks.async_key_updated_hook( @@ -1635,19 +1636,19 @@ async def _process_single_key_update( litellm_changed_by=litellm_changed_by, ) ) - + if response is None: raise ValueError("Failed to update key got response = None") - + # Extract and format updated key info updated_key_info = response.get("data", {}) if hasattr(updated_key_info, "model_dump"): updated_key_info = updated_key_info.model_dump() elif hasattr(updated_key_info, "dict"): updated_key_info = updated_key_info.dict() - + updated_key_info.pop("token", None) - + return updated_key_info @@ -1740,7 +1741,9 @@ async def update_key_fn( if data.max_budget is not None and data.max_budget < 0: raise HTTPException( status_code=400, - detail={"error": f"max_budget cannot be negative. Received: {data.max_budget}"} + detail={ + "error": f"max_budget cannot be negative. Received: {data.max_budget}" + }, ) data_json: dict = data.model_dump(exclude_unset=True, exclude_none=True) @@ -1959,13 +1962,11 @@ async def bulk_update_keys( proxy_logging_obj, user_api_key_cache, ) - + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: raise HTTPException( status_code=403, - detail={ - "error": "Only proxy admins can perform bulk key updates" - }, + detail={"error": "Only proxy admins can perform bulk key updates"}, ) if prisma_client is None: @@ -2381,10 +2382,10 @@ async def info_key_fn( # if using pydantic v1 key_info = key_info.dict() key_info.pop("token") - + # Attach object_permission if object_permission_id is set key_info = await attach_object_permission_to_dict(key_info, prisma_client) - + return {"key": key, "info": key_info} except Exception as e: raise handle_exception_on_proxy(e) @@ -2509,7 +2510,9 @@ async def generate_key_helper_fn( # noqa: PLR0915 aliases_json = json.dumps(aliases) config_json = json.dumps(config) permissions_json = json.dumps(permissions) - router_settings_json = safe_dumps(router_settings) if router_settings is not None else safe_dumps({}) + router_settings_json = ( + safe_dumps(router_settings) if router_settings is not None else safe_dumps({}) + ) # Add model_rpm_limit and model_tpm_limit to metadata if model_rpm_limit is not None: @@ -2676,10 +2679,12 @@ async def generate_key_helper_fn( # noqa: PLR0915 ) key_data["created_at"] = getattr(create_key_response, "created_at", None) key_data["updated_at"] = getattr(create_key_response, "updated_at", None) - + # Deserialize router_settings from JSON string to dict for response router_settings_value = key_data.get("router_settings") - if router_settings_value is not None and isinstance(router_settings_value, str): + if router_settings_value is not None and isinstance( + router_settings_value, str + ): try: key_data["router_settings"] = yaml.safe_load(router_settings_value) except yaml.YAMLError: @@ -2762,28 +2767,35 @@ async def can_modify_verification_token( ) -> bool: """ Check if user has permission to modify (delete/regenerate) a verification token. - + Rules: - Proxy admin can modify any key + - Internal jobs service account can modify any key (for auto-rotation) - For team keys: only team admin or key owner can modify - For personal keys: only key owner can modify - + Args: key_info: The verification token to check user_api_key_cache: Cache for user API keys user_api_key_dict: The user making the request prisma_client: Prisma client for database access - + Returns: True if user can modify the key, False otherwise """ + from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME + is_team_key = _is_team_key(data=key_info) - + # 1. Proxy admin can modify any key if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: return True - - # 2. For team keys: only team admin or key owner can modify + + # 2. Internal jobs service account can modify any key (for auto-rotation) + if user_api_key_dict.api_key == LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME: + return True + + # 3. For team keys: only team admin or key owner can modify if is_team_key and key_info.team_id is not None: # Get team object to check if user is team admin team_table = await get_team_object( @@ -2792,34 +2804,35 @@ async def can_modify_verification_token( user_api_key_cache=user_api_key_cache, check_db_only=True, ) - + if team_table is None: return False - + # Check if user is team admin if _is_user_team_admin( user_api_key_dict=user_api_key_dict, team_obj=team_table, ): return True - + # Check if the key belongs to the user (they own it) - if key_info.user_id is not None and key_info.user_id == user_api_key_dict.user_id: + if ( + key_info.user_id is not None + and key_info.user_id == user_api_key_dict.user_id + ): return True - + # Not team admin and doesn't own the key return False - - # 3. For personal keys: only key owner can modify + + # 4. For personal keys: only key owner can modify if key_info.user_id is not None and key_info.user_id == user_api_key_dict.user_id: return True - + # Default: deny return False - - async def delete_verification_tokens( tokens: List, user_api_key_cache: DualCache, @@ -2849,10 +2862,10 @@ 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 prisma_client.db.litellm_verificationtoken.find_many( - where={"token": {"in": tokens}} - ) + _keys_being_deleted: List[ + LiteLLM_VerificationToken + ] = await prisma_client.db.litellm_verificationtoken.find_many( + where={"token": {"in": tokens}} ) if len(_keys_being_deleted) == 0: @@ -2952,11 +2965,24 @@ def _transform_verification_tokens_to_deleted_records( if org_id_value is not None: record["organization_id"] = org_id_value - for json_field in ["aliases", "config", "permissions", "metadata", "model_spend", "model_max_budget", "router_settings"]: + for json_field in [ + "aliases", + "config", + "permissions", + "metadata", + "model_spend", + "model_max_budget", + "router_settings", + ]: if json_field in record and record[json_field] is not None: record[json_field] = json.dumps(record[json_field]) - for rel_key in ("litellm_budget_table", "litellm_organization_table", "object_permission", "id"): + for rel_key in ( + "litellm_budget_table", + "litellm_organization_table", + "object_permission", + "id", + ): record.pop(rel_key, None) records.append(record) @@ -2971,9 +2997,7 @@ async def _save_deleted_verification_token_records( """Save deleted verification token records to the database.""" if not records: return - await prisma_client.db.litellm_deletedverificationtoken.create_many( - data=records - ) + await prisma_client.db.litellm_deletedverificationtoken.create_many(data=records) async def _persist_deleted_verification_tokens( @@ -3036,9 +3060,9 @@ async def _rotate_master_key( from litellm.proxy.proxy_server import proxy_config try: - models: Optional[List] = ( - await prisma_client.db.litellm_proxymodeltable.find_many() - ) + models: Optional[ + List + ] = await prisma_client.db.litellm_proxymodeltable.find_many() except Exception: models = None # 2. process model table @@ -3115,7 +3139,9 @@ async def _rotate_master_key( updated_patch=decrypted_cred, new_encryption_key=new_master_key, ) - credential_object_jsonified = jsonify_object(encrypted_cred.model_dump()) + credential_object_jsonified = jsonify_object( + encrypted_cred.model_dump() + ) await prisma_client.db.litellm_credentialstable.update( where={"credential_name": cred.credential_name}, data={ @@ -3160,7 +3186,7 @@ def get_new_token(data: Optional[RegenerateKeyRequest]) -> str: dependencies=[Depends(user_api_key_auth)], ) @management_endpoint_wrapper -async def regenerate_key_fn( +async def regenerate_key_fn( # noqa: PLR0915 key: Optional[str] = None, data: Optional[RegenerateKeyRequest] = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -3311,6 +3337,10 @@ async def regenerate_key_fn( detail={"error": "You are not authorized to regenerate this key"}, ) + verbose_proxy_logger.info( + "Key regeneration requested: key_alias=%s", + getattr(_key_in_db, "key_alias", None), + ) verbose_proxy_logger.debug("key_in_db: %s", _key_in_db) new_token = get_new_token(data=data) @@ -3361,6 +3391,10 @@ async def regenerate_key_fn( **updated_token_dict, ) + verbose_proxy_logger.info( + "Key regeneration completed: key_alias=%s", + getattr(_key_in_db, "key_alias", None), + ) asyncio.create_task( KeyManagementEventHooks.async_key_rotated_hook( data=data, @@ -3427,7 +3461,9 @@ def _validate_reset_spend_value( if reset_to > current_spend: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": f"reset_to ({reset_to}) must be <= current spend ({current_spend})"}, + detail={ + "error": f"reset_to ({reset_to}) must be <= current spend ({current_spend})" + }, ) max_budget = key_in_db.max_budget @@ -3553,11 +3589,11 @@ async def validate_key_list_check( param="user_id", code=status.HTTP_403_FORBIDDEN, ) - complete_user_info_db_obj: Optional[BaseModel] = ( - await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": user_api_key_dict.user_id}, - include={"organization_memberships": True}, - ) + complete_user_info_db_obj: Optional[ + BaseModel + ] = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": user_api_key_dict.user_id}, + include={"organization_memberships": True}, ) if complete_user_info_db_obj is None: @@ -3643,10 +3679,10 @@ async def get_admin_team_ids( if complete_user_info is None: return [] # Get all teams that user is an admin of - teams: Optional[List[BaseModel]] = ( - await prisma_client.db.litellm_teamtable.find_many( - where={"team_id": {"in": complete_user_info.teams}} - ) + teams: Optional[ + List[BaseModel] + ] = await prisma_client.db.litellm_teamtable.find_many( + where={"team_id": {"in": complete_user_info.teams}} ) if teams is None: return [] @@ -3691,8 +3727,12 @@ async def list_keys( description="Column to sort by (e.g. 'user_id', 'created_at', 'spend')", ), sort_order: str = Query(default="desc", description="Sort order ('asc' or 'desc')"), - expand: Optional[List[str]] = Query(None, description="Expand related objects (e.g. 'user')"), - status: Optional[str] = Query(None, description="Filter by status (e.g. 'deleted')"), + expand: Optional[List[str]] = Query( + None, description="Expand related objects (e.g. 'user')" + ), + status: Optional[str] = Query( + None, description="Filter by status (e.g. 'deleted')" + ), ) -> KeyListResponseObject: """ List all keys for a given user / team / organization. @@ -3784,7 +3824,9 @@ async def list_keys( message=getattr(e, "detail", f"error({str(e)})"), type=ProxyErrorTypes.internal_server_error, param=getattr(e, "param", "None"), - code=getattr(e, "status_code", fastapi.status.HTTP_500_INTERNAL_SERVER_ERROR), + code=getattr( + e, "status_code", fastapi.status.HTTP_500_INTERNAL_SERVER_ERROR + ), ) elif isinstance(e, ProxyException): raise e diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index 8edfc48336..c1b4d019dc 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -3,7 +3,8 @@ This is a file for the AWS Secret Manager Integration Handles Async Operations for: - Read Secret -- Write Secret +- Write Secret (CreateSecret) +- Update Secret (PutSecretValue) - for in-place rotation when alias is preserved - Delete Secret Relevant issue: https://github.com/BerriAI/litellm/issues/1883 @@ -42,11 +43,11 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): aws_profile_name: Optional[str] = None, aws_web_identity_token: Optional[str] = None, aws_sts_endpoint: Optional[str] = None, - **kwargs + **kwargs, ): BaseSecretManager.__init__(self, **kwargs) BaseAWSLLM.__init__(self, **kwargs) - + # Store AWS authentication settings self.aws_region_name = aws_region_name self.aws_role_name = aws_role_name @@ -61,7 +62,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): # AWS_REGION_NAME is only strictly required if not using a profile or role # When using IAM roles, the region can come from multiple sources if ( - "AWS_REGION_NAME" not in os.environ + "AWS_REGION_NAME" not in os.environ and "AWS_REGION" not in os.environ and "AWS_DEFAULT_REGION" not in os.environ ): @@ -83,22 +84,36 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): return try: cls.validate_environment() - + # Extract AWS settings from key_management_settings if provided aws_kwargs = {} if key_management_settings is not None: aws_kwargs = { - "aws_region_name": getattr(key_management_settings, "aws_region_name", None), - "aws_role_name": getattr(key_management_settings, "aws_role_name", None), - "aws_session_name": getattr(key_management_settings, "aws_session_name", None), - "aws_external_id": getattr(key_management_settings, "aws_external_id", None), - "aws_profile_name": getattr(key_management_settings, "aws_profile_name", None), - "aws_web_identity_token": getattr(key_management_settings, "aws_web_identity_token", None), - "aws_sts_endpoint": getattr(key_management_settings, "aws_sts_endpoint", None), + "aws_region_name": getattr( + key_management_settings, "aws_region_name", None + ), + "aws_role_name": getattr( + key_management_settings, "aws_role_name", None + ), + "aws_session_name": getattr( + key_management_settings, "aws_session_name", None + ), + "aws_external_id": getattr( + key_management_settings, "aws_external_id", None + ), + "aws_profile_name": getattr( + key_management_settings, "aws_profile_name", None + ), + "aws_web_identity_token": getattr( + key_management_settings, "aws_web_identity_token", None + ), + "aws_sts_endpoint": getattr( + key_management_settings, "aws_sts_endpoint", None + ), } # Remove None values aws_kwargs = {k: v for k, v in aws_kwargs.items() if v is not None} - + litellm.secret_manager_client = cls(**aws_kwargs) litellm._key_management_system = KeyManagementSystem.AWS_SECRET_MANAGER @@ -246,13 +261,13 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): return primary_secret_kv_pairs.get(secret_name) async def async_write_secret( - self, - secret_name: str, - secret_value: str, - description: Optional[str] = None, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - tags: Optional[Union[dict, list]] = None + self, + secret_name: str, + secret_value: str, + description: Optional[str] = None, + optional_params: Optional[dict] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + tags: Optional[Union[dict, list]] = None, ) -> dict: """ Async function to write a secret to AWS Secrets Manager @@ -312,6 +327,94 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): except httpx.TimeoutException: raise ValueError("Timeout error occurred") + async def async_put_secret_value( + self, + secret_name: str, + secret_value: str, + optional_params: Optional[dict] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + ) -> dict: + """ + Async function to update an existing secret's value in AWS Secrets Manager. + + Uses PutSecretValue to update in place. Use this when rotating a secret + that keeps the same name (current_secret_name == new_secret_name). + + Args: + secret_name: Name of the existing secret to update + secret_value: New value to store + optional_params: Additional AWS parameters + timeout: Request timeout + + Returns: + dict: Response from AWS Secrets Manager containing update details + """ + from litellm._uuid import uuid + + data: Dict[str, Any] = { + "SecretId": secret_name, + "SecretString": secret_value, + "ClientRequestToken": str(uuid.uuid4()), + } + + endpoint_url, headers, body = self._prepare_request( + action="PutSecretValue", + secret_name=secret_name, + secret_value=secret_value, + optional_params=optional_params, + request_data=data, + ) + + async_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.SecretManager, + params={"timeout": timeout}, + ) + + try: + response = await async_client.post( + url=endpoint_url, headers=headers, data=body.decode("utf-8") + ) + response.raise_for_status() + return response.json() + except httpx.HTTPStatusError as err: + raise ValueError(f"HTTP error occurred: {err.response.text}") + except httpx.TimeoutException: + raise ValueError("Timeout error occurred") + + async def async_rotate_secret( + self, + current_secret_name: str, + new_secret_name: str, + new_secret_value: str, + optional_params: Optional[dict] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + ) -> dict: + """ + Rotate a secret. When current_secret_name == new_secret_name (in-place + update), uses PutSecretValue instead of create+delete to avoid + ResourceExistsException. + """ + if current_secret_name == new_secret_name: + # Same alias: update in place via PutSecretValue + verbose_logger.info( + "Secret rotated in-place (PutSecretValue): secret_name=%s", + current_secret_name, + ) + return await self.async_put_secret_value( + secret_name=current_secret_name, + secret_value=new_secret_value, + optional_params=optional_params, + timeout=timeout, + ) + # Different names: create new, delete old (base class logic) + return await super().async_rotate_secret( + current_secret_name=current_secret_name, + new_secret_name=new_secret_name, + new_secret_value=new_secret_value, + optional_params=optional_params, + timeout=timeout, + ) + async def async_delete_secret( self, secret_name: str, @@ -375,7 +478,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") optional_params = optional_params or {} - + # Build optional_params from instance settings if not provided # This allows the IAM role settings to be used for Secret Manager calls if not optional_params.get("aws_role_name") and self.aws_role_name: @@ -388,11 +491,14 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): optional_params["aws_external_id"] = self.aws_external_id if not optional_params.get("aws_profile_name") and self.aws_profile_name: optional_params["aws_profile_name"] = self.aws_profile_name - if not optional_params.get("aws_web_identity_token") and self.aws_web_identity_token: + if ( + not optional_params.get("aws_web_identity_token") + and self.aws_web_identity_token + ): optional_params["aws_web_identity_token"] = self.aws_web_identity_token if not optional_params.get("aws_sts_endpoint") and self.aws_sts_endpoint: optional_params["aws_sts_endpoint"] = self.aws_sts_endpoint - + boto3_credentials_info = self._get_boto_credentials_from_optional_params( optional_params ) @@ -431,12 +537,3 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): prepped = request.prepare() return endpoint_url, prepped.headers, body - - -# if __name__ == "__main__": -# print("loading aws secret manager v2") -# aws_secret_manager_v2 = AWSSecretsManagerV2() -# import asyncio -# print("writing secret to aws secret manager v2") -# asyncio.run(aws_secret_manager_v2.async_write_secret(secret_name="test_secret_3", secret_value="test_value_2")) -# print("reading secret from aws secret manager v2") diff --git a/tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py b/tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py new file mode 100644 index 0000000000..8398248262 --- /dev/null +++ b/tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py @@ -0,0 +1,109 @@ +""" +Regression tests for AWS Secrets Manager same-name in-place rotation fix. + +When current_secret_name == new_secret_name (e.g. key alias preserved during +rotation), AWS must use PutSecretValue to update in place instead of +create+delete, which would fail with ResourceExistsException. +""" +from unittest.mock import AsyncMock, patch + +import pytest + +from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2 + + +@pytest.mark.asyncio +async def test_rotate_secret_same_name_uses_put_secret_value(): + """ + When current_secret_name == new_secret_name, async_rotate_secret should + call PutSecretValue (async_put_secret_value) instead of create+delete. + """ + secret_name = "litellm/tenant/litellm-metis-key" + new_value = "sk-new-rotated-key-value" + + with patch.object( + AWSSecretsManagerV2, + "async_put_secret_value", + new_callable=AsyncMock, + return_value={"ARN": "arn:aws:secretsmanager:us-east-1:123:secret:test"}, + ) as mock_put: + with patch.object( + AWSSecretsManagerV2, + "async_write_secret", + new_callable=AsyncMock, + ) as mock_write: + with patch.object( + AWSSecretsManagerV2, + "async_delete_secret", + new_callable=AsyncMock, + ) as mock_delete: + manager = AWSSecretsManagerV2() + result = await manager.async_rotate_secret( + current_secret_name=secret_name, + new_secret_name=secret_name, + new_secret_value=new_value, + ) + + # PutSecretValue (in-place update) should be called + mock_put.assert_called_once_with( + secret_name=secret_name, + secret_value=new_value, + optional_params=None, + timeout=None, + ) + # Create + delete should NOT be called + mock_write.assert_not_called() + mock_delete.assert_not_called() + assert result["ARN"] == "arn:aws:secretsmanager:us-east-1:123:secret:test" + + +@pytest.mark.asyncio +async def test_rotate_secret_different_names_uses_create_delete(): + """ + When current_secret_name != new_secret_name, async_rotate_secret should + use base class logic (create new, delete old). + """ + current_name = "litellm/old-key-alias" + new_name = "litellm/virtual-key-new-token-id" + new_value = "sk-new-key-value" + + with patch.object( + AWSSecretsManagerV2, + "async_read_secret", + new_callable=AsyncMock, + side_effect=["sk-old-value", new_value], # read old, then read new + ): + with patch.object( + AWSSecretsManagerV2, + "async_write_secret", + new_callable=AsyncMock, + return_value={"ARN": "arn:new"}, + ) as mock_write: + with patch.object( + AWSSecretsManagerV2, + "async_delete_secret", + new_callable=AsyncMock, + return_value={}, + ) as mock_delete: + with patch.object( + AWSSecretsManagerV2, + "async_put_secret_value", + new_callable=AsyncMock, + ) as mock_put: + manager = AWSSecretsManagerV2() + await manager.async_rotate_secret( + current_secret_name=current_name, + new_secret_name=new_name, + new_secret_value=new_value, + ) + + # PutSecretValue should NOT be called (different names) + mock_put.assert_not_called() + # Create + delete should be called + mock_write.assert_called_once() + mock_delete.assert_called_once_with( + secret_name=current_name, + recovery_window_in_days=7, + optional_params=None, + timeout=None, + )