chore(proxy): extend SSRF/destination guards to provider endpoint overrides

The URL validation and destination-change credential clearing only covered api_base/base_url, but provider-specific endpoint URLs (aws_bedrock_runtime_endpoint, aws_sts_endpoint) are equally caller-controlled: a non-admin team model could point one at an internal/metadata address while reusing stored AWS credentials. Validate those endpoints too and treat them as destination overrides (model management + connection-test guard).

Also make the os.environ/ reference scan iterative (stack-based) instead of recursive to satisfy the recursive-function check.
This commit is contained in:
user 2026-05-31 06:37:44 +00:00
parent dc34d1d8cd
commit f9cd054445
No known key found for this signature in database
3 changed files with 30 additions and 9 deletions

View file

@ -88,6 +88,8 @@ _HEALTH_CREDENTIAL_FIELDS = (
_HEALTH_DESTINATION_FIELDS = (
"api_base",
"base_url",
"aws_bedrock_runtime_endpoint",
"aws_sts_endpoint",
"api_version",
"vertex_location",
"vertex_project",

View file

@ -77,7 +77,16 @@ _CREDENTIAL_LITELLM_PARAMS = (
"aws_session_token",
"vertex_credentials",
)
_DESTINATION_LITELLM_PARAMS = ("api_base", "base_url", "custom_llm_provider")
# Caller-controlled endpoint URLs that must be SSRF-validated on write.
_URL_LITELLM_PARAMS = (
"api_base",
"base_url",
"aws_bedrock_runtime_endpoint",
"aws_sts_endpoint",
)
# Destination fields whose change must drop an inherited credential (URLs above
# plus the provider selector, which isn't a URL so isn't SSRF-validated).
_DESTINATION_LITELLM_PARAMS = _URL_LITELLM_PARAMS + ("custom_llm_provider",)
def _field_explicitly_set(model: Optional[BaseModel], field: str) -> bool:
@ -92,7 +101,7 @@ def _validate_model_url_params(litellm_params: dict) -> None:
with user_url_allowed_hosts as the escape hatch for internal endpoints."""
if not getattr(litellm, "user_url_validation", False):
return
for url_field in ("api_base", "base_url"):
for url_field in _URL_LITELLM_PARAMS:
url_value = litellm_params.get(url_field)
if not url_value or not isinstance(url_value, str):
continue
@ -151,13 +160,18 @@ def _is_pricing_field(field: str) -> bool:
def _contains_env_reference(value: object) -> bool:
"""True if any (possibly nested) string is an `os.environ/` reference, which
resolves to a server-side environment secret at call time."""
if isinstance(value, str):
return value.startswith("os.environ/")
if isinstance(value, dict):
return any(_contains_env_reference(v) for v in value.values())
if isinstance(value, list):
return any(_contains_env_reference(v) for v in value)
resolves to a server-side environment secret at call time. Iterative (stack)
rather than recursive, matching _reject_os_environ_references."""
stack: List[object] = [value]
while stack:
item = stack.pop()
if isinstance(item, str):
if item.startswith("os.environ/"):
return True
elif isinstance(item, dict):
stack.extend(item.values())
elif isinstance(item, list):
stack.extend(item)
return False

View file

@ -2044,6 +2044,11 @@ class TestModelMgmtAuthzHardening:
_validate_model_url_params(
{"api_base": "http://169.254.169.254/latest/meta-data/"}
)
# Provider-specific endpoint URLs are guarded too, not just api_base.
with pytest.raises(ProxyException):
_validate_model_url_params(
{"aws_bedrock_runtime_endpoint": "http://169.254.169.254/"}
)
# Honors the opt-out toggle (internal endpoints / Ollama).
with (
patch.object(litellm, "user_url_validation", False),