diff --git a/docs/my-website/docs/proxy/virtual_keys.md b/docs/my-website/docs/proxy/virtual_keys.md index 28842e5e257..2be4b95c1f4 100644 --- a/docs/my-website/docs/proxy/virtual_keys.md +++ b/docs/my-website/docs/proxy/virtual_keys.md @@ -696,7 +696,9 @@ general_settings: "region_name": "us-west-2" "user_table_name": "your-user-table", "key_table_name": "your-token-table", - "config_table_name": "your-config-table" + "config_table_name": "your-config-table", + "aws_role_name": "your-aws_role_name", + "aws_session_name": "your-aws_session_name", } ``` diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0d7355dad78..c85564231e8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -234,6 +234,15 @@ class DynamoDBArgs(LiteLLMBase): key_table_name: str = "LiteLLM_VerificationToken" config_table_name: str = "LiteLLM_Config" spend_table_name: str = "LiteLLM_SpendLogs" + aws_role_name: Optional[str] = None + aws_session_name: Optional[str] = None + aws_web_identity_token: Optional[str] = None + aws_provider_id: Optional[str] = None + aws_policy_arns: Optional[List[str]] = None + aws_policy: Optional[str] = None + aws_duration_seconds: Optional[int] = None + assume_role_aws_role_name: Optional[str] = None + assume_role_aws_session_name: Optional[str] = None class ConfigGeneralSettings(LiteLLMBase): diff --git a/litellm/proxy/db/dynamo_db.py b/litellm/proxy/db/dynamo_db.py index b8d2b09ca4e..a9461d9225c 100644 --- a/litellm/proxy/db/dynamo_db.py +++ b/litellm/proxy/db/dynamo_db.py @@ -53,6 +53,41 @@ class DynamoDBWrapper(CustomDB): self.database_arguments = database_arguments self.region_name = database_arguments.region_name + def set_env_vars_based_on_arn(self): + if self.database_arguments.aws_role_name is None: + return + verbose_proxy_logger.debug( + f"DynamoDB: setting env vars based on arn={self.database_arguments.aws_role_name}" + ) + import boto3, os + + sts_client = boto3.client("sts") + + # call 1 + non_used_assumed_role = sts_client.assume_role_with_web_identity( + RoleArn=self.database_arguments.aws_role_name, + RoleSessionName=self.database_arguments.aws_session_name, + WebIdentityToken=self.database_arguments.aws_web_identity_token, + ) + + # call 2 + assumed_role = sts_client.assume_role( + RoleArn=self.database_arguments.assume_role_aws_role_name, + RoleSessionName=self.database_arguments.assume_role_aws_session_name, + ) + + aws_access_key_id = assumed_role["Credentials"]["AccessKeyId"] + aws_secret_access_key = assumed_role["Credentials"]["SecretAccessKey"] + aws_session_token = assumed_role["Credentials"]["SessionToken"] + + verbose_proxy_logger.debug( + f"Got STS assumed Role, aws_access_key_id={aws_access_key_id}" + ) + # set these in the env so aiodynamo can use them + os.environ["AWS_ACCESS_KEY_ID"] = aws_access_key_id + os.environ["AWS_SECRET_ACCESS_KEY"] = aws_secret_access_key + os.environ["AWS_SESSION_TOKEN"] = aws_session_token + async def connect(self): """ Connect to DB, and creating / updating any tables @@ -75,6 +110,7 @@ class DynamoDBWrapper(CustomDB): import aiohttp verbose_proxy_logger.debug("DynamoDB Wrapper - Attempting to connect") + self.set_env_vars_based_on_arn() # before making ClientSession check if ssl_verify=False if self.database_arguments.ssl_verify == False: client_session = ClientSession(connector=aiohttp.TCPConnector(ssl=False)) @@ -171,6 +207,8 @@ class DynamoDBWrapper(CustomDB): from aiohttp import ClientSession import aiohttp + self.set_env_vars_based_on_arn() + if self.database_arguments.ssl_verify == False: client_session = ClientSession(connector=aiohttp.TCPConnector(ssl=False)) else: @@ -214,6 +252,8 @@ class DynamoDBWrapper(CustomDB): from aiohttp import ClientSession import aiohttp + self.set_env_vars_based_on_arn() + if self.database_arguments.ssl_verify == False: client_session = ClientSession(connector=aiohttp.TCPConnector(ssl=False)) else: @@ -261,6 +301,7 @@ class DynamoDBWrapper(CustomDB): async def update_data( self, key: str, value: dict, table_name: Literal["user", "key", "config"] ): + self.set_env_vars_based_on_arn() from aiodynamo.client import Client from aiodynamo.credentials import Credentials, StaticCredentials from aiodynamo.http.httpx import HTTPX @@ -334,4 +375,5 @@ class DynamoDBWrapper(CustomDB): """ Not Implemented yet. """ + self.set_env_vars_based_on_arn() return super().delete_data(keys, table_name) diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 08a1e32ab67..8d35bcae815 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -62,6 +62,7 @@ general_settings: + environment_variables: # otel: True # OpenTelemetry Logger # master_key: sk-1234 # [OPTIONAL] Only use this if you to require all calls to contain this key (Authorization: Bearer sk-1234) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 621ef08e4e9..fdcc68dae54 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1444,6 +1444,24 @@ class ProxyConfig: database_type == "dynamo_db" or database_type == "dynamodb" ): database_args = general_settings.get("database_args", None) + ### LOAD FROM os.environ/ ### + for k, v in database_args.items(): + if isinstance(v, str) and v.startswith("os.environ/"): + database_args[k] = litellm.get_secret(v) + if isinstance(k, str) and k == "aws_web_identity_token": + value = database_args[k] + verbose_proxy_logger.debug( + f"Loading AWS Web Identity Token from file: {value}" + ) + if os.path.exists(value): + with open(value, "r") as file: + token_content = file.read() + database_args[k] = token_content + else: + verbose_proxy_logger.info( + f"DynamoDB Loading - {value} is not a valid file path" + ) + verbose_proxy_logger.debug(f"database_args: {database_args}") custom_db_client = DBClient( custom_db_args=database_args, custom_db_type=database_type )