mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: validate aws region name
This commit is contained in:
parent
609454d0f1
commit
d47948ab23
2 changed files with 164 additions and 0 deletions
|
|
@ -1,6 +1,7 @@
|
|||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import urllib.parse
|
||||
from datetime import datetime
|
||||
from typing import (
|
||||
|
|
@ -37,6 +38,11 @@ else:
|
|||
AWSPreparedRequest = Any
|
||||
|
||||
|
||||
# Real AWS region names are lowercase letters, digits, and hyphens
|
||||
# (e.g. "us-east-1", "eu-west-2", "us-gov-west-1", "cn-north-1").
|
||||
_VALID_AWS_REGION_PATTERN = re.compile(r"\A[a-z0-9-]+\Z")
|
||||
|
||||
|
||||
class Boto3CredentialsInfo(BaseModel):
|
||||
credentials: Credentials
|
||||
aws_region_name: str
|
||||
|
|
@ -284,6 +290,9 @@ class BaseAWSLLM:
|
|||
if not region: # Check if region is empty
|
||||
return None
|
||||
|
||||
if not _VALID_AWS_REGION_PATTERN.match(region):
|
||||
return None
|
||||
|
||||
return region
|
||||
except Exception:
|
||||
# Catch any unexpected errors and return None
|
||||
|
|
@ -481,6 +490,7 @@ class BaseAWSLLM:
|
|||
str: The AWS region name
|
||||
"""
|
||||
aws_region_name = optional_params.get("aws_region_name", None)
|
||||
self._validate_aws_region_name(aws_region_name)
|
||||
### SET REGION NAME ###
|
||||
if aws_region_name is None:
|
||||
# check model arn #
|
||||
|
|
@ -519,8 +529,25 @@ class BaseAWSLLM:
|
|||
except Exception:
|
||||
aws_region_name = "us-west-2"
|
||||
|
||||
self._validate_aws_region_name(aws_region_name)
|
||||
return aws_region_name
|
||||
|
||||
@staticmethod
|
||||
def _validate_aws_region_name(aws_region_name: Optional[str]) -> None:
|
||||
"""
|
||||
Validate that an AWS region name conforms to the expected format
|
||||
(lowercase alphanumerics and hyphens). Raises ValueError otherwise.
|
||||
"""
|
||||
if aws_region_name is None:
|
||||
return
|
||||
if not isinstance(aws_region_name, str) or not _VALID_AWS_REGION_PATTERN.match(
|
||||
aws_region_name
|
||||
):
|
||||
raise ValueError(
|
||||
f"Invalid AWS region format: {aws_region_name!r}. "
|
||||
"Region names must contain only lowercase letters, digits, and hyphens."
|
||||
)
|
||||
|
||||
def get_aws_region_name_for_non_llm_api_calls(
|
||||
self,
|
||||
aws_region_name: Optional[str] = None,
|
||||
|
|
@ -532,6 +559,7 @@ class BaseAWSLLM:
|
|||
|
||||
For non-llm api calls eg. Guardrails, Vector Stores we just need to check the dynamic param or env vars.
|
||||
"""
|
||||
self._validate_aws_region_name(aws_region_name)
|
||||
if aws_region_name is None:
|
||||
# check env #
|
||||
litellm_aws_region_name = get_secret("AWS_REGION_NAME", None)
|
||||
|
|
@ -549,6 +577,8 @@ class BaseAWSLLM:
|
|||
|
||||
if aws_region_name is None:
|
||||
aws_region_name = "us-west-2"
|
||||
|
||||
self._validate_aws_region_name(aws_region_name)
|
||||
return aws_region_name
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -180,6 +180,140 @@ def test_get_aws_region_name_boto3_fallback():
|
|||
mock_boto3_session.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_region",
|
||||
[
|
||||
"us-east-1@example.com/",
|
||||
"us-east-1@example.com",
|
||||
"us-east-1/path",
|
||||
"us-east-1.example.com",
|
||||
"us-east-1:8080",
|
||||
"us-east-1#fragment",
|
||||
"us-east-1?query=1",
|
||||
"us-east-1\\path",
|
||||
"US-EAST-1", # uppercase not allowed
|
||||
"us east 1", # spaces not allowed
|
||||
"", # empty string not allowed
|
||||
"us-east-1\n", # trailing newline must not slip past $
|
||||
],
|
||||
)
|
||||
def test_get_aws_region_name_rejects_malformed_region(bad_region):
|
||||
"""
|
||||
Region names are interpolated into endpoint URL templates, so any value
|
||||
containing characters that would alter URL parsing must be rejected.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid AWS region format"):
|
||||
base_aws_llm._get_aws_region_name(
|
||||
optional_params={"aws_region_name": bad_region}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"valid_region",
|
||||
[
|
||||
"us-east-1",
|
||||
"eu-west-2",
|
||||
"ap-southeast-1",
|
||||
"us-gov-west-1",
|
||||
"cn-north-1",
|
||||
"me-south-1",
|
||||
],
|
||||
)
|
||||
def test_get_aws_region_name_accepts_valid_regions(valid_region):
|
||||
"""Real AWS region formats must continue to work after the format guard."""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
result = base_aws_llm._get_aws_region_name(
|
||||
optional_params={"aws_region_name": valid_region}
|
||||
)
|
||||
assert result == valid_region
|
||||
|
||||
|
||||
def test_get_aws_region_name_rejects_malformed_region_from_env():
|
||||
"""
|
||||
A malformed AWS_REGION / AWS_REGION_NAME env value must also be rejected
|
||||
before it can flow into a URL template.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
with patch("litellm.llms.bedrock.base_aws_llm.get_secret") as mock_get_secret:
|
||||
|
||||
def side_effect(key, default=None):
|
||||
if key == "AWS_REGION_NAME":
|
||||
return "us-east-1@example.com/"
|
||||
return default
|
||||
|
||||
mock_get_secret.side_effect = side_effect
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid AWS region format"):
|
||||
base_aws_llm._get_aws_region_name(optional_params={})
|
||||
|
||||
|
||||
def test_get_aws_region_name_for_non_llm_api_calls_rejects_malformed_param():
|
||||
"""
|
||||
The non-LLM helper (used by Guardrails, Vector Stores, etc.) must validate
|
||||
a region passed in directly so it can't flow into a URL template.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid AWS region format"):
|
||||
base_aws_llm.get_aws_region_name_for_non_llm_api_calls(
|
||||
aws_region_name="us-east-1@example.com/"
|
||||
)
|
||||
|
||||
|
||||
def test_get_aws_region_name_for_non_llm_api_calls_rejects_malformed_env():
|
||||
"""
|
||||
A malformed AWS_REGION / AWS_REGION_NAME env value must be rejected on the
|
||||
non-LLM path too — Guardrails and Vector Stores read the same env vars.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
with patch("litellm.llms.bedrock.base_aws_llm.get_secret") as mock_get_secret:
|
||||
|
||||
def side_effect(key, default=None):
|
||||
if key == "AWS_REGION_NAME":
|
||||
return "us-east-1@example.com/"
|
||||
return default
|
||||
|
||||
mock_get_secret.side_effect = side_effect
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid AWS region format"):
|
||||
base_aws_llm.get_aws_region_name_for_non_llm_api_calls()
|
||||
|
||||
|
||||
def test_get_aws_region_name_for_non_llm_api_calls_accepts_valid_region():
|
||||
"""The non-LLM helper still returns valid regions unchanged."""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
assert (
|
||||
base_aws_llm.get_aws_region_name_for_non_llm_api_calls(
|
||||
aws_region_name="us-east-1"
|
||||
)
|
||||
== "us-east-1"
|
||||
)
|
||||
|
||||
|
||||
def test_get_aws_region_from_model_arn_rejects_malformed_region():
|
||||
"""
|
||||
If the region segment of a model ARN does not match the expected format,
|
||||
the helper must return None so the caller falls back to env / default.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
bad_arn = (
|
||||
"arn:aws:bedrock:us-east-1@example.com:123456789012"
|
||||
":foundation-model/anthropic.claude-3-sonnet"
|
||||
)
|
||||
assert base_aws_llm._get_aws_region_from_model_arn(bad_arn) is None
|
||||
|
||||
good_arn = (
|
||||
"arn:aws:bedrock:us-east-1:123456789012"
|
||||
":foundation-model/anthropic.claude-3-sonnet"
|
||||
)
|
||||
assert base_aws_llm._get_aws_region_from_model_arn(good_arn) == "us-east-1"
|
||||
|
||||
|
||||
def test_sign_request_with_env_var_bearer_token():
|
||||
# Create instance of actual class
|
||||
llm = BaseAWSLLM()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue