- 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
195 lines
7.9 KiB
Python
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()
|