Add websocket session support to kaya-session
- 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
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user