mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
refactor: extract _get_cors_config() for testability, fix no-op CORS tests
This commit is contained in:
parent
e01fe01d35
commit
f686d33408
2 changed files with 117 additions and 61 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue