diff --git a/docs/my-website/docs/proxy/virtual_keys.md b/docs/my-website/docs/proxy/virtual_keys.md index e1c89bbc216..999db6055d0 100644 --- a/docs/my-website/docs/proxy/virtual_keys.md +++ b/docs/my-website/docs/proxy/virtual_keys.md @@ -554,7 +554,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 d5dc841cb26..88ee88bcd0f 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -208,6 +208,8 @@ 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 class ConfigGeneralSettings(LiteLLMBase): diff --git a/litellm/proxy/db/dynamo_db.py b/litellm/proxy/db/dynamo_db.py index 83cf6b15724..50330ccf3bf 100644 --- a/litellm/proxy/db/dynamo_db.py +++ b/litellm/proxy/db/dynamo_db.py @@ -52,6 +52,32 @@ 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") + assumed_role = sts_client.assume_role( + RoleArn=self.database_arguments.aws_role_name, + RoleSessionName=self.database_arguments.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 @@ -74,6 +100,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)) @@ -170,6 +197,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: @@ -210,6 +239,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: @@ -253,6 +284,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 @@ -323,4 +355,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 97168b19f9a..2bed73bcb8b 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -67,12 +67,14 @@ litellm_settings: general_settings: master_key: sk-1234 - # database_type: "dynamo_db" - # database_args: { # 👈 all args - https://github.com/BerriAI/litellm/blob/befbcbb7ac8f59835ce47415c128decf37aac328/litellm/proxy/_types.py#L190 - # "billing_mode": "PAY_PER_REQUEST", - # "region_name": "us-west-2", - # "ssl_verify": False - # } + database_type: "dynamo_db" + database_args: { # 👈 all args - https://github.com/BerriAI/litellm/blob/befbcbb7ac8f59835ce47415c128decf37aac328/litellm/proxy/_types.py#L190 + "billing_mode": "PAY_PER_REQUEST", + "region_name": "us-west-2", + "ssl_verify": False, + "aws_role_name": "", + "aws_session_name": "", + }