fix: validate aws region name

This commit is contained in:
Yassin Kortam 2026-04-30 11:22:47 -07:00
parent 609454d0f1
commit d47948ab23
2 changed files with 164 additions and 0 deletions

View file

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

View file

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