import asyncio import unittest from typing import Any import httpx from pwo import async_test from kaya.core import KayaApp, HttpContext from kaya.session import InMemorySessionStore, Session, SessionMixin, SessionStore from kaya.session._cookie import format_set_cookie, parse_cookie_value class FakeClock: def __init__(self, start: float = 0.0) -> None: self._now = start def __call__(self) -> float: return self._now def advance(self, seconds: float) -> None: self._now += seconds class SessionTest(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: n = ctx.session.get('visits', 0) await ctx.send_str(200, f'visits: {n}') @self.app.GET('/write') async def write(ctx: HttpContext) -> None: ctx.session['foo'] = 'bar' await ctx.send_str(200, 'ok') @self.app.GET('/clear') async def clear(ctx: HttpContext) -> None: ctx.session.invalidate() await ctx.send_str(200, 'cleared') @self.app.GET('/rotate') async def rotate(ctx: HttpContext) -> None: ctx.session.regenerate_id() await ctx.send_str(200, 'rotated') @async_test async def test_session_persists_across_requests(self) -> None: transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/') self.assertEqual(200, r.status_code) self.assertEqual('visits: 1', r.text) self.assertIn('Set-Cookie', r.headers) r = await client.get('/') self.assertEqual(200, r.status_code) self.assertEqual('visits: 2', r.text) @async_test async def test_no_cookie_when_session_not_modified(self) -> None: transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/read') self.assertEqual(200, r.status_code) self.assertEqual('visits: 0', r.text) self.assertNotIn('Set-Cookie', r.headers) @async_test async def test_existing_session_refreshes_cookie(self) -> None: transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: await client.get('/') r = await client.get('/read') self.assertEqual(200, r.status_code) self.assertIn('Set-Cookie', r.headers) self.assertEqual('visits: 1', r.text) @async_test async def test_cookie_attributes(self) -> None: transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/') set_cookie = r.headers['Set-Cookie'] self.assertIn('HttpOnly', set_cookie) self.assertIn('SameSite=Lax', set_cookie) self.assertIn('Path=/', set_cookie) self.assertIn('Max-Age=', set_cookie) @async_test async def test_sessions_are_isolated(self) -> None: transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r1 = await client.get('/') r2 = await client.get('/') self.assertEqual('visits: 1', r1.text) self.assertEqual('visits: 2', r2.text) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/') self.assertEqual('visits: 1', r.text) @async_test async def test_invalidate(self) -> None: transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: await client.get('/') r = await client.get('/clear') self.assertEqual(200, r.status_code) self.assertEqual('cleared', r.text) set_cookie = r.headers['Set-Cookie'] self.assertIn('Max-Age=0', set_cookie) r = await client.get('/') self.assertEqual('visits: 1', r.text) @async_test async def test_regenerate_id(self) -> None: transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: await client.get('/') cookies_before = {c.name: c.value for c in client.cookies.jar} old_id = cookies_before.get('session_id') self.assertIsNotNone(old_id) self.assertIn(old_id, self.store._data) r = await client.get('/rotate') self.assertEqual(200, r.status_code) self.assertEqual('rotated', r.text) cookies_after = {c.name: c.value for c in client.cookies.jar} new_id = cookies_after.get('session_id') self.assertIsNotNone(new_id) self.assertNotEqual(old_id, new_id) self.assertNotIn(old_id, self.store._data) self.assertIn(new_id, self.store._data) r = await client.get('/read') self.assertEqual('visits: 1', r.text) @async_test async def test_invalid_cookie_creates_fresh_session(self) -> None: transport = httpx.ASGITransport(app=self.app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: client.cookies.set('session_id', 'not-a-real-id') r = await client.get('/') self.assertEqual('visits: 1', r.text) self.assertIn('Set-Cookie', r.headers) @async_test async def test_stale_cookie_cannot_access_old_data(self) -> None: clock = FakeClock() store = InMemorySessionStore(clock=clock) app = KayaApp(mixins=[SessionMixin(store, max_age=60)]) @app.GET('/') async def home(ctx: HttpContext) -> None: ctx.session['secret'] = 'super-sensitive' await ctx.send_str(200, 'ok') @app.GET('/read') async def read(ctx: HttpContext) -> None: await ctx.send_str(200, ctx.session.get('secret', 'none')) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url='http://127.0.0.1:80') as client: r = await client.get('/') self.assertEqual('ok', r.text) old_id = client.cookies['session_id'] self.assertIn(old_id, store._data) self.assertIn(old_id, store._expires) clock.advance(61) r = await client.get('/read') self.assertEqual('none', r.text) self.assertNotIn(old_id, store._data) self.assertNotIn(old_id, store._expires) class SessionUnitTest(unittest.TestCase): def test_session_is_dict_like(self) -> None: session = Session() session['a'] = 1 self.assertEqual(1, session['a']) self.assertTrue('a' in session) self.assertEqual({'a': 1}, dict(session)) self.assertTrue(session.modified) def test_session_modification_tracking(self) -> None: session = Session(data={'a': 1}) self.assertFalse(session.modified) session['a'] = 2 self.assertTrue(session.modified) def test_session_read_does_not_mark_modified(self) -> None: session = Session(data={'a': 1}) self.assertFalse(session.modified) _ = session['a'] self.assertFalse(session.modified) def test_session_clear_marks_modified(self) -> None: session = Session(data={'a': 1}) self.assertFalse(session.modified) session.clear() self.assertTrue(session.modified) self.assertEqual(0, len(session)) def test_session_invalidate(self) -> None: session = Session(session_id='abc', data={'a': 1}) self.assertFalse(session.modified) session.invalidate() self.assertTrue(session.invalidated) self.assertTrue(session.modified) self.assertEqual(0, len(session)) def test_session_regenerate_id(self) -> None: session = Session(session_id='abc', data={'a': 1}) session.regenerate_id() self.assertTrue(session.regenerate) self.assertIsNone(session.id) self.assertEqual('abc', session._old_id) self.assertTrue(session.modified) def test_session_store_is_abstract(self) -> None: with self.assertRaises(TypeError): SessionStore() # type: ignore[abstract] def test_in_memory_store_round_trip(self) -> None: store = InMemorySessionStore() session_id = store.new_session_id() session = Session(session_id, {'a': 1}) self.assertIsNone(store._data.get(session_id)) asyncio.run(store.save(session_id, session)) loaded = asyncio.run(store.load(session_id)) self.assertIsNotNone(loaded) assert loaded is not None self.assertEqual(1, loaded['a']) asyncio.run(store.delete(session_id)) self.assertIsNone(asyncio.run(store.load(session_id))) def test_in_memory_store_expires_after_max_age(self) -> None: clock = FakeClock() store = InMemorySessionStore(clock=clock) session_id = store.new_session_id() asyncio.run(store.save(session_id, Session(session_id, {'a': 1}), max_age=60)) clock.advance(61) self.assertIsNone(asyncio.run(store.load(session_id, max_age=60))) self.assertIsNone(store._data.get(session_id)) self.assertIsNone(store._expires.get(session_id)) def test_in_memory_store_slides_expiry_on_load(self) -> None: clock = FakeClock() store = InMemorySessionStore(clock=clock) session_id = store.new_session_id() asyncio.run(store.save(session_id, Session(session_id, {'a': 1}), max_age=60)) clock.advance(30) loaded = asyncio.run(store.load(session_id, max_age=60)) self.assertIsNotNone(loaded) assert loaded is not None self.assertEqual(1, loaded['a']) # Without sliding, the session would expire at t=60. With sliding it is now valid until t=90. clock.advance(35) loaded = asyncio.run(store.load(session_id, max_age=60)) self.assertIsNotNone(loaded) assert loaded is not None self.assertEqual(1, loaded['a']) def test_in_memory_store_no_expiry_without_max_age(self) -> None: clock = FakeClock() store = InMemorySessionStore(clock=clock) session_id = store.new_session_id() asyncio.run(store.save(session_id, Session(session_id, {'a': 1}))) clock.advance(1000000) loaded = asyncio.run(store.load(session_id)) self.assertIsNotNone(loaded) assert loaded is not None self.assertEqual(1, loaded['a']) def test_in_memory_store_invalidate_removes_expiry(self) -> None: clock = FakeClock() store = InMemorySessionStore(clock=clock) session_id = store.new_session_id() asyncio.run(store.save(session_id, Session(session_id, {'a': 1}), max_age=60)) asyncio.run(store.delete(session_id)) self.assertIsNone(store._data.get(session_id)) self.assertIsNone(store._expires.get(session_id)) class CookieUtilTest(unittest.TestCase): def test_parse_cookie_value(self) -> None: self.assertEqual('bar', parse_cookie_value('foo=bar; baz=qux', 'foo')) self.assertEqual('qux', parse_cookie_value('foo=bar; baz=qux', 'baz')) self.assertIsNone(parse_cookie_value('foo=bar', 'missing')) self.assertIsNone(parse_cookie_value('', 'foo')) def test_format_set_cookie(self) -> None: value = format_set_cookie('sid', 'abc123', path='/', max_age=3600, httponly=True, secure=True, samesite='Strict') self.assertIn('sid=abc123', value) self.assertIn('Path=/', value) self.assertIn('Max-Age=3600', value) self.assertIn('HttpOnly', value) self.assertIn('Secure', value) self.assertIn('SameSite=Strict', value) def test_format_set_cookie_without_secure(self) -> None: value = format_set_cookie('sid', 'abc123', max_age=3600) self.assertIn('sid=abc123', value) self.assertIn('HttpOnly', value) self.assertNotIn('Secure', value) self.assertIn('SameSite=Lax', value) if __name__ == '__main__': unittest.main()