diff --git a/pyproject.toml b/pyproject.toml index 92ace8df..150764ef 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -16,7 +16,7 @@ dependencies = [ "fastapi>=0.115.5", "httpx[http2]>=0.28.0", "jinja2>=3.1.6", - "pydantic-settings>=2.7.0", + "pydantic-settings>=2.8.0", "pyjwt>=2.10.1", "starlette>=1.0.1", "starlette-cramjam>=0.4.0", diff --git a/src/stac_auth_proxy/config.py b/src/stac_auth_proxy/config.py index ea140560..9e3a7778 100644 --- a/src/stac_auth_proxy/config.py +++ b/src/stac_auth_proxy/config.py @@ -47,11 +47,12 @@ def __call__(self): class CorsSettings(BaseModel): """CORS configuration settings.""" - allow_origins: Sequence[str] = ["*"] - allow_methods: Sequence[str] = ["*"] - allow_headers: Sequence[str] = ["*"] + # NoDecode: accept the documented comma-separated env form (see root_path_skip_prefixes) + allow_origins: Annotated[Sequence[str], NoDecode] = ["*"] + allow_methods: Annotated[Sequence[str], NoDecode] = ["*"] + allow_headers: Annotated[Sequence[str], NoDecode] = ["*"] allow_credentials: bool = True - expose_headers: Sequence[str] = [] + expose_headers: Annotated[Sequence[str], NoDecode] = [] max_age: int = 600 @field_validator( @@ -63,10 +64,8 @@ class CorsSettings(BaseModel): ) @classmethod def parse_list(cls, v) -> Sequence[str] | None: - """Parse a comma-separated string into a list.""" - if isinstance(v, str): - return [s.strip() for s in v.split(",") if s.strip()] - return v + """Parse a comma-separated or JSON-list string into a list.""" + return str2list(v) class Settings(BaseSettings): @@ -135,6 +134,9 @@ class Settings(BaseSettings): model_config = SettingsConfigDict( env_nested_delimiter="_", + # Split only once so CORS_ALLOW_ORIGINS -> cors.allow_origins, + # not cors.allow.origins (which is silently ignored). + env_nested_max_split=1, ) @model_validator(mode="before") diff --git a/tests/test_config.py b/tests/test_config.py index f0ac5334..93da5e28 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -151,3 +151,36 @@ def test_cors_model_config(): ] assert cors_settings.allow_methods == ["GET", "POST"] assert cors_settings.allow_headers == ["Authorization", "Content-Type"] + + +def test_cors_env_vars(monkeypatch): + """The documented CORS_* env vars are honored alongside other nested settings.""" + monkeypatch.setenv("UPSTREAM_URL", "http://upstream") + monkeypatch.setenv("OIDC_DISCOVERY_URL", "http://oidc/.well-known/x") + monkeypatch.setenv("CORS_ALLOW_ORIGINS", "https://a.com,https://b.com") + monkeypatch.setenv("CORS_ALLOW_CREDENTIALS", "false") + monkeypatch.setenv("CORS_MAX_AGE", "10") + monkeypatch.setenv("ITEMS_FILTER_CLS", "stac_auth_proxy.filters:Template") + monkeypatch.setenv("ITEMS_FILTER_ARGS", '["true"]') + settings = Settings() + assert list(settings.cors.allow_origins) == ["https://a.com", "https://b.com"] + assert settings.cors.allow_credentials is False + assert settings.cors.max_age == 10 + assert settings.items_filter.cls == "stac_auth_proxy.filters:Template" + assert list(settings.items_filter.args) == ["true"] + + +def test_cors_json_env_var(monkeypatch): + """The CORS={...} JSON env var form is honored.""" + monkeypatch.setenv("UPSTREAM_URL", "http://upstream") + monkeypatch.setenv("OIDC_DISCOVERY_URL", "http://oidc/.well-known/x") + monkeypatch.setenv("CORS", '{"allow_origins": ["https://a.com"]}') + assert list(Settings().cors.allow_origins) == ["https://a.com"] + + +def test_cors_env_var_json_list(monkeypatch): + """A CORS_* list env var also accepts the JSON-list form.""" + monkeypatch.setenv("UPSTREAM_URL", "http://upstream") + monkeypatch.setenv("OIDC_DISCOVERY_URL", "http://oidc/.well-known/x") + monkeypatch.setenv("CORS_ALLOW_ORIGINS", '["https://a.com", "https://b.com"]') + assert list(Settings().cors.allow_origins) == ["https://a.com", "https://b.com"] diff --git a/uv.lock b/uv.lock index be5286bb..eaaa9a5b 100644 --- a/uv.lock +++ b/uv.lock @@ -2075,7 +2075,7 @@ requires-dist = [ { name = "mkdocs-material", extras = ["imaging"], marker = "extra == 'docs'", specifier = ">=9.6.16" }, { name = "mkdocstrings", extras = ["python"], marker = "extra == 'docs'", specifier = ">=0.30.0" }, { name = "prometheus-fastapi-instrumentator", marker = "extra == 'metrics'", specifier = ">=8.0.2" }, - { name = "pydantic-settings", specifier = ">=2.7.0" }, + { name = "pydantic-settings", specifier = ">=2.8.0" }, { name = "pyjwt", specifier = ">=2.10.1" }, { name = "starlette", specifier = ">=1.0.1" }, { name = "starlette-cramjam", specifier = ">=0.4.0" },