diff --git a/litellm/proxy/_super_secret_config.yaml b/litellm/proxy/_super_secret_config.yaml index ec79cbbdf2..77bd6ee960 100644 --- a/litellm/proxy/_super_secret_config.yaml +++ b/litellm/proxy/_super_secret_config.yaml @@ -81,7 +81,3 @@ general_settings: enable_jwt_auth: True litellm_jwtauth: team_id_jwt_field: "client_id" -# key_management_system: "aws_kms" -# key_management_settings: -# hosted_keys: ["LITELLM_MASTER_KEY"] - diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 2960d9b1c1..70232c5be2 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -21,7 +21,7 @@ import shutil telemetry = None -def append_query_params(url, params): +def append_query_params(url, params) -> str: from litellm._logging import verbose_proxy_logger verbose_proxy_logger.debug(f"url: {url}") @@ -229,13 +229,29 @@ def run_server( ): args = locals() if local: - from proxy_server import app, save_worker_config, ProxyConfig + from proxy_server import ( + app, + save_worker_config, + ProxyConfig, + KeyManagementSystem, + KeyManagementSettings, + load_from_azure_key_vault, + load_aws_kms, + load_aws_secret_manager, + load_google_kms, + ) else: try: from .proxy_server import ( app, save_worker_config, ProxyConfig, + KeyManagementSystem, + KeyManagementSettings, + load_from_azure_key_vault, + load_aws_kms, + load_aws_secret_manager, + load_google_kms, ) except ImportError as e: if "litellm[proxy]" in str(e): @@ -247,6 +263,12 @@ def run_server( app, save_worker_config, ProxyConfig, + KeyManagementSystem, + KeyManagementSettings, + load_from_azure_key_vault, + load_aws_kms, + load_aws_secret_manager, + load_google_kms, ) if version == True: pkg_version = importlib.metadata.version("litellm") @@ -445,6 +467,40 @@ def run_server( general_settings = _config.get("general_settings", {}) if general_settings is None: general_settings = {} + if general_settings: + ### LOAD SECRET MANAGER ### + key_management_system = general_settings.get( + "key_management_system", None + ) + if key_management_system is not None: + if ( + key_management_system + == KeyManagementSystem.AZURE_KEY_VAULT.value + ): + ### LOAD FROM AZURE KEY VAULT ### + load_from_azure_key_vault(use_azure_key_vault=True) + elif key_management_system == KeyManagementSystem.GOOGLE_KMS.value: + ### LOAD FROM GOOGLE KMS ### + load_google_kms(use_google_kms=True) + elif ( + key_management_system + == KeyManagementSystem.AWS_SECRET_MANAGER.value # noqa: F405 + ): + ### LOAD FROM AWS SECRET MANAGER ### + load_aws_secret_manager(use_aws_secret_manager=True) + elif key_management_system == KeyManagementSystem.AWS_KMS.value: + load_aws_kms(use_aws_kms=True) + else: + raise ValueError("Invalid Key Management System selected") + key_management_settings = general_settings.get( + "key_management_settings", None + ) + if key_management_settings is not None: + import litellm + + litellm._key_management_settings = KeyManagementSettings( + **key_management_settings + ) database_url = general_settings.get("database_url", None) db_connection_pool_limit = general_settings.get( "database_connection_pool_limit", 100 @@ -460,7 +516,7 @@ def run_server( ) # Adds the parent directory to the system path - for litellm local dev import litellm - database_url = litellm.get_secret(database_url) + database_url = litellm.get_secret(database_url, default_value=None) os.chdir(original_dir) if database_url is not None and isinstance(database_url, str): os.environ["DATABASE_URL"] = database_url @@ -470,13 +526,15 @@ def run_server( or os.getenv("DIRECT_URL", None) is not None ): try: + from litellm import get_secret + if os.getenv("DATABASE_URL", None) is not None: ### add connection pool + pool timeout args params = { "connection_limit": db_connection_pool_limit, "pool_timeout": db_connection_timeout, } - database_url = os.getenv("DATABASE_URL") + database_url = get_secret("DATABASE_URL", default_value=None) modified_url = append_query_params(database_url, params) os.environ["DATABASE_URL"] = modified_url if os.getenv("DIRECT_URL", None) is not None: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0645f19ac1..8217a210a0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2626,7 +2626,7 @@ async def startup_event(): master_key = litellm.get_secret("LITELLM_MASTER_KEY", None) # check if DATABASE_URL in environment - load from there if prisma_client is None: - prisma_setup(database_url=os.getenv("DATABASE_URL")) + prisma_setup(database_url=litellm.get_secret("DATABASE_URL", None)) ### LOAD CONFIG ### worker_config = litellm.get_secret("WORKER_CONFIG") @@ -2752,9 +2752,6 @@ async def startup_event(): ) ) - verbose_proxy_logger.debug( - f"custom_db_client client {custom_db_client}. Master_key: {master_key}" - ) if custom_db_client is not None and master_key is not None: # add master key to db await generate_key_helper_fn( diff --git a/litellm/tests/test_text_completion.py b/litellm/tests/test_text_completion.py index 61f649a224..cac448c630 100644 --- a/litellm/tests/test_text_completion.py +++ b/litellm/tests/test_text_completion.py @@ -3990,7 +3990,7 @@ def test_async_text_completion(): asyncio.run(test_get_response()) -@pytest.mark.skip(reason="Tgai endpoints are unstable") +@pytest.mark.skip(reason="Skip flaky tgai test") def test_async_text_completion_together_ai(): litellm.set_verbose = True print("test_async_text_completion") @@ -3998,7 +3998,7 @@ def test_async_text_completion_together_ai(): async def test_get_response(): try: response = await litellm.atext_completion( - model="together_ai/codellama/CodeLlama-13b-Instruct-hf", + model="together_ai/mistralai/Mixtral-8x7B-Instruct-v0.1", prompt="good morning", max_tokens=10, ) diff --git a/litellm/utils.py b/litellm/utils.py index ae4c343ba0..885bd25334 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5652,6 +5652,14 @@ def exception_type( response=original_exception.response, litellm_debug_info=extra_information, ) + elif "A timeout occurred" in error_str: + exception_mapping_worked = True + raise Timeout( + message=f"{exception_provider} - {message}", + model=model, + llm_provider=custom_llm_provider, + litellm_debug_info=extra_information, + ) elif ( "invalid_request_error" in error_str and "content_policy_violation" in error_str @@ -6844,7 +6852,13 @@ def exception_type( llm_provider="together_ai", response=original_exception.response, ) - + elif "A timeout occurred" in error_str: + exception_mapping_worked = True + raise Timeout( + message=f"TogetherAIException - {error_str}", + model=model, + llm_provider="together_ai", + ) elif ( "error" in error_response and "API key doesn't match expected format." @@ -7283,7 +7297,6 @@ def get_secret( ): key_management_system = litellm._key_management_system key_management_settings = litellm._key_management_settings - args = locals() if secret_name.startswith("os.environ/"): secret_name = secret_name.replace("os.environ/", "") @@ -7417,19 +7430,24 @@ def get_secret( """ encrypted_value = os.getenv(secret_name, None) if encrypted_value is None: - raise Exception("encrypted value for AWS KMS cannot be None.") + raise Exception( + "AWS KMS - Encrypted Value of Key={} is None".format( + secret_name + ) + ) # Decode the base64 encoded ciphertext ciphertext_blob = base64.b64decode(encrypted_value) # Set up the parameters for the decrypt call params = {"CiphertextBlob": ciphertext_blob} - # Perform the decryption response = client.decrypt(**params) # Extract and decode the plaintext plaintext = response["Plaintext"] secret = plaintext.decode("utf-8") + if isinstance(secret, str): + secret = secret.strip() elif key_manager == KeyManagementSystem.AWS_SECRET_MANAGER.value: try: get_secret_value_response = client.get_secret_value( @@ -7456,7 +7474,7 @@ def get_secret( secret = client.get_secret(secret_name).secret_value except Exception as e: # check if it's in os.environ verbose_logger.error( - f"An exception occurred - {str(e)}\n\n{traceback.format_exc()}" + f"Defaulting to os.environ value for key={secret_name}. An exception occurred - {str(e)}.\n\n{traceback.format_exc()}" ) secret = os.getenv(secret_name) try: