Merge pull request #1589 from BerriAI/litellm_dynamo_use_arn

[Feat] Proxy DynamoDB - set arn number on dynamoDB /key/gen
This commit is contained in:
Ishaan Jaff 2024-02-13 21:27:53 -08:00 • committed by GitHub
commit 9b4cf8c91f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 73 additions and 1 deletions

View file

@ -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",
}
```

View file

@ -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):

View file

@ -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)

View file

@ -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)

View file

@ -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
)