Add kaya-cors package for CORS support
CI / Build Pip package (push) Successful in 3m44s

This commit is contained in:
2026-09-04 15:25:26 +08:00
parent 59f1a8227f
commit 270a0d87fc
9 changed files with 658 additions and 1 deletions
+63
View File
@@ -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.
+56
View File
@@ -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/"
@@ -0,0 +1,10 @@
from pkgutil import extend_path
__path__ = extend_path(__path__, __name__)
from ._mixin import CorsMixin
__all__ = [
'CorsMixin',
]
+274
View File
@@ -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))
+246
View File
@@ -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'))