refactor: extract _get_cors_config() for testability, fix no-op CORS tests

This commit is contained in:
shreyes19 2026-04-11 22:24:04 +05:30
parent e01fe01d35
commit f686d33408
2 changed files with 117 additions and 61 deletions

View file

@ -1140,23 +1140,54 @@ async def openai_exception_handler(request: Request, exc: ProxyException):
router = APIRouter()
_cors_origins_env = os.getenv("LITELLM_CORS_ORIGINS")
if _cors_origins_env is None or _cors_origins_env.strip() == "":
origins = ["*"]
else:
origins = [o.strip() for o in _cors_origins_env.split(",") if o.strip()]
# Disable credentials by default when wildcard origins are used — combining
# allow_origins=["*"] with allow_credentials=True causes Starlette to reflect
# the incoming Origin header, allowing any site to make credentialed requests.
# Set LITELLM_CORS_ALLOW_CREDENTIALS=true to explicitly restore the old behaviour
# (e.g. for non-browser clients that relied on the Access-Control-Allow-Credentials
# header being present regardless of origin).
_cors_credentials_env = os.getenv("LITELLM_CORS_ALLOW_CREDENTIALS")
if _cors_credentials_env is not None:
allow_cors_credentials = _cors_credentials_env.strip().lower() == "true"
else:
allow_cors_credentials = "*" not in origins
def _get_cors_config(
cors_origins_env: Optional[str] = None,
cors_credentials_env: Optional[str] = None,
):
"""
Compute CORS allowed origins and credentials flag from environment variables.
Extracted into a function so it can be unit-tested without reloading the module.
Args:
cors_origins_env: Value of LITELLM_CORS_ORIGINS (defaults to os.getenv).
cors_credentials_env: Value of LITELLM_CORS_ALLOW_CREDENTIALS (defaults to os.getenv).
Returns:
Tuple[List[str], bool]: (origins, allow_credentials)
"""
_origins_raw = (
cors_origins_env
if cors_origins_env is not None
else os.getenv("LITELLM_CORS_ORIGINS")
)
if _origins_raw is None or _origins_raw.strip() == "":
computed_origins = ["*"]
else:
computed_origins = [o.strip() for o in _origins_raw.split(",") if o.strip()]
# Disable credentials by default when wildcard origins are used — combining
# allow_origins=["*"] with allow_credentials=True causes Starlette to reflect
# the incoming Origin header, allowing any site to make credentialed requests.
# Set LITELLM_CORS_ALLOW_CREDENTIALS=true to explicitly restore the old behaviour
# (e.g. for non-browser clients that relied on the Access-Control-Allow-Credentials
# header being present regardless of origin).
_credentials_raw = (
cors_credentials_env
if cors_credentials_env is not None
else os.getenv("LITELLM_CORS_ALLOW_CREDENTIALS")
)
if _credentials_raw is not None:
computed_credentials = _credentials_raw.strip().lower() == "true"
else:
computed_credentials = "*" not in computed_origins
return computed_origins, computed_credentials
origins, allow_cors_credentials = _get_cors_config()
# get current directory

View file

