From e4e00762bbf3fcbe3c0c0ea955f05c318461569d Mon Sep 17 00:00:00 2001 From: Walter Oggioni Date: Sun, 19 Jul 2026 05:51:16 +0000 Subject: [PATCH] Add websocket session support to kaya-session MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - kaya-core: WebSocket ABC gains session attribute and accept(headers=...) - kaya-core: AsgiWebSocket injects headers into websocket.accept message - kaya-rsgi: RsgiWebSocket accepts headers param (ignored — Granian's accept() takes no args) - kaya-session: SessionWebSocket wrapper exposes ws.session and injects Set-Cookie on accept() - kaya-session: SessionMixin registers before/after websocket hooks; session loaded at connect, persisted on close if modified - 10 new WV session tests covering read, persist, handshake cookie, regenerate, invalidate, isolation - Example and README updated --- example/session.py | 10 +- packages/kaya-core/src/kaya/core/_asgi.py | 9 +- .../kaya-core/src/kaya/core/_websocket.py | 7 +- packages/kaya-core/tests/test_websocket.py | 53 +++++ packages/kaya-rsgi/src/kaya/rsgi/_rsgi.py | 4 +- packages/kaya-session/README.md | 23 +++ .../kaya-session/src/kaya/session/_mixin.py | 127 +++++++++--- .../tests/test_session_websocket.py | 194 ++++++++++++++++++ 8 files changed, 390 insertions(+), 37 deletions(-) create mode 100644 packages/kaya-session/tests/test_session_websocket.py diff --git a/example/session.py b/example/session.py index 7849558..58176d9 100644 --- a/example/session.py +++ b/example/session.py @@ -1,4 +1,4 @@ -from kaya.core import HttpContext, KayaApp +from kaya.core import HttpContext, KayaApp, WebSocket from kaya.session import InMemorySessionStore, SessionMixin app = KayaApp(mixins=[SessionMixin(InMemorySessionStore())]) @@ -35,11 +35,7 @@ async def echo(ws: WebSocket) -> None: @app.websocket('/ws/visits') async def ws_visits(ws: WebSocket) -> None: - # WebSocket handlers can read the existing session. Most ASGI servers - # (including Granian and Daphne) do not forward the `headers` field of the - # `websocket.accept` message into the HTTP 101 response, so a new session - # cookie cannot be set during the handshake. Use the HTTP `/` endpoint to - # set or refresh the session cookie before connecting here. + visits = ws.session.get('visits', 0) + 1 + ws.session['visits'] = visits await ws.accept() - visits = ws.session.get('visits', 0) await ws.send_text(f'visits: {visits}') diff --git a/packages/kaya-core/src/kaya/core/_asgi.py b/packages/kaya-core/src/kaya/core/_asgi.py index 25f9ba0..09ac05a 100644 --- a/packages/kaya-core/src/kaya/core/_asgi.py +++ b/packages/kaya-core/src/kaya/core/_asgi.py @@ -173,8 +173,13 @@ class AsgiWebSocket(WebSocket): self.server = scope['server'] self.headers = decode_headers(scope['headers']) - async def accept(self) -> None: - await self._send({'type': 'websocket.accept'}) + async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + message: Dict[str, Any] = {'type': 'websocket.accept'} + if headers is not None: + # Emit a list rather than a tuple: wsproto's AcceptConnection (used + # by httpx_ws and others) requires list concatenation. + message['headers'] = list(encode_headers(headers)) + await self._send(message) async def receive(self) -> WebSocketMessage: message = await self._receive() diff --git a/packages/kaya-core/src/kaya/core/_websocket.py b/packages/kaya-core/src/kaya/core/_websocket.py index 1398f52..a8041e5 100644 --- a/packages/kaya-core/src/kaya/core/_websocket.py +++ b/packages/kaya-core/src/kaya/core/_websocket.py @@ -1,6 +1,8 @@ from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import AsyncIterator, Literal, Mapping, Optional, Sequence, Tuple, Union +from typing import Any, AsyncIterator, Literal, Mapping, Optional, Sequence, Tuple, Union + +from ._types.base import StrOrStrings type WebSocketData = Optional[Union[str, bytes, int]] @@ -18,9 +20,10 @@ class WebSocket(ABC): headers: Mapping[str, Sequence[str]] client: Optional[Tuple[str, int]] server: Optional[Tuple[str, Optional[int]]] + session: Optional[Any] = None @abstractmethod - async def accept(self) -> None: + async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: pass @abstractmethod diff --git a/packages/kaya-core/tests/test_websocket.py b/packages/kaya-core/tests/test_websocket.py index ae2523b..08536ca 100644 --- a/packages/kaya-core/tests/test_websocket.py +++ b/packages/kaya-core/tests/test_websocket.py @@ -4,6 +4,7 @@ from pwo import async_test from httpx_ws import aconnect_ws, WebSocketDisconnect from httpx_ws.transport import ASGIWebSocketTransport from kaya.core import KayaApp, WebSocket +from kaya.core._asgi import AsgiWebSocket class WebSocketTest(unittest.TestCase): @@ -78,3 +79,55 @@ class WebSocketTest(unittest.TestCase): async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as client: async with aconnect_ws("/echo", client) as ws: pass + + @async_test + async def test_accept_with_headers(self): + sent_messages = [] + + async def send(message): + sent_messages.append(message) + + async def receive(): + return {'type': 'websocket.connect'} + + scope = { + 'type': 'websocket', + 'path': '/echo', + 'query_string': b'', + 'scheme': 'ws', + 'client': ('127.0.0.1', 12345), + 'server': ('127.0.0.1', 80), + 'headers': [], + } + ws = AsgiWebSocket(scope, receive, send) + await ws.accept(headers={'Set-Cookie': 'sid=abc; Path=/', 'X-Custom': ('a', 'b')}) + self.assertEqual(1, len(sent_messages)) + message = sent_messages[0] + self.assertEqual('websocket.accept', message['type']) + self.assertIn((b'Set-Cookie', b'sid=abc; Path=/'), message['headers']) + self.assertIn((b'X-Custom', b'a'), message['headers']) + self.assertIn((b'X-Custom', b'b'), message['headers']) + + @async_test + async def test_accept_without_headers(self): + sent_messages = [] + + async def send(message): + sent_messages.append(message) + + async def receive(): + return {'type': 'websocket.connect'} + + scope = { + 'type': 'websocket', + 'path': '/echo', + 'query_string': b'', + 'scheme': 'ws', + 'client': ('127.0.0.1', 12345), + 'server': ('127.0.0.1', 80), + 'headers': [], + } + ws = AsgiWebSocket(scope, receive, send) + await ws.accept() + self.assertEqual(1, len(sent_messages)) + self.assertEqual({'type': 'websocket.accept'}, sent_messages[0]) diff --git a/packages/kaya-rsgi/src/kaya/rsgi/_rsgi.py b/packages/kaya-rsgi/src/kaya/rsgi/_rsgi.py index 2595fdd..774cbe6 100644 --- a/packages/kaya-rsgi/src/kaya/rsgi/_rsgi.py +++ b/packages/kaya-rsgi/src/kaya/rsgi/_rsgi.py @@ -141,7 +141,9 @@ class RsgiWebSocket(WebSocket): .map(lambda it: (it[0], int(it[1]))) .or_else_throw(RuntimeError)) - async def accept(self) -> None: + async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + # RSGI's websocket accept() takes no arguments (https://github.com/emmett-framework/granian/blob/master/docs/spec/RSGI.md), + # so handshake headers (e.g. Set-Cookie) cannot be sent on RSGI; they are silently ignored here. self._transport = await self._protocol.accept() async def receive(self) -> WebSocketMessage: diff --git a/packages/kaya-session/README.md b/packages/kaya-session/README.md index 215ee84..bdcf3bc 100644 --- a/packages/kaya-session/README.md +++ b/packages/kaya-session/README.md @@ -27,6 +27,27 @@ session. `SessionMixin` is a `KayaMixin`, so the app stays a `KayaApp` and both ASGI and RSGI keep working. +## WebSocket sessions + +The same session is available in websocket handlers as `ws.session`: + +```python +@app.websocket('/ws/visits') +async def ws_visits(ws: WebSocket): + visits = ws.session.get('visits', 0) + 1 + ws.session['visits'] = visits + await ws.accept() + await ws.send_text(f'visits: {visits}') +``` + +The session is loaded from the cookie when the connection is opened and +persisted when the connection closes, if it was modified. The session cookie +can only be set or refreshed on the handshake response, so mutate the session +*before* calling `ws.accept()` if you want the cookie delivered with the +handshake. Handshake cookies require ASGI spec version 2.1+; RSGI websocket +handshakes cannot carry response headers, so on RSGI the session is loaded and +persisted but the cookie is only set or refreshed by HTTP responses. + ## Session expiry The cookie sent to the browser has a `Max-Age` (default 14 days), but that is @@ -50,6 +71,8 @@ attribute) entirely. - `SessionMixin`: composable Kaya mixin managing session cookies and persistence - Session ID regeneration (`session.regenerate_id()`) and invalidation (`session.invalidate()`) for authentication layers +- WebSocket support: the session is exposed as `ws.session` in websocket + handlers, loaded at connect time and persisted on close ## Notes diff --git a/packages/kaya-session/src/kaya/session/_mixin.py b/packages/kaya-session/src/kaya/session/_mixin.py index 87f04e9..0f14abb 100644 --- a/packages/kaya-session/src/kaya/session/_mixin.py +++ b/packages/kaya-session/src/kaya/session/_mixin.py @@ -1,7 +1,7 @@ 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 import HttpContext, KayaApp, KayaMixin, WebSocket from kaya.core._types import StrOrStrings from ._cookie import format_set_cookie, parse_cookie_value @@ -9,6 +9,23 @@ from ._session import Session from ._store import SessionStore +def _merge_cookie_header( + headers: Optional[Mapping[str, StrOrStrings]], + cookie_value: Optional[str], +) -> Optional[Mapping[str, StrOrStrings]]: + 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 + + class SessionHttpContext(HttpContext): """HttpContext wrapper that exposes ``session`` and injects the session cookie into response headers. @@ -37,18 +54,7 @@ class SessionHttpContext(HttpContext): return getattr(self._ctx, name) 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 + return _merge_cookie_header(headers, self._cookie_injector()) async def stream_body(self, status: int, @@ -69,6 +75,55 @@ class SessionHttpContext(HttpContext): await self._ctx.send_empty(status, self._inject_cookie(headers)) +class SessionWebSocket(WebSocket): + """WebSocket wrapper that exposes ``session`` and injects the session + cookie into the handshake response headers on ``accept()``. + + Works with any concrete ``WebSocket`` (ASGI or RSGI) because it only + relies on the abstract methods, which all implementations share. + Attributes not explicitly overridden are delegated to the wrapped socket + via ``__getattr__``. + + The cookie is only sent if the underlying transport supports handshake + response headers: ASGI does (spec version 2.1+), RSGI does not, so on + RSGI the session is still loaded and persisted but no cookie is set or + refreshed from a websocket connection. + """ + + def __init__( + self, + ws: WebSocket, + session: Session, + cookie_injector: Callable[[], Optional[str]], + ) -> None: + object.__setattr__(self, '_ws', ws) + object.__setattr__(self, 'session', session) + object.__setattr__(self, '_cookie_injector', cookie_injector) + + def __getattr__(self, name: str) -> Any: + if name == '_ws': + raise AttributeError(name) + return getattr(self._ws, name) + + async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None: + await self._ws.accept(_merge_cookie_header(headers, self._cookie_injector())) + + async def receive(self) -> Any: + return await self._ws.receive() + + async def send_text(self, data: str) -> None: + await self._ws.send_text(data) + + async def send_bytes(self, data: bytes) -> None: + await self._ws.send_bytes(data) + + async def close(self, code: int = 1000) -> None: + await self._ws.close(code) + + async def __anext__(self) -> Any: + return await self._ws.__anext__() + + class _CookieInjector: """Computes the Set-Cookie value once (on first response) and caches it.""" @@ -88,9 +143,16 @@ class _CookieInjector: 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. + Registers before/after request and websocket hooks that load and persist + the session, injecting the session cookie into HTTP responses via a + wrapped ``HttpContext`` and into websocket handshake responses via a + wrapped ``WebSocket``. Because the app stays a ``KayaApp``, both ASGI and + RSGI keep working. + + For websockets the session is loaded when the connection is opened and + persisted when it closes if modified. The session cookie can only be set + or refreshed on the handshake response (ASGI only; RSGI websocket + handshakes cannot carry response headers). Example:: @@ -124,15 +186,11 @@ class SessionMixin(KayaMixin): def apply(self, app: KayaApp) -> None: app.add_before_request_hook(self._before_request) app.add_after_request_hook(self._after_request) + app.add_before_websocket_hook(self._before_websocket) + app.add_after_websocket_hook(self._after_websocket) 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() + session = await self._load_session(ctx.headers) injector = _CookieInjector(self, session) return SessionHttpContext(ctx, session, injector) @@ -142,8 +200,27 @@ class SessionMixin(KayaMixin): return await self._persist(session) - def _extract_session_id(self, ctx: HttpContext) -> Optional[str]: - cookie_header_values = ctx.headers.get('cookie') + async def _before_websocket(self, ws: WebSocket) -> Optional[WebSocket]: + session = await self._load_session(ws.headers) + injector = _CookieInjector(self, session) + return SessionWebSocket(ws, session, injector) + + async def _after_websocket(self, ws: WebSocket) -> None: + session = ws.session + if not isinstance(session, Session): + return + await self._persist(session) + + async def _load_session(self, headers: Mapping[str, Sequence[str]]) -> Session: + session_id = self._extract_session_id(headers) + if session_id is not None: + loaded = await self._store.load(session_id, self._max_age) + if loaded is not None: + return loaded + return Session() + + def _extract_session_id(self, headers: Mapping[str, Sequence[str]]) -> Optional[str]: + cookie_header_values = headers.get('cookie') if cookie_header_values is None: return None if isinstance(cookie_header_values, str): diff --git a/packages/kaya-session/tests/test_session_websocket.py b/packages/kaya-session/tests/test_session_websocket.py new file mode 100644 index 0000000..eff69b4 --- /dev/null +++ b/packages/kaya-session/tests/test_session_websocket.py @@ -0,0 +1,194 @@ +import unittest + +import httpx +from pwo import async_test +from httpx_ws import aconnect_ws +from httpx_ws.transport import ASGIWebSocketTransport +from kaya.core import KayaApp, HttpContext, WebSocket + +from kaya.session import InMemorySessionStore, SessionMixin + + +class WebSocketSessionTest(unittest.TestCase): + app: KayaApp + store: InMemorySessionStore + + def setUp(self) -> None: + self.store = InMemorySessionStore() + self.app = KayaApp(mixins=[SessionMixin(self.store)]) + + @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.app.GET('/read') + async def read(ctx: HttpContext) -> None: + await ctx.send_str(200, f"visits: {ctx.session.get('visits', 0)}" + f" ws_seen: {ctx.session.get('ws_seen', False)}") + + @self.app.websocket('/visits') + async def ws_visits(ws: WebSocket) -> None: + await ws.accept() + await ws.send_text(f"visits: {ws.session.get('visits', 0)}") + + @self.app.websocket('/mark') + async def ws_mark(ws: WebSocket) -> None: + await ws.accept() + ws.session['ws_seen'] = True + await ws.send_text('marked') + + @self.app.websocket('/handshake-write') + async def ws_handshake_write(ws: WebSocket) -> None: + ws.session['ws_seen'] = True + await ws.accept() + await ws.send_text('marked') + + @self.app.websocket('/peek') + async def ws_peek(ws: WebSocket) -> None: + await ws.accept() + await ws.send_text(f"visits: {ws.session.get('visits', 0)}") + + @self.app.websocket('/rotate') + async def ws_rotate(ws: WebSocket) -> None: + old_id = ws.session.id + ws.session.regenerate_id() + await ws.accept() + await ws.send_text(f'old: {old_id} new: {ws.session.id}') + + @self.app.websocket('/clear') + async def ws_clear(ws: WebSocket) -> None: + ws.session.invalidate() + await ws.accept() + await ws.send_text('cleared') + + @async_test + async def test_ws_session_loaded_from_cookie(self) -> None: + transport = ASGIWebSocketTransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + r = await client.get('/') + self.assertEqual('visits: 1', r.text) + + async with aconnect_ws('/visits', client) as ws: + message = await ws.receive_text() + self.assertEqual('visits: 1', message) + + @async_test + async def test_ws_session_persisted_on_close(self) -> None: + transport = ASGIWebSocketTransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + await client.get('/') + async with aconnect_ws('/mark', client) as ws: + message = await ws.receive_text() + self.assertEqual('marked', message) + + r = await client.get('/read') + self.assertEqual('visits: 1 ws_seen: True', r.text) + + @async_test + async def test_ws_handshake_sets_cookie(self) -> None: + transport = ASGIWebSocketTransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + self.assertNotIn('session_id', client.cookies) + async with aconnect_ws('/handshake-write', client) as ws: + message = await ws.receive_text() + self.assertEqual('marked', message) + + self.assertIn('session_id', client.cookies) + self.assertIn(client.cookies['session_id'], self.store._data) + + @async_test + async def test_ws_no_cookie_when_session_not_modified(self) -> None: + transport = ASGIWebSocketTransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + async with aconnect_ws('/peek', client) as ws: + message = await ws.receive_text() + self.assertEqual('visits: 0', message) + + self.assertNotIn('session_id', client.cookies) + + @async_test + async def test_ws_session_unmodified_not_persisted(self) -> None: + transport = ASGIWebSocketTransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + async with aconnect_ws('/peek', client) as ws: + message = await ws.receive_text() + self.assertEqual('visits: 0', message) + + self.assertEqual(0, len(self.store._data)) + + @async_test + async def test_ws_session_new_session_saved_on_close(self) -> None: + transport = ASGIWebSocketTransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + self.assertNotIn('session_id', client.cookies) + async with aconnect_ws('/mark', client) as ws: + message = await ws.receive_text() + self.assertEqual('marked', message) + + self.assertEqual(1, len(self.store._data)) + + @async_test + async def test_ws_session_regenerate_id(self) -> None: + transport = ASGIWebSocketTransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + await client.get('/') + old_id = client.cookies['session_id'] + self.assertIn(old_id, self.store._data) + + new_id = None + async with aconnect_ws('/rotate', client) as ws: + message = await ws.receive_text() + old_part, new_part = message.split(' new: ') + self.assertEqual(f'old: {old_id}', old_part) + new_id = new_part + self.assertNotEqual(old_id, new_id) + + self.assertIsNotNone(new_id) + self.assertNotIn(old_id, self.store._data) + self.assertIn(new_id, self.store._data) + + @async_test + async def test_ws_session_invalidate(self) -> None: + transport = ASGIWebSocketTransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + await client.get('/') + old_id = client.cookies['session_id'] + self.assertIn(old_id, self.store._data) + + async with aconnect_ws('/clear', client) as ws: + message = await ws.receive_text() + self.assertEqual('cleared', message) + + self.assertNotIn(old_id, self.store._data) + + @async_test + async def test_ws_sessions_are_isolated(self) -> None: + transport = ASGIWebSocketTransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + await client.get('/') + async with aconnect_ws('/visits', client) as ws: + message = await ws.receive_text() + self.assertEqual('visits: 1', message) + + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + async with aconnect_ws('/visits', client) as ws: + message = await ws.receive_text() + self.assertEqual('visits: 0', message) + + @async_test + async def test_ws_invalid_cookie_creates_fresh_session(self) -> None: + transport = ASGIWebSocketTransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url='http://testserver') as client: + client.cookies.set('session_id', 'not-a-real-id') + async with aconnect_ws('/mark', client) as ws: + message = await ws.receive_text() + self.assertEqual('marked', message) + + self.assertNotIn('not-a-real-id', self.store._data) + self.assertEqual(1, len(self.store._data)) + + +if __name__ == '__main__': + unittest.main()