Files
kaya/packages/kaya-session/tests/test_session_websocket.py
T
woggioni e4e00762bb 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
2026-07-23 22:11:04 +08:00

195 lines
7.9 KiB
Python

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()