diff --git a/.env.example b/.env.example index 51eaa99..ed33d4b 100644 --- a/.env.example +++ b/.env.example @@ -20,6 +20,20 @@ AZURE_COMMUNICATION_SERVICE_CONNECTION_STRING="" # ] CLIENT_APPLICATION_DETAILS_FILE="path/to/client/applications/metadata/file.json" +# This file lists the origins that browser based applications, such as the IATI +# Dashboard, may call this API from. It should be a JSON list of origins. Each origin +# must be of the form scheme://host[:port] - lower case, with no trailing slash and no +# path - because the browser sends the Origin header in exactly that form and it is +# matched by string equality. Wildcards are not allowed. +# Example: +# [ +# "https://dev-dashboard.iatistandard.org", +# "http://localhost:5173" +# ] +# If this variable is omitted the API allows no cross-origin requests at all, which is +# the correct setting for a deployment that only serves server-to-server clients. +# CORS_ALLOWED_ORIGINS_FILE="path/to/cors/allowed/origins/file.json" + DATA_REGISTRY_SUITECRM_API_URL = "https://base_url_of_suitecrm_instance.org" DATA_REGISTRY_SUITECRM_CLIENT_ID = "RYD's client ID here" diff --git a/.gitignore b/.gitignore index 7418cb1..ae2cc1f 100644 --- a/.gitignore +++ b/.gitignore @@ -21,3 +21,5 @@ __pycache__/ # all key files *.pem + +/cors-allowed-origins.json diff --git a/CHANGELOG.md b/CHANGELOG.md index 432c0f9..4e13fa0 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,20 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added +### Changed + +### Deprecated + +### Fixed + +### Removed + +### Security + +## [0.3.9] - 2026-09-22 + +### Added + - Sentry error monitoring and request tracing, initialised in `src/main.py` before the FastAPI application is created. Configured by the optional `SENTRY_DSN`, `SENTRY_ENVIRONMENT`, and `SENTRY_TRACES_SAMPLE_RATE` environment variables; when no @@ -19,16 +33,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 entirely, since its records are written at `CRITICAL` and carry the contents of the `Authorization` header, and a failed startup is now logged rather than printed so that it is reported as well. - -### Changed - -### Deprecated - -### Fixed - -### Removed - -### Security +- CORS support, so that browser based applications such as the IATI Dashboard can call the + API. The origins that are allowed to do so are listed in a JSON file named by the new, + optional, `CORS_ALLOWED_ORIGINS_FILE` environment variable. If it is not set, no + cross-origin requests are allowed. ## [0.3.8] - 2026-07-01 diff --git a/README.md b/README.md index 02b44a9..a064b9d 100644 --- a/README.md +++ b/README.md @@ -46,6 +46,7 @@ The application is configured using a set of environment variables in a `.env` f | `SENTRY_DSN` | DSN of the Sentry project to report errors to. **Optional** — leave empty or unset and Sentry is not initialised, and the application runs normally. | | `SENTRY_ENVIRONMENT` | Environment name that events are tagged with in Sentry, e.g. `"dev"` or `"prod"`. Defaults to `"local-development"` rather than the SDK's own default of `"production"`, so that an unconfigured environment can never be mistaken for the live one. | | `SENTRY_TRACES_SAMPLE_RATE` | Proportion of requests traced for performance monitoring, `0.0`–`1.0`. Defaults to `1.0`; lower it (e.g. `0.1`) if trace volume becomes a problem. | +| `CORS_ALLOWED_ORIGINS_FILE` | Path to a JSON file listing the origins that browser based applications, such as the IATI Dashboard, may call this API from. Omit it to allow no cross-origin requests at all. See `cors-allowed-origins.example.json` and the notes in `.env.example` for the required format. | ### Error monitoring with Sentry @@ -169,6 +170,7 @@ Running the API in production can be run directly from its Docker container with ``` docker run \ --mount type=bind,source=.env,target=/api/.env,readonly \ + --mount type=bind,source=cors-allowed-origins.json,target=/api/cors-allowed-origins.json,readonly \ --mount type=bind,source=./logs/,target=/api/logs/ \ --mount type=bind,source=./keys/,target=/api/keys/,readonly \ -p 8000:8000 \ @@ -178,12 +180,13 @@ docker run \ This performs the following setup: * `--mount type=bind,source=.env,target=/api/.env,readonly`: shares `.env` configuration file on the host with the container. +* `--mount type=bind,source=cors-allowed-origins.json,target=/api/cors-allowed-origins.json,readonly`: shares the CORS allowed origins file on the host with the container. Mounting it, rather than baking it into the image, means origins can be added by editing the file and restarting the container. Omit this mount if the deployment does not serve any browser based clients. * `--mount type=bind,source=./logs/,target=/api/logs/`: allows the container to write logs to the /logs directory on the host. * `--mount type=bind,source=./keys/,target=/api/keys/,readonly`: shares the `/keys` directory on the host with the container so public and private keys can be used by the container. * `-p 8000:8000`: shares port 8000 for API traffic. * `-p 9000:9000`: shares port 9000 for Prometheus metrics (assuming that this is the port as specified by the `.env` file.) -Care should be taken to make sure that the `.env` variables match the log (`APP_LOG_PATH` and `AUDIT_LOG_PATH`) and key directories (`AUDIT_LOG_PUBLIC_KEY_PATH`) and the Prometheus metric port (`PROMETHEUS_PORT`). +Care should be taken to make sure that the `.env` variables match the log (`APP_LOG_PATH` and `AUDIT_LOG_PATH`) and key directories (`AUDIT_LOG_PUBLIC_KEY_PATH`) and the Prometheus metric port (`PROMETHEUS_PORT`). If CORS is in use, `CORS_ALLOWED_ORIGINS_FILE` must also match the target path of the mount above. ## Development diff --git a/cors-allowed-origins.example.json b/cors-allowed-origins.example.json new file mode 100644 index 0000000..94d77b1 --- /dev/null +++ b/cors-allowed-origins.example.json @@ -0,0 +1,3 @@ +[ + "https://dev-dashboard.iatistandard.org" +] diff --git a/pyproject.toml b/pyproject.toml index 24bf132..5d09562 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "register-your-data-api" -version = "0.3.8" +version = "0.3.9" requires-python = ">= 3.12.11" readme = "README.md" authors = [{name="IATI Secretariat", email="support@iatistandard.org"}] diff --git a/src/main.py b/src/main.py index 61344be..d20a99c 100644 --- a/src/main.py +++ b/src/main.py @@ -8,14 +8,27 @@ import prometheus_client from fastapi import FastAPI +import register_your_data_api.cors as cors import register_your_data_api.exception_handlers import register_your_data_api.util as util from register_your_data_api.routers import datasets, discoverable_reporting_orgs, misc, reporting_orgs, users from register_your_data_api.sentry import setup_sentry +# Assigned at the bottom of this module and acted on in prod_lifespan; the comment on the +# assignment explains why those are two different places. +_cors_configuration_error: RuntimeError | None = None + @contextlib.asynccontextmanager async def prod_lifespan(app: FastAPI) -> AsyncIterator[None]: + if _cors_configuration_error is not None: + # Logged rather than printed, and with the exception attached, for the same reasons + # as the context failure below. + logging.getLogger(__name__).error( + "Could not initialise application - error configuring CORS", exc_info=_cors_configuration_error + ) + sys.exit("Could not startup") + try: context = util.Context() context.setup() @@ -53,4 +66,14 @@ def add_routers_and_general_exception_handling(app: FastAPI) -> None: app = FastAPI(title="Register Your Data", lifespan=prod_lifespan, redirect_slashes=False) +# Middleware has to be registered before the application starts, so this runs at import. +# Exiting here would take down every importer of this module - which includes the whole +# test suite, via tests/helpers/mocking.py - and SystemExit during collection aborts pytest +# without reporting a cause. The failure is therefore carried into prod_lifespan, where +# the equivalent context failure is already reported. +try: + cors.add_cors_middleware(app, cors.load_allowed_origins()) +except RuntimeError as err: + _cors_configuration_error = err + add_routers_and_general_exception_handling(app) diff --git a/src/register_your_data_api/config.py b/src/register_your_data_api/config.py new file mode 100644 index 0000000..cc21e99 --- /dev/null +++ b/src/register_your_data_api/config.py @@ -0,0 +1,32 @@ +"""Configuration read straight from the environment. + +Both Sentry and CORS have to be configured before the FastAPI application object is +created, which is earlier than the ``Context`` (created in the application lifespan) +exists. They therefore read the environment here rather than going through ``Context``, +using the same precedence that ``Context`` uses: values from the ``.env`` file in the +current directory, overridden by real environment variables. +""" + +import os + +import dotenv + + +def get_environment_config(env_file: str = ".env") -> dict[str, str]: + """Read configuration with the same precedence as Context: the env file, then os.environ. + + Parameters + ---------- + env_file : str + Path to the environment file. A missing file is not an error, in which case only + os.environ is used. + + Returns + ------- + dict[str, str] + """ + + env: dict[str, str] = {key: value for key, value in dotenv.dotenv_values(env_file).items() if value is not None} + env.update(os.environ) + + return env diff --git a/src/register_your_data_api/cors.py b/src/register_your_data_api/cors.py new file mode 100644 index 0000000..5105f97 --- /dev/null +++ b/src/register_your_data_api/cors.py @@ -0,0 +1,132 @@ +"""Cross-Origin Resource Sharing (CORS) configuration. + +The allowed origins have to be read before the FastAPI app object is handed to uvicorn, +because Starlette refuses to add middleware once an application has started - and that +includes the lifespan in which the Context is built. This module therefore reads the +environment directly rather than going through Context. +""" + +import json +import re +from typing import Final + +from fastapi import FastAPI +from fastapi.middleware.cors import CORSMiddleware + +from .config import get_environment_config + +CORS_ALLOWED_ORIGINS_FILE: Final[str] = "CORS_ALLOWED_ORIGINS_FILE" + +# The verbs the API serves, and no others. +ALLOWED_METHODS: Final[list[str]] = ["GET", "POST", "PATCH", "PUT", "DELETE"] + +ALLOWED_HEADERS: Final[list[str]] = ["Authorization", "Content-Type"] + +# Scheme, host and an optional port, and nothing else. Browsers send an Origin header in +# exactly this form and CORSMiddleware compares it by string equality, so an entry with a +# trailing slash, a path, or an upper case host would silently never match. A wildcard is +# rejected by this pattern too: a single "*" entry would switch CORSMiddleware into +# allow-all mode, which is precisely what we are avoiding. +_ORIGIN_PATTERN: Final[re.Pattern[str]] = re.compile(r"https?://[a-z0-9.-]+(:[0-9]+)?") + + +def load_allowed_origins(env: dict[str, str] | None = None) -> list[str]: + """Load the origins the API will accept cross-origin requests from. + + Parameters + ---------- + env : dict[str, str] | None + Environment variables to read the configuration from. Defaults to the values + returned by get_environment_config(). + + Returns + ------- + list[str] + The configured origins, or an empty list if CORS_ALLOWED_ORIGINS_FILE is not set, + in which case no cross-origin requests are allowed at all. + + Raises + ------ + RuntimeError + If CORS_ALLOWED_ORIGINS_FILE is set but the file it names cannot be read, is not + valid JSON, or does not contain a list of well formed origins. + """ + + if env is None: + env = get_environment_config() + + filename = env.get(CORS_ALLOWED_ORIGINS_FILE, "").strip() + + if not filename: + return [] + + origins = _read_origins_file(filename) + + if not isinstance(origins, list): + raise RuntimeError(f"CORS allowed origins file {filename} must contain a JSON list of origins") + + invalid_origins = [origin for origin in origins if not _is_valid_origin(origin)] + + if invalid_origins: + raise RuntimeError( + f"CORS allowed origins file {filename} contains invalid origins: {invalid_origins}. Each must be " + "of the form scheme://host[:port], lower case, with no trailing slash and no path. Wildcards are " + "not allowed." + ) + + return origins + + +def add_cors_middleware(app: FastAPI, allowed_origins: list[str]) -> None: + """Add CORS middleware to a FastAPI app instance. + + Preflight OPTIONS requests are answered by this middleware before the router, and so + before the per-endpoint Security() dependencies, which is what stops browsers being + given a 401 for a preflight. + + Note that CORS is enforced by the browser, not by us: a request from an origin that is + not allowed is still served as normal, it just comes back without the headers the + browser needs in order to hand the response to the calling script. + + An empty list registers no middleware at all, so a deployment that has not opted into + CORS is left exactly as it was. Registering the middleware with an empty allowlist + would not be equivalent: it would still intercept preflights, answering OPTIONS with a + 400 where the router would previously have returned 405. + + Parameters + ---------- + app : FastAPI + allowed_origins : list[str] + Origins to accept cross-origin requests from. An empty list disables CORS. + """ + + if not allowed_origins: + return + + app.add_middleware( + CORSMiddleware, + allow_origins=allowed_origins, + allow_credentials=False, + allow_methods=ALLOWED_METHODS, + allow_headers=ALLOWED_HEADERS, + ) + + +def _read_origins_file(filename: str) -> object: + """Read and parse the JSON origins file, translating any failure into a RuntimeError.""" + + try: + with open(filename, "r") as file: + return json.load(file) + except OSError as err: + raise RuntimeError(f"Could not read CORS allowed origins file {filename}: {err}") from err + except ValueError as err: + # Covers json.JSONDecodeError and the UnicodeDecodeError raised when the file is + # not readable as text; both are ValueError subclasses. + raise RuntimeError(f"CORS allowed origins file {filename} is not valid JSON: {err}") from err + + +def _is_valid_origin(origin: object) -> bool: + """Check that a single entry from the origins file is an origin we can match against.""" + + return isinstance(origin, str) and _ORIGIN_PATTERN.fullmatch(origin) is not None diff --git a/src/register_your_data_api/sentry.py b/src/register_your_data_api/sentry.py index 45186ff..8becd7e 100644 --- a/src/register_your_data_api/sentry.py +++ b/src/register_your_data_api/sentry.py @@ -11,14 +11,14 @@ """ import importlib.metadata -import os from typing import Final -import dotenv import sentry_sdk from sentry_sdk.integrations.logging import ignore_logger, ignore_logger_for_sentry_logs from sentry_sdk.scrubber import DEFAULT_DENYLIST, EventScrubber +from .config import get_environment_config + # Sentry's SDK sends no traces at all unless a sample rate is set, so a default is # supplied here rather than leaving it to the SDK DEFAULT_TRACES_SAMPLE_RATE: Final[float] = 1.0 @@ -52,14 +52,6 @@ ] -def get_environment_config() -> dict[str, str]: - """Reads configuration with the same precedence as ``Context``: ``.env`` then os.environ.""" - - env: dict[str, str] = {k: v for k, v in dotenv.dotenv_values(".env").items() if v is not None} - env.update(os.environ) - return env - - def get_release() -> str | None: """Builds the Sentry release identifier from the installed package version.""" diff --git a/tests/helpers/mocking.py b/tests/helpers/mocking.py index 515fa3a..5d9a9c3 100644 --- a/tests/helpers/mocking.py +++ b/tests/helpers/mocking.py @@ -25,6 +25,7 @@ ) from register_your_data_api.auth.fga.models import FineGrainedAuthorisationRole from register_your_data_api.client_application_details_provider import ClientApplicationDetails +from register_your_data_api.cors import add_cors_middleware from tests.helpers.keys import KeyDict from ..helpers import prom @@ -233,7 +234,21 @@ def _create_test_key(self, key_name: str) -> None: public_key, ) - def get_test_app(self) -> FastAPI: + def get_test_app(self, cors_allowed_origins: list[str] | None = None) -> FastAPI: + """Build (or return the already built) test app. + + Parameters + ---------- + cors_allowed_origins : list[str] | None + Origins to configure CORS middleware with. By default no CORS middleware is + added at all, matching an API deployed without CORS_ALLOWED_ORIGINS_FILE set. + The app is built once and then cached, so the first call wins - a later call + passing different origins has no effect. + + Returns + ------- + FastAPI + """ @contextlib.asynccontextmanager async def test_lifespan(app: FastAPI) -> AsyncIterator[None]: @@ -247,6 +262,8 @@ async def test_lifespan(app: FastAPI) -> AsyncIterator[None]: if not self._app_is_created: self._app = FastAPI(title="Register Your Data", lifespan=test_lifespan) + if cors_allowed_origins is not None: + add_cors_middleware(self._app, cors_allowed_origins) add_routers_and_general_exception_handling(self._app) self._app_is_created = True diff --git a/tests/integration/test_cors.py b/tests/integration/test_cors.py new file mode 100644 index 0000000..b4cef1c --- /dev/null +++ b/tests/integration/test_cors.py @@ -0,0 +1,284 @@ +"""Tests for the CORS support that lets browser based clients call the API. + +CORS is enforced by the browser, not by the API. These tests therefore check which +Access-Control-* headers come back, rather than whether a request succeeds: an origin we +do not allow is still served as normal, it just comes back without the headers the browser +needs in order to hand the response to the calling script. +""" + +import importlib +import json +from pathlib import Path + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from starlette.middleware.cors import CORSMiddleware + +from register_your_data_api.cors import CORS_ALLOWED_ORIGINS_FILE, add_cors_middleware + +from ..helpers.mocking import MockedAppAndContext + +ALLOWED_ORIGIN = "https://client.example.org" + +DISALLOWED_ORIGIN = "https://other.example.com" + +TOOLS_ENDPOINT = "/api/v1/tools" + + +def preflight_headers(origin: str, method: str = "GET", request_headers: str = "authorization") -> dict[str, str]: + """Build the headers a browser sends when it preflights a cross-origin request.""" + + return { + "Origin": origin, + "Access-Control-Request-Method": method, + "Access-Control-Request-Headers": request_headers, + } + + +def test_preflight_succeeds_without_credentials() -> None: + """A browser preflight carries no Authorization header, so it must not be made to authenticate.""" + + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[ALLOWED_ORIGIN]) + + with TestClient(fastAPIapp) as client: + response = client.options(TOOLS_ENDPOINT, headers=preflight_headers(ALLOWED_ORIGIN)) + + assert response.status_code == 200 + assert response.headers["access-control-allow-origin"] == ALLOWED_ORIGIN + allowed_methods = {method.strip() for method in response.headers["access-control-allow-methods"].split(",")} + assert allowed_methods == {"GET", "POST", "PATCH", "PUT", "DELETE"} + assert "Authorization" in response.headers["access-control-allow-headers"] + assert "Content-Type" in response.headers["access-control-allow-headers"] + assert "access-control-allow-credentials" not in response.headers + + +def test_authenticated_request_from_an_allowed_origin() -> None: + """A browser client sends a Bearer token on every call, and must be able to read the response.""" + + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[ALLOWED_ORIGIN]) + + with TestClient(fastAPIapp) as client: + response = client.get( + TOOLS_ENDPOINT, + headers={"Origin": ALLOWED_ORIGIN, **appAndContext.get_valid_authorization_header(0)}, + ) + + assert response.status_code == 200 + assert response.headers["access-control-allow-origin"] == ALLOWED_ORIGIN + assert "Origin" in response.headers["vary"] + assert "data" in response.json() + + +def test_preflight_from_a_disallowed_origin_is_refused() -> None: + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[ALLOWED_ORIGIN]) + + with TestClient(fastAPIapp) as client: + response = client.options(TOOLS_ENDPOINT, headers=preflight_headers(DISALLOWED_ORIGIN)) + + assert "access-control-allow-origin" not in response.headers + + +def test_request_from_a_disallowed_origin_is_still_served() -> None: + """CORS is not server side access control: the response is produced, the browser blocks it.""" + + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[ALLOWED_ORIGIN]) + + with TestClient(fastAPIapp) as client: + response = client.get( + TOOLS_ENDPOINT, + headers={"Origin": DISALLOWED_ORIGIN, **appAndContext.get_valid_authorization_header(0)}, + ) + + assert response.status_code == 200 + assert "data" in response.json() + assert "access-control-allow-origin" not in response.headers + + +def test_request_without_an_origin_is_unaffected() -> None: + """Server to server clients send no Origin header and must see no change at all.""" + + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[ALLOWED_ORIGIN]) + + with TestClient(fastAPIapp) as client: + response = client.get(TOOLS_ENDPOINT, headers=appAndContext.get_valid_authorization_header(0)) + + assert response.status_code == 200 + assert "data" in response.json() + assert "access-control-allow-origin" not in response.headers + assert "vary" not in response.headers + + +def test_unauthenticated_request_carries_cors_headers() -> None: + """A browser client needs to be able to read a 401 in order to handle an expired token.""" + + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[ALLOWED_ORIGIN]) + + with TestClient(fastAPIapp) as client: + response = client.get(TOOLS_ENDPOINT, headers={"Origin": ALLOWED_ORIGIN}) + + assert response.status_code == 401 + assert response.headers["access-control-allow-origin"] == ALLOWED_ORIGIN + + +def test_error_response_carries_cors_headers() -> None: + """Handled errors travel through ExceptionMiddleware, which sits inside CORSMiddleware. + + Note that a genuinely unhandled exception would NOT carry these headers: that 500 is + produced by ServerErrorMiddleware, which Starlette places outside CORSMiddleware. + """ + + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[ALLOWED_ORIGIN]) + + # get_valid_authorization_header(6) gets Person Seven, a failure case: see mocking.py. + with TestClient(fastAPIapp) as client: + response = client.get( + TOOLS_ENDPOINT, + headers={"Origin": ALLOWED_ORIGIN, **appAndContext.get_valid_authorization_header(6)}, + ) + + assert response.status_code == 500 + assert response.headers["access-control-allow-origin"] == ALLOWED_ORIGIN + + +@pytest.mark.parametrize("method", ["GET", "POST", "PATCH", "PUT", "DELETE"]) +def test_preflight_succeeds_for_every_method_the_api_serves(method: str) -> None: + """These are the five verbs the routers actually serve, so a preflight for any of them + has to succeed or that part of the API is unreachable from a browser. + """ + + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[ALLOWED_ORIGIN]) + + with TestClient(fastAPIapp) as client: + response = client.options(TOOLS_ENDPOINT, headers=preflight_headers(ALLOWED_ORIGIN, method=method)) + + assert response.status_code == 200 + assert response.headers["access-control-allow-origin"] == ALLOWED_ORIGIN + + +def test_preflight_for_a_method_that_is_not_allowed() -> None: + """The allowlist is a real allowlist: a verb outside it is still refused. + + HEAD is the useful case to pin. None of the API's own endpoints serve it, because + FastAPI - unlike bare Starlette, which adds HEAD alongside GET - does not do so for an + APIRoute, so HEAD is correctly absent from ALLOWED_METHODS. FastAPI's own /docs, + /redoc and /openapi.json routes are plain Starlette Routes and do answer HEAD, but + they are documentation pages rather than endpoints a browser client calls, so they are + not a reason to widen the list. + """ + + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[ALLOWED_ORIGIN]) + + with TestClient(fastAPIapp) as client: + response = client.options(TOOLS_ENDPOINT, headers=preflight_headers(ALLOWED_ORIGIN, method="HEAD")) + + # A preflight only succeeds on a 2xx, so a refusal status is what makes the browser + # block the request. The exact code and body are Starlette's choice, so assert on + # the refusal and on the verb being unadvertised rather than on its wording. + assert not response.is_success + assert "HEAD" not in response.headers["access-control-allow-methods"] + + +def test_preflight_for_a_header_that_is_not_allowed() -> None: + """A request header the API does not permit is refused, the same way a verb is.""" + + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[ALLOWED_ORIGIN]) + + with TestClient(fastAPIapp) as client: + response = client.options( + TOOLS_ENDPOINT, headers=preflight_headers(ALLOWED_ORIGIN, request_headers="X-Custom-Thing") + ) + + assert not response.is_success + assert "x-custom-thing" not in response.headers["access-control-allow-headers"].lower() + + +def test_app_without_cors_configured_emits_no_cors_headers() -> None: + """An API deployed without CORS_ALLOWED_ORIGINS_FILE set behaves exactly as it did before. + + The empty list is what production actually builds when the variable is unset, so that + is what this test passes - not None, which would exercise a configuration the deployed + application never has. + """ + + appAndContext = MockedAppAndContext() + + fastAPIapp = appAndContext.get_test_app(cors_allowed_origins=[]) + + with TestClient(fastAPIapp) as client: + response = client.get( + TOOLS_ENDPOINT, + headers={"Origin": ALLOWED_ORIGIN, **appAndContext.get_valid_authorization_header(0)}, + ) + + assert response.status_code == 200 + assert "access-control-allow-origin" not in response.headers + + # No CORS middleware means preflights are not intercepted, so OPTIONS still falls + # through to the router's 405 exactly as it did before this feature existed. + preflight = client.options(TOOLS_ENDPOINT, headers=preflight_headers(ALLOWED_ORIGIN)) + + assert preflight.status_code == 405 + + +def test_cors_middleware_is_registered_only_when_origins_are_configured() -> None: + """The empty case must register nothing, or an opted-out deployment still gets CORS.""" + + configured = FastAPI() + add_cors_middleware(configured, [ALLOWED_ORIGIN]) + + unconfigured = FastAPI() + add_cors_middleware(unconfigured, []) + + configured_classes: list[object] = [middleware.cls for middleware in configured.user_middleware] + unconfigured_classes: list[object] = [middleware.cls for middleware in unconfigured.user_middleware] + + assert CORSMiddleware in configured_classes + assert unconfigured_classes == [] + + +def test_the_application_module_wires_cors_up(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + """main.py must actually call add_cors_middleware, not merely be able to. + + main does its CORS setup at import, so by now it has already been imported - via + tests/helpers/mocking.py - under whatever the ambient .env says, which on a developer + machine or in CI is usually no origins at all. Reloading it with an origins file + configured is what lets the real wiring run against a real configuration. It is reloaded + again afterwards because main.app is module global and would otherwise leak into later tests. + """ + + origins_file = tmp_path / "cors-allowed-origins.json" + origins_file.write_text(json.dumps([ALLOWED_ORIGIN])) + monkeypatch.setenv(CORS_ALLOWED_ORIGINS_FILE, str(origins_file)) + + import main + + try: + reloaded = importlib.reload(main) + + middleware_classes: list[object] = [middleware.cls for middleware in reloaded.app.user_middleware] + + assert CORSMiddleware in middleware_classes + finally: + monkeypatch.delenv(CORS_ALLOWED_ORIGINS_FILE, raising=False) + importlib.reload(main) diff --git a/tests/unit/test_config.py b/tests/unit/test_config.py new file mode 100644 index 0000000..2018819 --- /dev/null +++ b/tests/unit/test_config.py @@ -0,0 +1,42 @@ +"""Tests for configuration read straight from the environment. + +Sentry and CORS both rely on this, because each is configured before the Context exists. +""" + +from pathlib import Path + +import pytest + +from register_your_data_api.config import get_environment_config + +SOME_VARIABLE = "SOME_VARIABLE" + + +def test_environment_config_reads_the_env_file(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + # os.environ wins over the file, so a real variable of this name in the ambient + # environment would mask what this test is checking. + monkeypatch.delenv(SOME_VARIABLE, raising=False) + + env_file = tmp_path / ".env" + env_file.write_text('SOME_VARIABLE="a-value"\n') + + env = get_environment_config(str(env_file)) + + assert env[SOME_VARIABLE] == "a-value" + + +def test_environment_config_prefers_os_environ(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + env_file = tmp_path / ".env" + env_file.write_text('SOME_VARIABLE="from-the-env-file"\n') + + monkeypatch.setenv(SOME_VARIABLE, "from-os-environ") + + env = get_environment_config(str(env_file)) + + assert env[SOME_VARIABLE] == "from-os-environ" + + +def test_environment_config_tolerates_a_missing_env_file(tmp_path: Path) -> None: + env = get_environment_config(str(tmp_path / "does-not-exist")) + + assert isinstance(env, dict) diff --git a/tests/unit/test_cors_config.py b/tests/unit/test_cors_config.py new file mode 100644 index 0000000..653bbb8 --- /dev/null +++ b/tests/unit/test_cors_config.py @@ -0,0 +1,174 @@ +import json +from pathlib import Path +from typing import Any + +import pytest + +from register_your_data_api.cors import CORS_ALLOWED_ORIGINS_FILE, load_allowed_origins + + +def write_origins_file(tmp_path: Path, contents: Any) -> str: + """Write an origins file containing the given JSON-serialisable contents.""" + + origins_file = tmp_path / "cors-allowed-origins.json" + origins_file.write_text(json.dumps(contents)) + + return str(origins_file) + + +def write_raw_origins_file(tmp_path: Path, contents: str) -> str: + """Write an origins file with contents that are not necessarily valid JSON.""" + + origins_file = tmp_path / "cors-allowed-origins.json" + origins_file.write_text(contents) + + return str(origins_file) + + +def test_no_origins_configured() -> None: + assert load_allowed_origins({}) == [] + + +@pytest.mark.parametrize("filename", ["", " "]) +def test_blank_filename_configured(filename: str) -> None: + assert load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: filename}) == [] + + +def test_origins_are_loaded(tmp_path: Path) -> None: + origins = ["https://client.example.org", "https://dev.client.example.org"] + + filename = write_origins_file(tmp_path, origins) + + assert load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: filename}) == origins + + +def test_origins_file_may_be_an_empty_list(tmp_path: Path) -> None: + """An empty list is valid configuration and means no cross-origin requests are allowed.""" + + filename = write_origins_file(tmp_path, []) + + assert load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: filename}) == [] + + +@pytest.mark.parametrize( + "origin", + [ + "http://localhost:5173", + "https://client.example.org", + "https://client.example.org:8443", + ], +) +def test_valid_origin_forms(tmp_path: Path, origin: str) -> None: + filename = write_origins_file(tmp_path, [origin]) + + assert load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: filename}) == [origin] + + +def test_missing_origins_file(tmp_path: Path) -> None: + filename = str(tmp_path / "does-not-exist.json") + + with pytest.raises(RuntimeError, match="Could not read CORS allowed origins file"): + load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: filename}) + + +def test_origins_file_is_not_valid_json(tmp_path: Path) -> None: + filename = write_raw_origins_file(tmp_path, "['not json',]") + + with pytest.raises(RuntimeError, match="is not valid JSON"): + load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: filename}) + + +def test_origins_file_is_not_readable_as_text(tmp_path: Path) -> None: + """A file in the wrong encoding raises UnicodeDecodeError, which must not escape.""" + + origins_file = tmp_path / "cors-allowed-origins.json" + origins_file.write_bytes(json.dumps(["https://client.example.org"]).encode("utf-16")) + + with pytest.raises(RuntimeError, match="is not valid JSON"): + load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: str(origins_file)}) + + +@pytest.mark.parametrize( + "contents", + [ + pytest.param({"origins": ["https://client.example.org"]}, id="object_not_list"), + pytest.param("https://client.example.org", id="string_not_list"), + pytest.param(123, id="number_not_list"), + ], +) +def test_origins_file_must_contain_a_list(tmp_path: Path, contents: Any) -> None: + """Valid JSON of the wrong shape is a different failure from a bad entry within a list.""" + + filename = write_origins_file(tmp_path, contents) + + with pytest.raises(RuntimeError, match="must contain a JSON list of origins"): + load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: filename}) + + +@pytest.mark.parametrize( + "origin", + [ + pytest.param(123, id="number_entry"), + pytest.param(None, id="null_entry"), + pytest.param(["https://client.example.org"], id="nested_list_entry"), + pytest.param("", id="empty_entry"), + pytest.param(" ", id="whitespace_entry"), + pytest.param("https://client.example.org/", id="trailing_slash"), + pytest.param("https://client.example.org/api/v1", id="path"), + pytest.param("https://Client.Example.org", id="not_lower_case"), + pytest.param("client.example.org", id="no_scheme"), + pytest.param("ftp://client.example.org", id="unsupported_scheme"), + ], +) +def test_malformed_origins_are_rejected(tmp_path: Path, origin: Any) -> None: + """Each of these shapes must be refused. + + What varies between the cases is only whether the origin is rejected at all - the + wording of the failure is the same for every one of them, and is covered by + test_the_error_names_only_the_invalid_origins below. + """ + + filename = write_origins_file(tmp_path, [origin]) + + with pytest.raises(RuntimeError, match="contains invalid origins"): + load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: filename}) + + +def test_the_error_names_only_the_invalid_origins(tmp_path: Path) -> None: + """Whoever is editing a long list needs to be told which entries are wrong.""" + + filename = write_origins_file( + tmp_path, + [ + "https://good-one.example.org", + "https://bad-one.example.org/", + "https://good-two.example.org", + "ftp://bad-two.example.org", + ], + ) + + with pytest.raises(RuntimeError, match="contains invalid origins") as raised: + load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: filename}) + + message = str(raised.value) + + assert "https://bad-one.example.org/" in message + assert "ftp://bad-two.example.org" in message + assert "good-one" not in message + assert "good-two" not in message + + +@pytest.mark.parametrize( + "contents", + [ + pytest.param(["*"], id="only_wildcard"), + pytest.param(["https://client.example.org", "*"], id="wildcard_alongside_an_origin"), + ], +) +def test_wildcard_origins_are_rejected(tmp_path: Path, contents: Any) -> None: + """A single "*" entry would put CORSMiddleware into allow-all mode, so it must not load.""" + + filename = write_origins_file(tmp_path, contents) + + with pytest.raises(RuntimeError, match="Wildcards are not allowed"): + load_allowed_origins({CORS_ALLOWED_ORIGINS_FILE: filename}) diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py index 14b5699..9f680c0 100644 --- a/tests/unit/test_main.py +++ b/tests/unit/test_main.py @@ -30,3 +30,25 @@ async def enter_lifespan() -> None: # only that startup failed and not why. Logger.exception delegates to Logger.error # with exc_info=True, which is why the assertion is on `error`. logged.assert_called_once_with("Could not initialise application - error setting up context", exc_info=True) + + +def test_a_failed_cors_configuration_is_logged_with_the_exception_attached() -> None: + """The CORS allowlist is read at import, so the failure is carried into the lifespan. + + It is reported the same way as a failed context setup: an application which will not + start because its CORS configuration is broken is exactly as worth being told about, + and without `exc_info` the reported event would not say which origin was at fault. + """ + + cors_error = RuntimeError("CORS allowed origins file origins.json contains invalid origins") + + async def enter_lifespan() -> None: + async with main.prod_lifespan(mock.MagicMock()): + pass + + with mock.patch.object(main, "_cors_configuration_error", cors_error): + with mock.patch.object(logging.getLogger("main"), "error") as logged: + with pytest.raises(SystemExit): + asyncio.run(enter_lifespan()) + + logged.assert_called_once_with("Could not initialise application - error configuring CORS", exc_info=cors_error)