diff --git a/README.md b/README.md index 061f758..ab1ad71 100644 --- a/README.md +++ b/README.md @@ -13,6 +13,7 @@ This repository is a monorepo for the Kaya framework. The code is split into ind - **kaya-session-memcache** — memcached-backed session storage (`packages/kaya-session-memcache/`) - **kaya-oidc** — OpenID Connect authentication (`packages/kaya-oidc/`) - **kaya-openapi** — automatic OpenAPI specification generation (`packages/kaya-openapi/`) +- **kaya-cors** — CORS (Cross-Origin Resource Sharing) support (`packages/kaya-cors/`) Additional `kaya-*` packages can be added as new directories under `packages/`. @@ -27,7 +28,7 @@ pip install --index-url https://gitea.woggioni.net/api/packages/woggioni/pypi/si Install the packages in development mode: ```bash -pip install -e packages/kaya-core -e packages/kaya-rsgi -e packages/kaya-session -e packages/kaya-session-redis -e packages/kaya-session-memcache -e packages/kaya-oidc -e packages/kaya-openapi +pip install -e packages/kaya-core -e packages/kaya-rsgi -e packages/kaya-session -e packages/kaya-session-redis -e packages/kaya-session-memcache -e packages/kaya-oidc -e packages/kaya-openapi -e packages/kaya-cors ``` Run the example: @@ -46,6 +47,7 @@ python -m unittest discover -s packages/kaya-session-redis/tests python -m unittest discover -s packages/kaya-session-memcache/tests python -m unittest discover -s packages/kaya-oidc/tests python -m unittest discover -s packages/kaya-openapi/tests +python -m unittest discover -s packages/kaya-cors/tests ``` ## Static analysis @@ -58,6 +60,7 @@ mypy -p kaya.session.redis mypy -p kaya.session.memcache mypy -p kaya.oidc mypy -p kaya.openapi +mypy -p kaya.cors ``` ## Building packages @@ -70,4 +73,5 @@ python -m build packages/kaya-session-redis python -m build packages/kaya-session-memcache python -m build packages/kaya-oidc python -m build packages/kaya-openapi +python -m build packages/kaya-cors ``` diff --git a/packages/kaya-cors/README.md b/packages/kaya-cors/README.md new file mode 100644 index 0000000..9e14ab6 --- /dev/null +++ b/packages/kaya-cors/README.md @@ -0,0 +1,63 @@ +# kaya-cors + +CORS (Cross-Origin Resource Sharing) support for the Kaya web framework. + +Provides `CorsMixin`, a `KayaMixin` that adds CORS response headers to outgoing +responses and answers CORS preflight (`OPTIONS`) requests, with the same +configuration parameters and semantics as FastAPI/Starlette's `CORSMiddleware`. + +## Usage + +```python +from kaya.core import KayaApp, HttpContext +from kaya.cors import CorsMixin + +app = KayaApp(mixins=[ + CorsMixin( + allow_origins=['https://example.com'], + allow_methods=('GET', 'POST'), + allow_headers=('X-Custom-Header',), + allow_credentials=True, + max_age=600, + ) +]) + +@app.GET('/') +async def home(ctx: HttpContext): + await ctx.send_str(200, 'Hello World!') +``` + +## Parameters + +- `allow_origins`: list of origins allowed to make cross-origin requests. + Use `['*']` to allow any origin. +- `allow_origin_regex`: optional regex string matched (fullmatch) against the + request origin. +- `allow_methods`: HTTP methods allowed for cross-origin requests + (default `('GET',)`); use `'*'` to allow all standard methods. +- `allow_headers`: request headers allowed in cross-origin requests + (default `()`); use `'*'` to mirror back any requested headers. +- `allow_credentials`: allow cookies/credentials in cross-origin requests + (default `False`). When enabled, the allowed origin is always echoed + explicitly instead of `'*'`. +- `expose_headers`: response headers made accessible to the browser. +- `max_age`: seconds browsers may cache the preflight response + (default `600`). + +## Behavior + +- Requests without an `Origin` header pass through untouched. +- Simple cross-origin requests with an allowed origin get + `Access-Control-Allow-Origin` (plus `Access-Control-Allow-Credentials` and + `Access-Control-Expose-Headers` when configured) added to the response. + Headers already set by the handler are never overwritten. +- Preflight requests (`OPTIONS` with `Origin` and + `Access-Control-Request-Method` headers) are answered directly by the mixin + with `200 OK` (or `400` with a `Disallowed CORS ...` body when the origin, + method or headers are not allowed). The preflight response is the only one + delivered to the client: if the routing tree matches the request anyway + (including user-registered `OPTIONS` handlers or the 404 fallback), its + output is discarded. + +`CorsMixin` is a `KayaMixin`, so the app stays a `KayaApp` and both ASGI and +RSGI keep working. diff --git a/packages/kaya-cors/pyproject.toml b/packages/kaya-cors/pyproject.toml new file mode 100644 index 0000000..3f0b5d2 --- /dev/null +++ b/packages/kaya-cors/pyproject.toml @@ -0,0 +1,56 @@ +[build-system] +requires = ["setuptools>=61.0", "setuptools-scm>=8"] +build-backend = "setuptools.build_meta" + +[project] +name = "kaya-cors" +dynamic = ["version"] +authors = [ + { name="Walter Oggioni", email="oggioni.walter@gmail.com" }, +] +description = "CORS support for the Kaya lightweight ASGI web framework" +readme = "README.md" +requires-python = ">=3.10" +license = "MIT" +classifiers = [ + 'Development Status :: 3 - Alpha', + 'Topic :: Utilities', + 'Intended Audience :: System Administrators', + 'Intended Audience :: Developers', + 'Environment :: Console', + 'Programming Language :: Python :: 3', +] + +dependencies = [ + "kaya-core", +] + +[project.optional-dependencies] +dev = [ + "build", "mypy", "ipdb", "twine", "httpx", "httpx-ws", "kaya-rsgi" +] + +[project.urls] +"Homepage" = "https://github.com/woggioni/kaya" +"Bug Tracker" = "https://github.com/woggioni/kaya/issues" + +[tool.setuptools.packages.find] +where = ["src"] +namespaces = true + +[tool.mypy] +python_version = "3.12" +disallow_untyped_defs = true +show_error_codes = true +no_implicit_optional = true +warn_return_any = true +warn_unused_ignores = true +exclude = ["scripts", "docs", "test"] +strict = true + +[tool.setuptools_scm] +root = "../.." +version_file = "src/kaya/cors/_version.py" + +[tool.setuptools_scm.tag] +prefix = "release/" diff --git a/packages/kaya-cors/src/kaya/cors/__init__.py b/packages/kaya-cors/src/kaya/cors/__init__.py new file mode 100644 index 0000000..476a2db --- /dev/null +++ b/packages/kaya-cors/src/kaya/cors/__init__.py @@ -0,0 +1,10 @@ +from pkgutil import extend_path + +__path__ = extend_path(__path__, __name__) + +from ._mixin import CorsMixin + + +__all__ = [ + 'CorsMixin', +] diff --git a/packages/kaya-cors/src/kaya/cors/_mixin.py b/packages/kaya-cors/src/kaya/cors/_mixin.py new file mode 100644 index 0000000..70b7499 --- /dev/null +++ b/packages/kaya-cors/src/kaya/cors/_mixin.py @@ -0,0 +1,274 @@ +import re +from pathlib import Path +from typing import ( + Any, + AsyncGenerator, + Dict, + List, + Mapping, + Optional, + Sequence, + Tuple, +) + +from kaya.core import HttpContext, HttpMethod, KayaApp, KayaMixin +from kaya.core._types import StrOrStrings + +ALL_METHODS: Tuple[str, ...] = ("DELETE", "GET", "HEAD", "OPTIONS", "PATCH", "POST", "PUT", "QUERY") +SAFELISTED_HEADERS = frozenset({"Accept", "Accept-Language", "Content-Language", "Content-Type"}) + + +def _first_header(headers: Mapping[str, Sequence[str]], name: str) -> Optional[str]: + values = headers.get(name) + if not values: + return None + return values[0] + + +def _merge_headers(headers: Optional[Mapping[str, StrOrStrings]], + cors_headers: Mapping[str, str]) -> Mapping[str, StrOrStrings]: + """Merge CORS headers into the response headers. + + Header names are matched case-insensitively; headers already set by the + handler are never overwritten. A CORS ``Vary`` value is appended to an + existing ``Vary`` header when not already present. + """ + result: Dict[str, StrOrStrings] = dict(headers) if headers else {} + key_by_lower: Dict[str, str] = {k.lower(): k for k in result} + for key, value in cors_headers.items(): + existing_key = key_by_lower.get(key.lower()) + if existing_key is None: + result[key] = value + key_by_lower[key.lower()] = key + elif key.lower() == 'vary': + previous = result[existing_key] + previous_values = [previous] if isinstance(previous, str) else list(previous) + present = {v.strip().lower() for part in previous_values for v in part.split(',')} + if value.lower() not in present: + if isinstance(previous, str): + result[existing_key] = f"{previous}, {value}" + else: + result[existing_key] = (*previous_values, value) + return result + + +class CorsHttpContext(HttpContext): + """HttpContext wrapper that injects CORS headers into response headers. + + Works with any concrete ``HttpContext`` (ASGI or RSGI) because it only + relies on the abstract send methods, which all implementations share. + Attributes not explicitly overridden are delegated to the wrapped context + via ``__getattr__``, so protocol-specific fields (``pathsend``, + ``receive``/``send`` for ASGI, ``protocol`` for RSGI, etc.) are passed + through transparently. + """ + + def __init__(self, ctx: HttpContext, cors_headers: Mapping[str, str]) -> None: + object.__setattr__(self, '_ctx', ctx) + object.__setattr__(self, 'session', ctx.session) + object.__setattr__(self, '_cors_headers', cors_headers) + + def __getattr__(self, name: str) -> Any: + if name == '_ctx': + raise AttributeError(name) + return getattr(self._ctx, name) + + def _merge(self, headers: Optional[Mapping[str, StrOrStrings]]) -> Mapping[str, StrOrStrings]: + return _merge_headers(headers, self._cors_headers) + + async def stream_body(self, + status: int, + body_generator: AsyncGenerator[bytes, None], + headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + await self._ctx.stream_body(status, body_generator, self._merge(headers)) + + async def send_bytes(self, status: int, body: bytes, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + await self._ctx.send_bytes(status, body, self._merge(headers)) + + async def send_str(self, status: int, body: str, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + await self._ctx.send_str(status, body, self._merge(headers)) + + async def send_file(self, status: int, path: Path, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + await self._ctx.send_file(status, path, self._merge(headers)) + + async def send_empty(self, status: int, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + await self._ctx.send_empty(status, self._merge(headers)) + + +class _SwallowedHttpContext(HttpContext): + """HttpContext wrapper whose send methods are no-ops. + + Returned by the CORS hook after a preflight response has already been sent, + so that routing (or the 404 fallback) does not attempt to send a second + response for the same request. + """ + + def __init__(self, ctx: HttpContext) -> None: + object.__setattr__(self, '_ctx', ctx) + object.__setattr__(self, 'session', ctx.session) + + def __getattr__(self, name: str) -> Any: + if name == '_ctx': + raise AttributeError(name) + return getattr(self._ctx, name) + + async def stream_body(self, + status: int, + body_generator: AsyncGenerator[bytes, None], + headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + pass + + async def send_bytes(self, status: int, body: bytes, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + pass + + async def send_file(self, status: int, path: Path, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + pass + + async def send_empty(self, status: int, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + pass + + +class CorsMixin(KayaMixin): + """Kaya mixin adding CORS headers to responses, modeled after + FastAPI/Starlette's ``CORSMiddleware``. + + Registers a before-request hook that: + + - answers CORS preflight requests (``OPTIONS`` requests carrying ``Origin`` + and ``Access-Control-Request-Method`` headers) directly: ``200 OK`` when + the origin, method and headers are allowed, ``400`` with a + ``Disallowed CORS ...`` body otherwise. The preflight response is the + only one delivered to the client; if the routing tree matches the + request anyway, its output is discarded; + - wraps the request context in a :class:`CorsHttpContext` for simple + cross-origin requests, injecting ``Access-Control-Allow-Origin`` (and + the configured credentials/expose headers) into the response. + + Requests without an ``Origin`` header pass through untouched. Because the + app stays a ``KayaApp``, both ASGI and RSGI keep working. + + Example:: + + app = KayaApp(mixins=[ + CorsMixin( + allow_origins=['https://example.com'], + allow_methods=('GET', 'POST'), + allow_headers=('X-Custom-Header',), + allow_credentials=True, + ) + ]) + """ + + def __init__(self, + allow_origins: Sequence[str] = (), + allow_methods: Sequence[str] = ('GET',), + allow_headers: Sequence[str] = (), + allow_credentials: bool = False, + allow_origin_regex: Optional[str] = None, + expose_headers: Sequence[str] = (), + max_age: int = 600) -> None: + methods: Sequence[str] = ALL_METHODS if '*' in allow_methods else allow_methods + self._allow_origin_regex = re.compile(allow_origin_regex) if allow_origin_regex is not None else None + self._allow_all_origins = '*' in allow_origins + self._allow_all_headers = '*' in allow_headers + self._allow_credentials = allow_credentials + self._allow_origins = tuple(allow_origins) + self._allow_methods = tuple(methods) + sorted_allow_headers = sorted(SAFELISTED_HEADERS | set(allow_headers)) + self._allow_headers = [h.lower() for h in sorted_allow_headers] + self._preflight_explicit_allow_origin = not self._allow_all_origins or allow_credentials + + simple_headers: Dict[str, str] = {} + if self._allow_all_origins: + simple_headers['Access-Control-Allow-Origin'] = '*' + if allow_credentials: + simple_headers['Access-Control-Allow-Credentials'] = 'true' + if expose_headers: + simple_headers['Access-Control-Expose-Headers'] = ', '.join(expose_headers) + self._simple_headers = simple_headers + + preflight_headers: Dict[str, str] = {} + if self._preflight_explicit_allow_origin: + # the origin value is set dynamically in _preflight_response() + preflight_headers['Vary'] = 'Origin' + else: + preflight_headers['Access-Control-Allow-Origin'] = '*' + preflight_headers['Access-Control-Allow-Methods'] = ', '.join(self._allow_methods) + preflight_headers['Access-Control-Max-Age'] = str(max_age) + if sorted_allow_headers and not self._allow_all_headers: + preflight_headers['Access-Control-Allow-Headers'] = ', '.join(sorted_allow_headers) + if allow_credentials: + preflight_headers['Access-Control-Allow-Credentials'] = 'true' + self._preflight_headers = preflight_headers + + def apply(self, app: KayaApp) -> None: + app.add_before_request_hook(self._before_request) + + def _is_allowed_origin(self, origin: str) -> bool: + if self._allow_all_origins: + return True + if self._allow_origin_regex is not None and self._allow_origin_regex.fullmatch(origin): + return True + return origin in self._allow_origins + + def _simple_response_headers(self, origin: str) -> Mapping[str, str]: + headers = dict(self._simple_headers) + if self._allow_all_origins and self._allow_credentials: + # credentials require the specific origin instead of '*' + headers['Access-Control-Allow-Origin'] = origin + headers['Vary'] = 'Origin' + elif not self._allow_all_origins and self._is_allowed_origin(origin): + # specific origins must be mirrored back in the response + headers['Access-Control-Allow-Origin'] = origin + headers['Vary'] = 'Origin' + return headers + + def _preflight_response(self, + request_headers: Mapping[str, Sequence[str]], + origin: str) -> Tuple[int, str, Mapping[str, str]]: + requested_method = _first_header(request_headers, 'access-control-request-method') + requested_headers = _first_header(request_headers, 'access-control-request-headers') + + headers = dict(self._preflight_headers) + failures: List[str] = [] + + if self._is_allowed_origin(origin): + if self._preflight_explicit_allow_origin: + # the "else" case is already accounted for in self._preflight_headers + # and the value would be '*' + headers['Access-Control-Allow-Origin'] = origin + else: + failures.append('origin') + + if requested_method not in self._allow_methods: + failures.append('method') + + # if we allow all headers, then we have to mirror back any requested + # headers in the response + if self._allow_all_headers and requested_headers is not None: + headers['Access-Control-Allow-Headers'] = requested_headers + elif requested_headers is not None: + for header in [h.lower() for h in requested_headers.split(',')]: + if header.strip() not in self._allow_headers: + failures.append('headers') + break + + # we don't strictly need to use 400 responses here, since it's up to + # the browser to enforce the CORS policy, but it's more informative + # if we do + if failures: + return 400, 'Disallowed CORS ' + ', '.join(failures), headers + return 200, 'OK', headers + + async def _before_request(self, ctx: HttpContext) -> Optional[HttpContext]: + origin = _first_header(ctx.headers, 'origin') + if origin is None: + return None + if ctx.method == HttpMethod.OPTIONS \ + and _first_header(ctx.headers, 'access-control-request-method') is not None: + status, body, headers = self._preflight_response(ctx.headers, origin) + response_headers: Dict[str, StrOrStrings] = {'Content-Type': 'text/plain; charset=utf-8'} + response_headers.update(headers) + await ctx.send_str(status, body, response_headers) + return _SwallowedHttpContext(ctx) + return CorsHttpContext(ctx, self._simple_response_headers(origin)) diff --git a/packages/kaya-cors/src/kaya/cors/py.typed b/packages/kaya-cors/src/kaya/cors/py.typed new file mode 100644 index 0000000..e69de29 diff --git a/packages/kaya-cors/tests/test_cors.py b/packages/kaya-cors/tests/test_cors.py new file mode 100644 index 0000000..280c3c2 --- /dev/null +++ b/packages/kaya-cors/tests/test_cors.py @@ -0,0 +1,246 @@ +import asyncio +import unittest +from typing import Any, Optional + +import httpx +from pwo import async_test + +from kaya.core import HttpContext, KayaApp +from kaya.cors import CorsMixin + + +def make_app(**cors_kwargs: Any) -> KayaApp: + app = KayaApp(mixins=[CorsMixin(**cors_kwargs)]) + + @app.GET('/hello') + async def hello(ctx: HttpContext) -> None: + await ctx.send_str(200, 'Hello World!') + + return app + + +async def request(app: KayaApp, + method: str, + path: str = '/hello', + headers: Optional[dict[str, str]] = None) -> httpx.Response: + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as client: + return await client.request(method, path, headers=headers) + + +class CorsSimpleRequestTest(unittest.TestCase): + + @async_test + async def test_allowed_origin(self): + app = make_app(allow_origins=['https://example.com']) + r = await request(app, 'GET', headers={'Origin': 'https://example.com'}) + self.assertEqual(200, r.status_code) + self.assertEqual('Hello World!', r.text) + self.assertEqual('https://example.com', r.headers.get('Access-Control-Allow-Origin')) + self.assertEqual('Origin', r.headers.get('Vary')) + + @async_test + async def test_disallowed_origin(self): + app = make_app(allow_origins=['https://example.com']) + r = await request(app, 'GET', headers={'Origin': 'https://evil.com'}) + self.assertEqual(200, r.status_code) + self.assertNotIn('Access-Control-Allow-Origin', r.headers) + + @async_test + async def test_wildcard_origin(self): + app = make_app(allow_origins=['*']) + r = await request(app, 'GET', headers={'Origin': 'https://anything.example.com'}) + self.assertEqual('*', r.headers.get('Access-Control-Allow-Origin')) + + @async_test + async def test_wildcard_origin_with_credentials_echoes_origin(self): + app = make_app(allow_origins=['*'], allow_credentials=True) + r = await request(app, 'GET', headers={'Origin': 'https://example.com'}) + self.assertEqual('https://example.com', r.headers.get('Access-Control-Allow-Origin')) + self.assertEqual('true', r.headers.get('Access-Control-Allow-Credentials')) + self.assertEqual('Origin', r.headers.get('Vary')) + + @async_test + async def test_origin_regex(self): + app = make_app(allow_origin_regex=r'https://.*\.example\.com') + r = await request(app, 'GET', headers={'Origin': 'https://api.example.com'}) + self.assertEqual('https://api.example.com', r.headers.get('Access-Control-Allow-Origin')) + r = await request(app, 'GET', headers={'Origin': 'https://example.com.evil.org'}) + self.assertNotIn('Access-Control-Allow-Origin', r.headers) + + @async_test + async def test_no_origin_header(self): + app = make_app(allow_origins=['*']) + r = await request(app, 'GET') + self.assertEqual(200, r.status_code) + self.assertNotIn('Access-Control-Allow-Origin', r.headers) + + @async_test + async def test_expose_headers(self): + app = make_app(allow_origins=['*'], expose_headers=['X-Total-Count']) + r = await request(app, 'GET', headers={'Origin': 'https://example.com'}) + self.assertEqual('X-Total-Count', r.headers.get('Access-Control-Expose-Headers')) + + @async_test + async def test_handler_set_cors_header_not_overwritten(self): + app = KayaApp(mixins=[CorsMixin(allow_origins=['*'])]) + + @app.GET('/custom') + async def custom(ctx: HttpContext) -> None: + await ctx.send_str(200, 'custom', headers={'Access-Control-Allow-Origin': 'https://custom.example.com'}) + + r = await request(app, 'GET', '/custom', headers={'Origin': 'https://example.com'}) + self.assertEqual('https://custom.example.com', r.headers.get('Access-Control-Allow-Origin')) + + +class CorsPreflightTest(unittest.TestCase): + + @staticmethod + def preflight_headers(origin: str = 'https://example.com', + method: str = 'POST', + headers: Optional[str] = None) -> dict[str, str]: + result = { + 'Origin': origin, + 'Access-Control-Request-Method': method, + } + if headers is not None: + result['Access-Control-Request-Headers'] = headers + return result + + @async_test + async def test_preflight_allowed(self): + app = make_app(allow_origins=['https://example.com'], + allow_methods=('GET', 'POST'), + allow_credentials=True) + r = await request(app, 'OPTIONS', headers=self.preflight_headers()) + self.assertEqual(200, r.status_code) + self.assertEqual('OK', r.text) + self.assertEqual('https://example.com', r.headers.get('Access-Control-Allow-Origin')) + self.assertEqual('GET, POST', r.headers.get('Access-Control-Allow-Methods')) + self.assertEqual('600', r.headers.get('Access-Control-Max-Age')) + self.assertEqual('true', r.headers.get('Access-Control-Allow-Credentials')) + self.assertEqual('Origin', r.headers.get('Vary')) + + @async_test + async def test_preflight_wildcard_origin(self): + app = make_app(allow_origins=['*'], allow_methods=('GET', 'POST')) + r = await request(app, 'OPTIONS', headers=self.preflight_headers()) + self.assertEqual(200, r.status_code) + self.assertEqual('*', r.headers.get('Access-Control-Allow-Origin')) + + @async_test + async def test_preflight_disallowed_origin(self): + app = make_app(allow_origins=['https://example.com'], allow_methods=('GET', 'POST')) + r = await request(app, 'OPTIONS', headers=self.preflight_headers(origin='https://evil.com')) + self.assertEqual(400, r.status_code) + self.assertEqual('Disallowed CORS origin', r.text) + + @async_test + async def test_preflight_disallowed_method(self): + app = make_app(allow_origins=['https://example.com'], allow_methods=('GET',)) + r = await request(app, 'OPTIONS', headers=self.preflight_headers(method='DELETE')) + self.assertEqual(400, r.status_code) + self.assertEqual('Disallowed CORS method', r.text) + + @async_test + async def test_preflight_disallowed_headers(self): + app = make_app(allow_origins=['https://example.com'], allow_methods=('GET', 'POST')) + r = await request(app, 'OPTIONS', headers=self.preflight_headers(headers='X-Custom')) + self.assertEqual(400, r.status_code) + self.assertEqual('Disallowed CORS headers', r.text) + + @async_test + async def test_preflight_safelisted_headers_allowed(self): + app = make_app(allow_origins=['https://example.com'], allow_methods=('GET', 'POST')) + r = await request(app, 'OPTIONS', headers=self.preflight_headers(headers='Content-Type')) + self.assertEqual(200, r.status_code) + + @async_test + async def test_preflight_allow_all_headers_mirrors_request(self): + app = make_app(allow_origins=['*'], allow_methods=('GET', 'POST'), allow_headers=['*']) + r = await request(app, 'OPTIONS', headers=self.preflight_headers(headers='X-Custom, X-Other')) + self.assertEqual(200, r.status_code) + self.assertEqual('X-Custom, X-Other', r.headers.get('Access-Control-Allow-Headers')) + + @async_test + async def test_preflight_configured_allow_headers(self): + app = make_app(allow_origins=['https://example.com'], + allow_methods=('GET', 'POST'), + allow_headers=('X-Custom',)) + r = await request(app, 'OPTIONS', headers=self.preflight_headers(headers='X-Custom')) + self.assertEqual(200, r.status_code) + allow_headers = r.headers.get('Access-Control-Allow-Headers') + self.assertIsNotNone(allow_headers) + assert allow_headers is not None + self.assertIn('X-Custom', allow_headers) + + @async_test + async def test_preflight_response_not_overwritten_by_handler(self): + # even when a user-registered OPTIONS handler matches, the preflight + # response sent by the mixin is the only one delivered to the client + app = KayaApp(mixins=[CorsMixin(allow_origins=['https://example.com'], allow_methods=('GET', 'POST'))]) + + @app.GET('/hello') + async def hello(ctx: HttpContext) -> None: + await ctx.send_str(200, 'Hello World!') + + @app.OPTIONS('/hello') + async def options(ctx: HttpContext) -> None: + await ctx.send_str(200, 'custom OPTIONS handler') + + r = await request(app, 'OPTIONS', headers=self.preflight_headers()) + self.assertEqual(200, r.status_code) + self.assertEqual('OK', r.text) + + @async_test + async def test_options_without_preflight_headers_routes_normally(self): + app = KayaApp(mixins=[CorsMixin(allow_origins=['https://example.com'])]) + + @app.OPTIONS('/hello') + async def options(ctx: HttpContext) -> None: + await ctx.send_str(200, 'custom OPTIONS handler') + + # an OPTIONS request without Access-Control-Request-Method is not a + # preflight request and is routed normally + r = await request(app, 'OPTIONS', headers={'Origin': 'https://example.com'}) + self.assertEqual(200, r.status_code) + self.assertEqual('custom OPTIONS handler', r.text) + self.assertEqual('https://example.com', r.headers.get('Access-Control-Allow-Origin')) + + +class CorsRsgiTest(unittest.TestCase): + + def test_rsgi_context_header_injection(self): + from kaya.rsgi import RsgiContext + + class FakeScope: + scheme = 'http' + method = 'GET' + path = '/' + query_string = '' + headers = {'origin': 'https://example.com'} + client = '127.0.0.1:12345' + server = '127.0.0.1:80' + + class FakeProtocol: + def __init__(self) -> None: + self.responses = [] + + def response_str(self, status: int, headers: list, body: str) -> None: + self.responses.append((status, dict(headers), body)) + + mixin = CorsMixin(allow_origins=['https://example.com']) + protocol = FakeProtocol() + ctx = RsgiContext(FakeScope(), protocol) # type: ignore[arg-type] + + async def run() -> None: + wrapped = await mixin._before_request(ctx) + assert wrapped is not None + await wrapped.send_str(200, 'hi') + + asyncio.run(run()) + self.assertEqual(1, len(protocol.responses)) + status, headers, body = protocol.responses[0] + self.assertEqual(200, status) + self.assertEqual('https://example.com', headers.get('Access-Control-Allow-Origin')) + self.assertEqual('Origin', headers.get('Vary')) diff --git a/requirements-dev.in b/requirements-dev.in index 49159e2..40fdde9 100644 --- a/requirements-dev.in +++ b/requirements-dev.in @@ -5,6 +5,7 @@ kaya-session-redis @ file:./packages/kaya-session-redis kaya-session-memcache @ file:./packages/kaya-session-memcache kaya-oidc @ file:./packages/kaya-oidc kaya-openapi @ file:./packages/kaya-openapi +kaya-cors @ file:./packages/kaya-cors build fakeredis mypy diff --git a/requirements-dev.txt b/requirements-dev.txt index 7e7c199..8ca0c52 100644 --- a/requirements-dev.txt +++ b/requirements-dev.txt @@ -89,10 +89,13 @@ jeepney==0.9.0 file:./packages/kaya-core # via # -r requirements-dev.in + # kaya-cors # kaya-oidc # kaya-openapi # kaya-rsgi # kaya-session +file:./packages/kaya-cors + # via -r requirements-dev.in file:./packages/kaya-oidc # via -r requirements-dev.in file:./packages/kaya-openapi