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