Replace wrapper-based SessionMiddleware/OIDCApp with KayaMixin subclasses applied via KayaApp(mixins=[...]). Mixins hook into handle_request and handle_websocket via before/after hooks, so both ASGI and RSGI keep working. Mixin dependencies are applied automatically and deduplicated.
332 lines
13 KiB
Python
332 lines
13 KiB
Python
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()
|