From 3ebf079533f096ef8f383653ff20a4c06c9c45b3 Mon Sep 17 00:00:00 2001 From: Walter Oggioni Date: Sat, 18 Jul 2026 16:39:34 +0000 Subject: [PATCH] Refactor to composable KayaMixin architecture Replace wrapper-based SessionMiddleware/OIDCApp with KayaMixin subclasses applied via KayaApp(mixins=[...]). Mixins hook into handle_request and handle_websocket via before/after hooks, so both ASGI and RSGI keep working. Mixin dependencies are applied automatically and deduplicated. --- example/oidc.py | 20 +- example/session.py | 4 +- packages/kaya-core/src/kaya/core/__init__.py | 2 + packages/kaya-core/src/kaya/core/_app.py | 92 ++++++-- packages/kaya-core/src/kaya/core/_mixin.py | 37 ++++ packages/kaya-oidc/README.md | 23 +- packages/kaya-oidc/src/kaya/oidc/__init__.py | 4 +- .../src/kaya/oidc/{_app.py => _mixin.py} | 115 ++++------ packages/kaya-oidc/tests/test_oidc.py | 100 +++++---- packages/kaya-session/README.md | 19 +- .../kaya-session/src/kaya/session/__init__.py | 4 +- .../src/kaya/session/_middleware.py | 178 ---------------- .../kaya-session/src/kaya/session/_mixin.py | 200 ++++++++++++++++++ packages/kaya-session/tests/test_session.py | 41 ++-- 14 files changed, 480 insertions(+), 359 deletions(-) create mode 100644 packages/kaya-core/src/kaya/core/_mixin.py rename packages/kaya-oidc/src/kaya/oidc/{_app.py => _mixin.py} (67%) delete mode 100644 packages/kaya-session/src/kaya/session/_middleware.py create mode 100644 packages/kaya-session/src/kaya/session/_mixin.py diff --git a/example/oidc.py b/example/oidc.py index 304c606..9f9266c 100644 --- a/example/oidc.py +++ b/example/oidc.py @@ -1,30 +1,30 @@ import os from kaya.core import HttpContext, KayaApp -from kaya.oidc import OIDCConfig, OIDCApp -from kaya.session import InMemorySessionStore, SessionMiddleware +from kaya.oidc import OIDCConfig, OIDCMixin +from kaya.session import InMemorySessionStore, SessionMixin -app = KayaApp() -session_app = SessionMiddleware(app, InMemorySessionStore()) - -oidc = OIDCApp( - session_app, +session = SessionMixin(InMemorySessionStore()) +oidc = OIDCMixin( OIDCConfig( issuer=os.environ.get('OIDC_ISSUER', 'https://accounts.google.com'), client_id=os.environ.get('OIDC_CLIENT_ID', 'replace-me'), client_secret=os.environ.get('OIDC_CLIENT_SECRET'), redirect_uri=os.environ.get('OIDC_REDIRECT_URI', 'http://localhost:8000/auth/callback'), fetch_userinfo=True, - ) + ), + session=session, ) +app = KayaApp(mixins=[session, oidc]) -@oidc.GET('/') + +@app.GET('/') async def home(ctx: HttpContext) -> None: await ctx.send_str(200, 'public home') -@oidc.GET('/profile') +@app.GET('/profile') @oidc.require_auth async def profile(ctx: HttpContext) -> None: user = oidc.get_user(ctx) diff --git a/example/session.py b/example/session.py index bdb47d5..7849558 100644 --- a/example/session.py +++ b/example/session.py @@ -1,7 +1,7 @@ from kaya.core import HttpContext, KayaApp -from kaya.session import InMemorySessionStore, SessionMiddleware +from kaya.session import InMemorySessionStore, SessionMixin -app = SessionMiddleware(KayaApp(), InMemorySessionStore()) +app = KayaApp(mixins=[SessionMixin(InMemorySessionStore())]) @app.GET('/') diff --git a/packages/kaya-core/src/kaya/core/__init__.py b/packages/kaya-core/src/kaya/core/__init__.py index f30c647..76b56e6 100644 --- a/packages/kaya-core/src/kaya/core/__init__.py +++ b/packages/kaya-core/src/kaya/core/__init__.py @@ -1,6 +1,7 @@ from ._app import AbstractKayaApp, KayaApp from ._http_method import HttpMethod from ._http_context import HttpContext +from ._mixin import KayaMixin from ._tree import Tree, PathIterator from ._path_handler import PathHandler, Matches from ._websocket import WebSocket, WebSocketMessage @@ -10,6 +11,7 @@ __all__ = [ 'AbstractKayaApp', 'HttpMethod', 'KayaApp', + 'KayaMixin', 'HttpContext', 'Tree', 'PathHandler', diff --git a/packages/kaya-core/src/kaya/core/_app.py b/packages/kaya-core/src/kaya/core/_app.py index 0602671..61133f0 100644 --- a/packages/kaya-core/src/kaya/core/_app.py +++ b/packages/kaya-core/src/kaya/core/_app.py @@ -6,9 +6,10 @@ from typing import Callable, Awaitable, Any, Mapping, Sequence, Optional, Tuple, from pwo import Maybe, AsyncQueueIterator from ._http_context import HttpContext from ._http_method import HttpMethod +from ._mixin import KayaMixin from ._path_handler import Context from ._types import StrOrStrings -from ._websocket import WebSocket +from ._websocket import WebSocket, WebSocketMessage from ._asgi import AsgiContext, AsgiWebSocket from ._tree import Tree from ._types.asgi import LifespanScope, HTTPScope as ASGIHTTPScope, WebSocketScope as ASGIWebSocketScope @@ -18,6 +19,10 @@ log = getLogger(__name__) type HttpHandler = Callable[[HttpContext, Unpack[Any]], Awaitable[None]] type WebSocketHandler = Callable[[WebSocket, Unpack[Any]], Awaitable[None]] +type BeforeRequestHook = Callable[[HttpContext], Awaitable[Optional[HttpContext]]] +type AfterRequestHook = Callable[[HttpContext], Awaitable[None]] +type BeforeWebSocketHook = Callable[[WebSocket], Awaitable[Optional[WebSocket]]] +type AfterWebSocketHook = Callable[[WebSocket], Awaitable[None]] class AbstractKayaApp(ABC): @@ -89,25 +94,84 @@ class AbstractKayaApp(ABC): class KayaApp(AbstractKayaApp): _tree: Tree + _mixins: list[KayaMixin] + _applied_ids: set[int] + _before_request_hooks: list[BeforeRequestHook] + _after_request_hooks: list[AfterRequestHook] + _before_websocket_hooks: list[BeforeWebSocketHook] + _after_websocket_hooks: list[AfterWebSocketHook] - def __init__(self) -> None: + def __init__(self, mixins: Sequence[KayaMixin] = ()) -> None: self._tree = Tree() + self._mixins = [] + self._applied_ids = set() + self._before_request_hooks = [] + self._after_request_hooks = [] + self._before_websocket_hooks = [] + self._after_websocket_hooks = [] + for mixin in mixins: + self._apply_mixin(mixin) + + def _apply_mixin(self, mixin: KayaMixin) -> None: + if id(mixin) in self._applied_ids: + return + for dependency in mixin.dependencies: + self._apply_mixin(dependency) + mixin.apply(self) + self._applied_ids.add(id(mixin)) + self._mixins.append(mixin) + + def add_before_request_hook(self, hook: BeforeRequestHook) -> None: + self._before_request_hooks.append(hook) + + def add_after_request_hook(self, hook: AfterRequestHook) -> None: + self._after_request_hooks.append(hook) + + def add_before_websocket_hook(self, hook: BeforeWebSocketHook) -> None: + self._before_websocket_hooks.append(hook) + + def add_after_websocket_hook(self, hook: AfterWebSocketHook) -> None: + self._after_websocket_hooks.append(hook) async def handle_request(self, ctx: HttpContext) -> None: - result = self._tree.get_handler(ctx.path, ctx.method) - if result is not None: - handler, captured = result - await handler.handle_request(ctx, captured) - else: - await ctx.send_empty(404) + for hook in self._before_request_hooks: + new_ctx = await hook(ctx) + if new_ctx is not None: + ctx = new_ctx + try: + result = self._tree.get_handler(ctx.path, ctx.method) + if result is not None: + handler, captured = result + await handler.handle_request(ctx, captured) + else: + await ctx.send_empty(404) + finally: + for hook in reversed(self._after_request_hooks): + await hook(ctx) async def handle_websocket(self, ws: WebSocket) -> None: - result = self._tree.get_handler(ws.path, HttpMethod.WS) - if result is not None: - handler, captured = result - await handler.handle_request(ws, captured) - else: - await ws.close(1000) + for hook in self._before_websocket_hooks: + new_ws = await hook(ws) + if new_ws is not None: + ws = new_ws + try: + result = self._tree.get_handler(ws.path, HttpMethod.WS) + if result is not None: + handler, captured = result + await handler.handle_request(ws, captured) + else: + await ws.close(1000) + finally: + for hook in reversed(self._after_websocket_hooks): + await hook(ws) + + def setup(self, loop: AbstractEventLoop) -> None: + for mixin in self._mixins: + mixin.setup(loop) + + def shutdown(self, loop: AbstractEventLoop) -> None: + for mixin in self._mixins: + mixin.shutdown(loop) def route(self, paths: StrOrStrings, diff --git a/packages/kaya-core/src/kaya/core/_mixin.py b/packages/kaya-core/src/kaya/core/_mixin.py new file mode 100644 index 0000000..fd9444b --- /dev/null +++ b/packages/kaya-core/src/kaya/core/_mixin.py @@ -0,0 +1,37 @@ +from abc import ABC, abstractmethod +from asyncio import AbstractEventLoop +from typing import TYPE_CHECKING, Sequence + +if TYPE_CHECKING: + from ._app import KayaApp + + +class KayaMixin(ABC): + """Base class for composable Kaya app extensions. + + A mixin modifies a ``KayaApp`` instance in place by registering routes, + adding request/websocket hooks, or exposing helper methods on the mixin + instance itself. Because the app remains a ``KayaApp``, both ASGI and RSGI + protocols keep working regardless of which mixins are applied. + + Mixins may declare other mixins they depend on via ``dependencies``. The + ``KayaApp`` constructor applies dependencies first and guarantees each + mixin is applied at most once. + """ + + @property + def dependencies(self) -> Sequence['KayaMixin']: + return () + + @abstractmethod + def apply(self, app: 'KayaApp') -> None: + """Configure the app: register routes, add hooks, etc.""" + pass + + def setup(self, loop: AbstractEventLoop) -> None: + """Called on lifespan startup (default: no-op).""" + pass + + def shutdown(self, loop: AbstractEventLoop) -> None: + """Called on lifespan shutdown (default: no-op).""" + pass diff --git a/packages/kaya-oidc/README.md b/packages/kaya-oidc/README.md index bfa5ca5..8544178 100644 --- a/packages/kaya-oidc/README.md +++ b/packages/kaya-oidc/README.md @@ -10,34 +10,36 @@ Flow with PKCE**. ```python import os from kaya.core import HttpContext, KayaApp -from kaya.session import SessionMiddleware, InMemorySessionStore -from kaya.oidc import OIDCConfig, OIDCApp +from kaya.session import SessionMixin, InMemorySessionStore +from kaya.oidc import OIDCConfig, OIDCMixin -app = KayaApp() -session_app = SessionMiddleware(app, InMemorySessionStore()) - -oidc = OIDCApp( - session_app, +session = SessionMixin(InMemorySessionStore()) +oidc = OIDCMixin( OIDCConfig( issuer=os.environ['OIDC_ISSUER'], client_id=os.environ['OIDC_CLIENT_ID'], client_secret=os.environ.get('OIDC_CLIENT_SECRET'), redirect_uri='http://localhost:8000/auth/callback', fetch_userinfo=True, - ) + ), + session=session, ) +app = KayaApp(mixins=[session, oidc]) -@oidc.GET('/') +@app.GET('/') async def home(ctx: HttpContext): await ctx.send_str(200, 'public home') -@oidc.GET('/profile') +@app.GET('/profile') @oidc.require_auth async def profile(ctx: HttpContext): user = oidc.get_user(ctx) await ctx.send_str(200, f'Hello {user.email or user.sub}') ``` +`OIDCMixin` depends on `SessionMixin`; passing only `oidc` to `KayaApp(mixins=...)` +also works because the app applies mixin dependencies automatically. + ## Features - Generic OIDC discovery @@ -48,6 +50,7 @@ async def profile(ctx: HttpContext): - Optional userinfo endpoint fetch - Refresh token support - RP-initiated logout (when provider advertises `end_session_endpoint`) +- Composable with any other `KayaMixin` (RSGI, MCP, etc.) ## Security notes diff --git a/packages/kaya-oidc/src/kaya/oidc/__init__.py b/packages/kaya-oidc/src/kaya/oidc/__init__.py index 329e888..36d88a3 100644 --- a/packages/kaya-oidc/src/kaya/oidc/__init__.py +++ b/packages/kaya-oidc/src/kaya/oidc/__init__.py @@ -1,11 +1,11 @@ -from ._app import OIDCApp, OIDCUser from ._client import OIDCClient from ._config import OIDCConfig +from ._mixin import OIDCMixin, OIDCUser __all__ = [ - 'OIDCApp', 'OIDCClient', 'OIDCConfig', + 'OIDCMixin', 'OIDCUser', ] diff --git a/packages/kaya-oidc/src/kaya/oidc/_app.py b/packages/kaya-oidc/src/kaya/oidc/_mixin.py similarity index 67% rename from packages/kaya-oidc/src/kaya/oidc/_app.py rename to packages/kaya-oidc/src/kaya/oidc/_mixin.py index cec8ec1..198b24d 100644 --- a/packages/kaya-oidc/src/kaya/oidc/_app.py +++ b/packages/kaya-oidc/src/kaya/oidc/_mixin.py @@ -1,21 +1,14 @@ -from typing import Any, Awaitable, Callable, Mapping, MutableMapping, Optional, Sequence, cast -from urllib.parse import parse_qs, urlencode +from typing import Any, Awaitable, Callable, Mapping, Optional, Sequence -from kaya.core import HttpContext, HttpMethod, KayaApp -from kaya.session import Session, SessionMiddleware +from kaya.core import HttpContext, KayaApp, KayaMixin +from kaya.session import Session, SessionMixin +from urllib.parse import parse_qs from ._client import OIDCClient from ._config import OIDCConfig type HttpHandler = Callable[..., Awaitable[None]] -type WebSocketHandler = Callable[..., Awaitable[None]] -type RouteDecorator = Callable[[HttpHandler], HttpHandler] -type WebSocketDecorator = Callable[[WebSocketHandler], WebSocketHandler] -type ASGIApp = Callable[ - [MutableMapping[str, Any], Callable[[], Awaitable[Any]], Callable[[MutableMapping[str, Any]], Awaitable[None]]], - Awaitable[None], -] class OIDCUser(Mapping[str, Any]): @@ -50,42 +43,52 @@ class OIDCUser(Mapping[str, Any]): return self._data.get('picture') -class OIDCApp: - """ASGI app wrapper that adds OIDC authentication routes to a Kaya app. +class OIDCMixin(KayaMixin): + """Kaya mixin adding OpenID Connect authentication. - The wrapped app must be a ``SessionMiddleware`` instance so that OIDC state, - nonce, and user data can be stored in ``ctx.session``. + Depends on :class:`~kaya.session.SessionMixin` so that OIDC state, nonce, + and user data can be stored in ``ctx.session``. The dependency is applied + automatically by ``KayaApp``. - Built-in routes: + Registers three routes on the app: - ``login_path`` (default ``/auth/login``): redirects to the OIDC provider. - ``callback_path`` (default ``/auth/callback``): handles the provider callback. - ``logout_path`` (default ``/auth/logout``): logs the user out. + + Example:: + + session = SessionMixin(InMemorySessionStore()) + oidc = OIDCMixin(config, session=session) + app = KayaApp(mixins=[session, oidc]) + + @app.GET('/profile') + @oidc.require_auth + async def profile(ctx: HttpContext): + user = oidc.get_user(ctx) + ... """ - def __init__(self, app: SessionMiddleware, config: OIDCConfig) -> None: - self._app = app + def __init__(self, config: OIDCConfig, session: SessionMixin) -> None: self._config = config + self._session = session self._client = OIDCClient(config) - self._register_routes() - @staticmethod - def _session(ctx: HttpContext) -> Session: - session = ctx.session - assert isinstance(session, Session) - return session + @property + def dependencies(self) -> Sequence[KayaMixin]: + return [self._session] - def _register_routes(self) -> None: - @self._app.GET(self._config.login_path) + def apply(self, app: KayaApp) -> None: + @app.GET(self._config.login_path) async def login(ctx: HttpContext) -> None: auth_url, state, nonce, code_verifier = await self._client.build_authorization_url() - session = self._session(ctx) + session = self._session_of(ctx) session['oidc_state'] = state session['oidc_nonce'] = nonce session['oidc_code_verifier'] = code_verifier await ctx.send_empty(302, {'Location': auth_url}) - @self._app.GET(self._config.callback_path) + @app.GET(self._config.callback_path) async def callback(ctx: HttpContext) -> None: query = parse_qs(ctx.query_string) code = self._first_value(query.get('code')) @@ -104,7 +107,7 @@ class OIDCApp: await ctx.send_str(400, 'Missing code or state') return - session = self._session(ctx) + session = self._session_of(ctx) expected_state = session.get('oidc_state') if state != expected_state: await ctx.send_str(400, 'Invalid state') @@ -144,66 +147,32 @@ class OIDCApp: except ValueError as exc: await ctx.send_str(400, f'Authentication failed: {exc}') - @self._app.GET(self._config.logout_path) + @app.GET(self._config.logout_path) async def logout(ctx: HttpContext) -> None: - session = self._session(ctx) + session = self._session_of(ctx) id_token = session.get('oidc_id_token') session.invalidate() logout_url = await self._client.build_logout_url(id_token if isinstance(id_token, str) else None) location = logout_url if logout_url is not None else self._config.post_logout_redirect await ctx.send_empty(302, {'Location': location}) + @staticmethod + def _session_of(ctx: HttpContext) -> Session: + session = ctx.session + assert isinstance(session, Session) + return session + @staticmethod def _first_value(values: Optional[Sequence[str]]) -> Optional[str]: if values and len(values) > 0: return values[0] return None - def route( - self, - paths: str | Sequence[str], - methods: Optional[HttpMethod | Sequence[HttpMethod]] = None, - recursive: bool = False, - ) -> RouteDecorator: - return self._app.route(paths, methods, recursive) - - def GET(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.GET(path, recursive) - - def POST(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.POST(path, recursive) - - def PUT(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.PUT(path, recursive) - - def DELETE(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.DELETE(path, recursive) - - def OPTIONS(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.OPTIONS(path, recursive) - - def HEAD(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.HEAD(path, recursive) - - def PATCH(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.PATCH(path, recursive) - - def websocket(self, path: str, recursive: bool = False) -> WebSocketDecorator: - return self._app.websocket(path, recursive) - - async def __call__( - self, - scope: MutableMapping[str, Any], - receive: Callable[[], Awaitable[Any]], - send: Callable[[MutableMapping[str, Any]], Awaitable[None]], - ) -> None: - await cast(ASGIApp, self._app)(scope, receive, send) - def is_authenticated(self, ctx: HttpContext) -> bool: - return 'oidc_user' in self._session(ctx) + return 'oidc_user' in self._session_of(ctx) def get_user(self, ctx: HttpContext) -> Optional[OIDCUser]: - user = self._session(ctx).get('oidc_user') + user = self._session_of(ctx).get('oidc_user') if user is None or not isinstance(user, Mapping): return None return OIDCUser(user) diff --git a/packages/kaya-oidc/tests/test_oidc.py b/packages/kaya-oidc/tests/test_oidc.py index 12749aa..0c1c0aa 100644 --- a/packages/kaya-oidc/tests/test_oidc.py +++ b/packages/kaya-oidc/tests/test_oidc.py @@ -1,19 +1,17 @@ import base64 -import json import unittest from time import time -from typing import Any, Mapping +from typing import Mapping from urllib.parse import parse_qs, urlparse import httpx import jwt -from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import rsa from pwo import async_test from kaya.core import HttpContext, KayaApp -from kaya.session import InMemorySessionStore, SessionMiddleware +from kaya.session import InMemorySessionStore, SessionMixin -from kaya.oidc import OIDCApp, OIDCConfig +from kaya.oidc import OIDCConfig, OIDCMixin from kaya.oidc._client import OIDCClient @@ -197,14 +195,13 @@ class OIDCClientTest(unittest.TestCase): self.assertEqual(['/'], query['post_logout_redirect_uri']) -class OIDCAppTest(unittest.TestCase): +class OIDCMixinTest(unittest.TestCase): - def _build_app(self, fetch_userinfo: bool = False) -> tuple[OIDCApp, MockOIDCProvider, InMemorySessionStore]: + def _build_app(self, fetch_userinfo: bool = False) -> tuple[KayaApp, OIDCMixin, MockOIDCProvider, InMemorySessionStore]: provider = MockOIDCProvider() http_client = httpx.AsyncClient(transport=MockTransport(provider)) store = InMemorySessionStore() - app = KayaApp() - session_app = SessionMiddleware(app, store) + session = SessionMixin(store) config = OIDCConfig( issuer=provider.issuer, client_id='client', @@ -213,34 +210,35 @@ class OIDCAppTest(unittest.TestCase): http_client=http_client, fetch_userinfo=fetch_userinfo, ) - oidc_app = OIDCApp(session_app, config) - return oidc_app, provider, store + oidc = OIDCMixin(config, session=session) + app = KayaApp(mixins=[session, oidc]) + return app, oidc, provider, store - def _setup_routes(self, oidc_app: OIDCApp) -> None: - @oidc_app.GET('/') + def _setup_routes(self, app: KayaApp, oidc: OIDCMixin) -> None: + @app.GET('/') async def home(ctx: HttpContext) -> None: await ctx.send_str(200, 'home') - @oidc_app.GET('/profile') - @oidc_app.require_auth + @app.GET('/profile') + @oidc.require_auth async def profile(ctx: HttpContext) -> None: - user = oidc_app.get_user(ctx) + user = oidc.get_user(ctx) if user is None: await ctx.send_empty(401) return await ctx.send_str(200, f'Hello {user.email}') - @oidc_app.GET('/refresh') - @oidc_app.require_auth + @app.GET('/refresh') + @oidc.require_auth async def refresh(ctx: HttpContext) -> None: - new_token = await oidc_app.refresh_access_token(ctx.session) + new_token = await oidc.refresh_access_token(ctx.session) await ctx.send_str(200, new_token or 'no-token') @async_test async def test_login_redirect(self) -> None: - oidc_app, provider, store = self._build_app() - self._setup_routes(oidc_app) - transport = httpx.ASGITransport(app=oidc_app) + app, oidc, provider, store = self._build_app() + self._setup_routes(app, oidc) + transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/auth/login', follow_redirects=False) self.assertEqual(302, r.status_code) @@ -250,9 +248,9 @@ class OIDCAppTest(unittest.TestCase): @async_test async def test_callback_success(self) -> None: - oidc_app, provider, store = self._build_app(fetch_userinfo=True) - self._setup_routes(oidc_app) - transport = httpx.ASGITransport(app=oidc_app) + app, oidc, provider, store = self._build_app(fetch_userinfo=True) + self._setup_routes(app, oidc) + transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/auth/login', follow_redirects=False) self.assertEqual(302, r.status_code) @@ -274,18 +272,18 @@ class OIDCAppTest(unittest.TestCase): @async_test async def test_callback_invalid_state(self) -> None: - oidc_app, provider, store = self._build_app() - self._setup_routes(oidc_app) - transport = httpx.ASGITransport(app=oidc_app) + app, oidc, provider, store = self._build_app() + self._setup_routes(app, oidc) + transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/auth/callback', params={'code': 'mock-code', 'state': 'wrong'}, follow_redirects=False) self.assertEqual(400, r.status_code) @async_test async def test_logout(self) -> None: - oidc_app, provider, store = self._build_app() - self._setup_routes(oidc_app) - transport = httpx.ASGITransport(app=oidc_app) + app, oidc, provider, store = self._build_app() + self._setup_routes(app, oidc) + transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/auth/login', follow_redirects=False) parsed = urlparse(r.headers['Location']) @@ -304,9 +302,9 @@ class OIDCAppTest(unittest.TestCase): @async_test async def test_require_auth_redirect(self) -> None: - oidc_app, provider, store = self._build_app() - self._setup_routes(oidc_app) - transport = httpx.ASGITransport(app=oidc_app) + app, oidc, provider, store = self._build_app() + self._setup_routes(app, oidc) + transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/profile', follow_redirects=False) self.assertEqual(302, r.status_code) @@ -314,9 +312,9 @@ class OIDCAppTest(unittest.TestCase): @async_test async def test_refresh_access_token(self) -> None: - oidc_app, provider, store = self._build_app() - self._setup_routes(oidc_app) - transport = httpx.ASGITransport(app=oidc_app) + app, oidc, provider, store = self._build_app() + self._setup_routes(app, oidc) + transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/auth/login', follow_redirects=False) parsed = urlparse(r.headers['Location']) @@ -329,6 +327,34 @@ class OIDCAppTest(unittest.TestCase): self.assertEqual(200, r.status_code) self.assertEqual('new-access-token', r.text) + @async_test + async def test_dependency_applied_automatically(self) -> None: + # Only pass oidc to KayaApp; SessionMixin should be applied via dependencies. + provider = MockOIDCProvider() + http_client = httpx.AsyncClient(transport=MockTransport(provider)) + store = InMemorySessionStore() + session = SessionMixin(store) + config = OIDCConfig( + issuer=provider.issuer, + client_id='client', + client_secret='secret', + redirect_uri='http://localhost:8000/auth/callback', + http_client=http_client, + ) + oidc = OIDCMixin(config, session=session) + app = KayaApp(mixins=[oidc]) + + @app.GET('/') + async def home(ctx: HttpContext) -> None: + ctx.session['x'] = 1 + await ctx.send_str(200, 'ok') + + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: + r = await client.get('/') + self.assertEqual(200, r.status_code) + self.assertIn('Set-Cookie', r.headers) + if __name__ == '__main__': unittest.main() diff --git a/packages/kaya-session/README.md b/packages/kaya-session/README.md index 6ee507c..215ee84 100644 --- a/packages/kaya-session/README.md +++ b/packages/kaya-session/README.md @@ -9,12 +9,12 @@ session data is accessible from request handlers as `ctx.session`. ```python from kaya.core import KayaApp, HttpContext -from kaya.session import SessionMiddleware, InMemorySessionStore +from kaya.session import SessionMixin, InMemorySessionStore -app = KayaApp() -session_app = SessionMiddleware(app, InMemorySessionStore()) +session = SessionMixin(InMemorySessionStore()) +app = KayaApp(mixins=[session]) -@session_app.GET('/') +@app.GET('/') async def home(ctx: HttpContext): n = ctx.session.get('visits', 0) + 1 ctx.session['visits'] = n @@ -24,11 +24,14 @@ async def home(ctx: HttpContext): Sessions are created lazily: a cookie is only set when the handler modifies the session. +`SessionMixin` is a `KayaMixin`, so the app stays a `KayaApp` and both ASGI and +RSGI keep working. + ## Session expiry The cookie sent to the browser has a `Max-Age` (default 14 days), but that is only a client-side hint. The real boundary is the store's server-side TTL, -which the middleware keeps in sync with the cookie `Max-Age`. +which the mixin keeps in sync with the cookie `Max-Age`. For `InMemorySessionStore`, a session expires if it is idle for longer than `max_age`. Active sessions have their expiry slid forward on every access, so @@ -44,14 +47,12 @@ attribute) entirely. - `Session`: dict-like session object with modification tracking - `SessionStore`: abstract store interface - `InMemorySessionStore`: simple in-memory store for development/single-process -- `SessionMiddleware`: ASGI middleware managing session cookies and persistence +- `SessionMixin`: composable Kaya mixin managing session cookies and persistence - Session ID regeneration (`session.regenerate_id()`) and invalidation - (`session.invalidate()`) for future authentication layers + (`session.invalidate()`) for authentication layers ## Notes -- This release supports HTTP requests only; WebSocket and RSGI propagation is - planned for future releases. - `InMemorySessionStore` does not survive process restarts and is not shared across processes. Production deployments should use a store backed by a shared storage system (planned). diff --git a/packages/kaya-session/src/kaya/session/__init__.py b/packages/kaya-session/src/kaya/session/__init__.py index 2191e47..21a729c 100644 --- a/packages/kaya-session/src/kaya/session/__init__.py +++ b/packages/kaya-session/src/kaya/session/__init__.py @@ -1,4 +1,4 @@ -from ._middleware import SessionMiddleware +from ._mixin import SessionMixin from ._session import Session from ._store import InMemorySessionStore, SessionStore @@ -6,6 +6,6 @@ from ._store import InMemorySessionStore, SessionStore __all__ = [ 'InMemorySessionStore', 'Session', - 'SessionMiddleware', + 'SessionMixin', 'SessionStore', ] diff --git a/packages/kaya-session/src/kaya/session/_middleware.py b/packages/kaya-session/src/kaya/session/_middleware.py deleted file mode 100644 index a00fa8b..0000000 --- a/packages/kaya-session/src/kaya/session/_middleware.py +++ /dev/null @@ -1,178 +0,0 @@ -from typing import Any, Awaitable, Callable, MutableMapping, Optional, Sequence, cast - -from kaya.core import HttpMethod, KayaApp - -from ._cookie import format_set_cookie, parse_cookie_value -from ._session import Session -from ._store import SessionStore - - -type HttpHandler = Callable[..., Awaitable[None]] -type WebSocketHandler = Callable[..., Awaitable[None]] -type RouteDecorator = Callable[[HttpHandler], HttpHandler] -type WebSocketDecorator = Callable[[WebSocketHandler], WebSocketHandler] -type ASGIApp = Callable[ - [MutableMapping[str, Any], Callable[[], Awaitable[Any]], Callable[[MutableMapping[str, Any]], Awaitable[None]]], - Awaitable[None], -] - - -class SessionMiddleware: - """ASGI middleware providing server-side HTTP sessions. - - The middleware reads a session cookie from each HTTP request, loads the - session data via the configured store, and makes the session available to - Kaya handlers as ``ctx.session``. After the request it persists the session - and refreshes the cookie. - - Routing methods are delegated to the wrapped ``KayaApp`` so the middleware - can be used as a drop-in replacement when registering handlers. - - Example:: - - app = KayaApp() - session_app = SessionMiddleware(app, InMemorySessionStore()) - - @session_app.GET('/') - async def home(ctx: HttpContext): - ctx.session['visits'] = ctx.session.get('visits', 0) + 1 - await ctx.send_str(200, f"visits: {ctx.session['visits']}") - """ - - def __init__( - self, - app: KayaApp, - store: SessionStore, - cookie_name: str = 'session_id', - path: str = '/', - max_age: Optional[int] = 14 * 24 * 60 * 60, - httponly: bool = True, - secure: bool = False, - samesite: Optional[str] = 'Lax', - ) -> None: - self._app = app - self._store = store - self._cookie_name = cookie_name - self._path = path - self._max_age = max_age - self._httponly = httponly - self._secure = secure - self._samesite = samesite - - def route( - self, - paths: str | Sequence[str], - methods: Optional[HttpMethod | Sequence[HttpMethod]] = None, - recursive: bool = False, - ) -> RouteDecorator: - return self._app.route(paths, methods, recursive) - - def GET(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.GET(path, recursive) - - def POST(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.POST(path, recursive) - - def PUT(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.PUT(path, recursive) - - def DELETE(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.DELETE(path, recursive) - - def OPTIONS(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.OPTIONS(path, recursive) - - def HEAD(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.HEAD(path, recursive) - - def PATCH(self, path: str, recursive: bool = False) -> RouteDecorator: - return self._app.PATCH(path, recursive) - - def websocket(self, path: str, recursive: bool = False) -> WebSocketDecorator: - return self._app.websocket(path, recursive) - - async def __call__( - self, - scope: MutableMapping[str, Any], - receive: Callable[[], Awaitable[Any]], - send: Callable[[MutableMapping[str, Any]], Awaitable[None]], - ) -> None: - if scope['type'] != 'http': - await cast(ASGIApp, self._app)(scope, receive, send) - return - - session_id = self._extract_session_id(scope) - session: Session - if session_id is not None: - loaded = await self._store.load(session_id, self._max_age) - session = loaded if loaded is not None else Session() - else: - session = Session() - - state = scope.get('state') - if not isinstance(state, dict): - state = {} - scope['state'] = state - state['kaya_session'] = session - - async def wrapped_send(message: MutableMapping[str, Any]) -> None: - if message['type'] == 'http.response.start': - message = dict(message) - final_session_id = self._finalize_session_id(session) - if final_session_id is not None: - cookie_value = format_set_cookie( - self._cookie_name, - final_session_id, - path=self._path, - max_age=0 if session.invalidated else self._max_age, - httponly=self._httponly, - secure=self._secure, - samesite=self._samesite, - ) - headers = list(message.get('headers', [])) - headers.append((b'Set-Cookie', cookie_value.encode())) - message['headers'] = headers - await send(message) - - try: - await cast(ASGIApp, self._app)(scope, receive, wrapped_send) - finally: - await self._persist(session) - - def _extract_session_id(self, scope: MutableMapping[str, Any]) -> Optional[str]: - headers = scope.get('headers', []) - for key, value in headers: - key_bytes: bytes = key if isinstance(key, bytes) else key.encode() - if key_bytes.lower() == b'cookie': - value_str: str = value.decode() if isinstance(value, bytes) else value - return parse_cookie_value(value_str, self._cookie_name) - return None - - def _finalize_session_id(self, session: Session) -> Optional[str]: - if session.invalidated: - return session.id - if session.id is None: - if session.modified or session.regenerate: - session.set_id(self._store.new_session_id()) - elif session.regenerate: - session._old_id = session.id - session.set_id(self._store.new_session_id()) - session._regenerate = False - return session.id - - async def _persist(self, session: Session) -> None: - if session.invalidated: - old_id = session._old_id or session.id - if old_id is not None: - await self._store.delete(old_id) - return - - if session._old_id is not None and session._old_id != session.id: - await self._store.delete(session._old_id) - session._old_id = None - - if session.id is None and session.modified: - session.set_id(self._store.new_session_id()) - - if session.id is not None: - await self._store.save(session.id, session, self._max_age) diff --git a/packages/kaya-session/src/kaya/session/_mixin.py b/packages/kaya-session/src/kaya/session/_mixin.py new file mode 100644 index 0000000..1b6d818 --- /dev/null +++ b/packages/kaya-session/src/kaya/session/_mixin.py @@ -0,0 +1,200 @@ +from pathlib import Path +from typing import Any, AsyncGenerator, AsyncIterator, Awaitable, Callable, Mapping, Optional, Sequence + +from kaya.core import HttpContext, KayaApp, KayaMixin +from kaya.core._types import StrOrStrings + +from ._cookie import format_set_cookie, parse_cookie_value +from ._session import Session +from ._store import SessionStore + + +class SessionHttpContext(HttpContext): + """HttpContext wrapper that exposes ``session`` and injects the session + cookie into response headers. + + Works with any concrete ``HttpContext`` (ASGI or RSGI) because it only + relies on the abstract send methods, which all implementations share. + """ + + def __init__( + self, + ctx: HttpContext, + session: Session, + cookie_injector: Callable[[], Optional[str]], + ) -> None: + self._ctx = ctx + self.session = session + self._cookie_injector = cookie_injector + self.pathsend = ctx.pathsend + self.receive = ctx.receive + self.send = ctx.send + self.scheme = ctx.scheme + self.method = ctx.method + self.path = ctx.path + self.query_string = ctx.query_string + self.headers = ctx.headers + self.client = ctx.client + self.server = ctx.server + self.request_body = ctx.request_body + + def _inject_cookie(self, headers: Optional[Mapping[str, StrOrStrings]]) -> Optional[Mapping[str, StrOrStrings]]: + cookie_value = self._cookie_injector() + if cookie_value is None: + return headers + new_headers: dict[str, StrOrStrings] = dict(headers) if headers else {} + existing = new_headers.get('Set-Cookie') + if existing is None: + new_headers['Set-Cookie'] = cookie_value + elif isinstance(existing, str): + new_headers['Set-Cookie'] = (existing, cookie_value) + else: + new_headers['Set-Cookie'] = (*existing, cookie_value) + return new_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._inject_cookie(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._inject_cookie(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._inject_cookie(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._inject_cookie(headers)) + + async def send_empty(self, status: int, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + await self._ctx.send_empty(status, self._inject_cookie(headers)) + + +class _CookieInjector: + """Computes the Set-Cookie value once (on first response) and caches it.""" + + def __init__(self, mixin: 'SessionMixin', session: Session) -> None: + self._mixin = mixin + self._session = session + self._value: Optional[str] = None + self._computed = False + + def __call__(self) -> Optional[str]: + if not self._computed: + self._value = self._mixin._compute_cookie(self._session) + self._computed = True + return self._value + + +class SessionMixin(KayaMixin): + """Kaya mixin providing server-side HTTP sessions. + + Registers before/after request hooks that load and persist the session and + injects the session cookie into responses via a wrapped ``HttpContext``. + Because the app stays a ``KayaApp``, both ASGI and RSGI keep working. + + Example:: + + session = SessionMixin(InMemorySessionStore()) + app = KayaApp(mixins=[session]) + + @app.GET('/') + async def home(ctx: HttpContext): + ctx.session['visits'] = ctx.session.get('visits', 0) + 1 + await ctx.send_str(200, f"visits: {ctx.session['visits']}") + """ + + def __init__( + self, + store: SessionStore, + cookie_name: str = 'session_id', + path: str = '/', + max_age: Optional[int] = 14 * 24 * 60 * 60, + httponly: bool = True, + secure: bool = False, + samesite: Optional[str] = 'Lax', + ) -> None: + self._store = store + self._cookie_name = cookie_name + self._path = path + self._max_age = max_age + self._httponly = httponly + self._secure = secure + self._samesite = samesite + + def apply(self, app: KayaApp) -> None: + app.add_before_request_hook(self._before_request) + app.add_after_request_hook(self._after_request) + + async def _before_request(self, ctx: HttpContext) -> Optional[HttpContext]: + session_id = self._extract_session_id(ctx) + session: Session + if session_id is not None: + loaded = await self._store.load(session_id, self._max_age) + session = loaded if loaded is not None else Session() + else: + session = Session() + injector = _CookieInjector(self, session) + return SessionHttpContext(ctx, session, injector) + + async def _after_request(self, ctx: HttpContext) -> None: + session = ctx.session + if not isinstance(session, Session): + return + await self._persist(session) + + def _extract_session_id(self, ctx: HttpContext) -> Optional[str]: + cookie_header_values = ctx.headers.get('cookie') + if cookie_header_values is None: + return None + if isinstance(cookie_header_values, str): + return parse_cookie_value(cookie_header_values, self._cookie_name) + for value in cookie_header_values: + found = parse_cookie_value(value, self._cookie_name) + if found is not None: + return found + return None + + def _compute_cookie(self, session: Session) -> Optional[str]: + final_session_id = self._finalize_session_id(session) + if final_session_id is None: + return None + return format_set_cookie( + self._cookie_name, + final_session_id, + path=self._path, + max_age=0 if session.invalidated else self._max_age, + httponly=self._httponly, + secure=self._secure, + samesite=self._samesite, + ) + + def _finalize_session_id(self, session: Session) -> Optional[str]: + if session.invalidated: + return session.id + if session.id is None: + if session.modified or session.regenerate: + session.set_id(self._store.new_session_id()) + elif session.regenerate: + session._old_id = session.id + session.set_id(self._store.new_session_id()) + session._regenerate = False + return session.id + + async def _persist(self, session: Session) -> None: + if session.invalidated: + old_id = session._old_id or session.id + if old_id is not None: + await self._store.delete(old_id) + return + + if session._old_id is not None and session._old_id != session.id: + await self._store.delete(session._old_id) + session._old_id = None + + if session.id is None and session.modified: + session.set_id(self._store.new_session_id()) + + if session.id is not None: + await self._store.save(session.id, session, self._max_age) diff --git a/packages/kaya-session/tests/test_session.py b/packages/kaya-session/tests/test_session.py index aab1244..e760158 100644 --- a/packages/kaya-session/tests/test_session.py +++ b/packages/kaya-session/tests/test_session.py @@ -6,7 +6,7 @@ import httpx from pwo import async_test from kaya.core import KayaApp, HttpContext -from kaya.session import InMemorySessionStore, Session, SessionMiddleware, SessionStore +from kaya.session import InMemorySessionStore, Session, SessionMixin, SessionStore from kaya.session._cookie import format_set_cookie, parse_cookie_value @@ -24,42 +24,40 @@ class FakeClock: class SessionTest(unittest.TestCase): app: KayaApp store: InMemorySessionStore - session_app: SessionMiddleware def setUp(self) -> None: - self.app = KayaApp() self.store = InMemorySessionStore() - self.session_app = SessionMiddleware(self.app, self.store) + self.app = KayaApp(mixins=[SessionMixin(self.store)]) - @self.session_app.GET('/') + @self.app.GET('/') async def home(ctx: HttpContext) -> None: n = ctx.session.get('visits', 0) + 1 ctx.session['visits'] = n await ctx.send_str(200, f'visits: {n}') - @self.session_app.GET('/read') + @self.app.GET('/read') async def read(ctx: HttpContext) -> None: n = ctx.session.get('visits', 0) await ctx.send_str(200, f'visits: {n}') - @self.session_app.GET('/write') + @self.app.GET('/write') async def write(ctx: HttpContext) -> None: ctx.session['foo'] = 'bar' await ctx.send_str(200, 'ok') - @self.session_app.GET('/clear') + @self.app.GET('/clear') async def clear(ctx: HttpContext) -> None: ctx.session.invalidate() await ctx.send_str(200, 'cleared') - @self.session_app.GET('/rotate') + @self.app.GET('/rotate') async def rotate(ctx: HttpContext) -> None: ctx.session.regenerate_id() await ctx.send_str(200, 'rotated') @async_test async def test_session_persists_across_requests(self) -> None: - transport = httpx.ASGITransport(app=self.session_app) + transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/') self.assertEqual(200, r.status_code) @@ -72,7 +70,7 @@ class SessionTest(unittest.TestCase): @async_test async def test_no_cookie_when_session_not_modified(self) -> None: - transport = httpx.ASGITransport(app=self.session_app) + transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/read') self.assertEqual(200, r.status_code) @@ -81,7 +79,7 @@ class SessionTest(unittest.TestCase): @async_test async def test_existing_session_refreshes_cookie(self) -> None: - transport = httpx.ASGITransport(app=self.session_app) + transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: await client.get('/') r = await client.get('/read') @@ -91,7 +89,7 @@ class SessionTest(unittest.TestCase): @async_test async def test_cookie_attributes(self) -> None: - transport = httpx.ASGITransport(app=self.session_app) + transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/') set_cookie = r.headers['Set-Cookie'] @@ -102,7 +100,7 @@ class SessionTest(unittest.TestCase): @async_test async def test_sessions_are_isolated(self) -> None: - transport = httpx.ASGITransport(app=self.session_app) + transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r1 = await client.get('/') r2 = await client.get('/') @@ -115,7 +113,7 @@ class SessionTest(unittest.TestCase): @async_test async def test_invalidate(self) -> None: - transport = httpx.ASGITransport(app=self.session_app) + transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: await client.get('/') r = await client.get('/clear') @@ -129,7 +127,7 @@ class SessionTest(unittest.TestCase): @async_test async def test_regenerate_id(self) -> None: - transport = httpx.ASGITransport(app=self.session_app) + transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: await client.get('/') cookies_before = {c.name: c.value for c in client.cookies.jar} @@ -153,7 +151,7 @@ class SessionTest(unittest.TestCase): @async_test async def test_invalid_cookie_creates_fresh_session(self) -> None: - transport = httpx.ASGITransport(app=self.session_app) + transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: client.cookies.set('session_id', 'not-a-real-id') r = await client.get('/') @@ -164,19 +162,18 @@ class SessionTest(unittest.TestCase): async def test_stale_cookie_cannot_access_old_data(self) -> None: clock = FakeClock() store = InMemorySessionStore(clock=clock) - app = KayaApp() - session_app = SessionMiddleware(app, store, max_age=60) + app = KayaApp(mixins=[SessionMixin(store, max_age=60)]) - @session_app.GET('/') + @app.GET('/') async def home(ctx: HttpContext) -> None: ctx.session['secret'] = 'super-sensitive' await ctx.send_str(200, 'ok') - @session_app.GET('/read') + @app.GET('/read') async def read(ctx: HttpContext) -> None: await ctx.send_str(200, ctx.session.get('secret', 'none')) - transport = httpx.ASGITransport(app=session_app) + transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/') self.assertEqual('ok', r.text)