@ -1,38 +1,28 @@
"""
Tests for CORS configuration security fix.
Verifies that allow_credentials is automatically disabled when
allow_origins=["*"] (wildcard) to prevent credentialed cross-origin
requests from arbitrary origins.
All tests import _get_cors_config directly from proxy_server so they exercise
real production code rather than a local mirror.
"""
import pytest
def _compute_cors_config(cors_origins_env):
"""
Mirror of the CORS config logic in proxy_server.py.
Kept here so tests remain isolated from module-level side-effects.
"""
if cors_origins_env is None or cors_origins_env.strip() == "":
origins = ["*"]
else:
origins = [o.strip() for o in cors_origins_env.split(",") if o.strip()]
allow_cors_credentials = "*" not in origins
return origins, allow_cors_credentials
def test_cors_wildcard_disables_credentials():
"""should disable credentials when LITELLM_CORS_ORIGINS is not set (defaults to wildcard)."""
origins, allow_credentials = _compute_cors_config(None)
from litellm.proxy.proxy_server import _get_cors_config
origins, allow_credentials = _get_cors_config(cors_origins_env="")
assert origins == ["*"]
assert allow_credentials is False
def test_cors_empty_string_disables_credentials():
"""should disable credentials when LITELLM_CORS_ORIGINS is an empty or whitespace string."""
"""should disable credentials when LITELLM_CORS_ORIGINS is empty or whitespace."""
from litellm.proxy.proxy_server import _get_cors_config
for empty in ("", " ", "\t"):
origins, allow_credentials = _compute_cors_config(empty)
origins, allow_credentials = _get_cors_config(cors_origins_env=empty)
assert origins == ["*"], f"Expected wildcard for input {repr(empty)}"
assert (
allow_credentials is False
@ -41,15 +31,21 @@ def test_cors_empty_string_disables_credentials():
def test_cors_single_specific_origin_enables_credentials():
"""should enable credentials when a single explicit origin is configured."""
origins, allow_credentials = _compute_cors_config("https://admin.example.com")
from litellm.proxy.proxy_server import _get_cors_config
origins, allow_credentials = _get_cors_config(
cors_origins_env="https://admin.example.com"
)
assert origins == ["https://admin.example.com"]
assert allow_credentials is True
def test_cors_multiple_specific_origins_enables_credentials():
"""should enable credentials and correctly parse comma-separated origins."""
origins, allow_credentials = _compute_cors_config(
"https://app.example.com, https://admin.example.com, https://api.example.com"
from litellm.proxy.proxy_server import _get_cors_config
origins, allow_credentials = _get_cors_config(
cors_origins_env="https://app.example.com, https://admin.example.com, https://api.example.com"
)
assert origins == [
"https://app.example.com",
@ -61,51 +57,80 @@ def test_cors_multiple_specific_origins_enables_credentials():
def test_cors_wildcard_string_in_env_disables_credentials():
"""should disable credentials when LITELLM_CORS_ORIGINS is explicitly set to '*'."""
origins, allow_credentials = _compute_cors_config("*")
from litellm.proxy.proxy_server import _get_cors_config
origins, allow_credentials = _get_cors_config(cors_origins_env="*")
assert "*" in origins
assert allow_credentials is False
def test_cors_origins_strips_whitespace():
"""should strip surrounding whitespace from each origin entry."""
origins, _ = _compute_cors_config(" https://a.com , https://b.com ")
from litellm.proxy.proxy_server import _get_cors_config
origins, _ = _get_cors_config(
cors_origins_env=" https://a.com , https://b.com "
)
assert origins == ["https://a.com", "https://b.com"]
def test_cors_origins_skips_blank_entries():
"""should skip blank entries caused by trailing/double commas."""
origins, allow_credentials = _compute_cors_config("https://a.com,,https://b.com,")
from litellm.proxy.proxy_server import _get_cors_config
origins, allow_credentials = _get_cors_config(
cors_origins_env="https://a.com,,https://b.com,"
)
assert origins == ["https://a.com", "https://b.com"]
assert allow_credentials is True
def test_cors_explicit_credentials_override_true(monkeypatch):
"""should allow LITELLM_CORS_ALLOW_CREDENTIALS=true to explicitly re-enable
credentials even when wildcard origins are used (opt-in for existing deployments).
"""
monkeypatch.setenv("LITELLM_CORS_ALLOW_CREDENTIALS", "true")
_cors_credentials_env = "true"
_cors_credentials_env = _cors_credentials_env.strip().lower() == "true"
assert _cors_credentials_env is True
def test_cors_explicit_credentials_true_overrides_wildcard():
"""should enable credentials when LITELLM_CORS_ALLOW_CREDENTIALS=true even
if wildcard origins are in use (opt-in for existing deployments)."""
from litellm.proxy.proxy_server import _get_cors_config
origins, allow_credentials = _get_cors_config(
cors_origins_env="",
cors_credentials_env="true",
)
assert "*" in origins
assert allow_credentials is True
def test_cors_explicit_credentials_override_false(monkeypatch):
"""should allow LITELLM_CORS_ALLOW_CREDENTIALS=false to explicitly disable
credentials even when specific origins are configured."""
monkeypatch.setenv("LITELLM_CORS_ALLOW_CREDENTIALS", "false")
_cors_credentials_env = "false"
result = _cors_credentials_env.strip().lower() == "true"
assert result is False
def test_cors_explicit_credentials_false_overrides_specific_origins():
"""should disable credentials when LITELLM_CORS_ALLOW_CREDENTIALS=false even
if specific origins are configured."""
from litellm.proxy.proxy_server import _get_cors_config
origins, allow_credentials = _get_cors_config(
cors_origins_env="https://admin.example.com",
cors_credentials_env="false",
)
assert origins == ["https://admin.example.com"]
assert allow_credentials is False
def test_cors_explicit_credentials_case_insensitive():
"""should accept TRUE/FALSE case-insensitively for LITELLM_CORS_ALLOW_CREDENTIALS."""
from litellm.proxy.proxy_server import _get_cors_config
_, allow_true = _get_cors_config(cors_origins_env="", cors_credentials_env="TRUE")
_, allow_false = _get_cors_config(
cors_origins_env="https://x.com", cors_credentials_env="FALSE"
)
assert allow_true is True
assert allow_false is False
def test_proxy_server_cors_invariant():
"""should verify that proxy_server.allow_cors_credentials is always consistent
with proxy_server.origins — catches any future drift between the two variables."""
import litellm.proxy.proxy_server as proxy_server
# When LITELLM_CORS_ALLOW_CREDENTIALS is not explicitly set, the invariant must hold
"""should verify that proxy_server module-level origins and allow_cors_credentials
are consistent — catches any future drift in the module-level call to _get_cors_config.
"""
import os
import litellm.proxy.proxy_server as proxy_server
if os.getenv("LITELLM_CORS_ALLOW_CREDENTIALS") is None:
assert proxy_server.allow_cors_credentials == (
"*" not in proxy_server.origins