import unittest from datetime import datetime, timezone from typing import Optional import httpx from pwo import async_test from kaya.core import KayaApp, HttpContext from kaya.session import Session, SessionMixin from kaya.session.memcache import MemcacheSessionStore class FakeClock: def __init__(self, start: float = 1_000_000.0) -> None: self._now = start def __call__(self) -> float: return self._now def advance(self, seconds: float) -> None: self._now += seconds class StubMemcacheClient: """In-memory stub implementing the aiomcache.Client subset used by the store.""" def __init__(self, clock: FakeClock) -> None: self._clock = clock self._data: dict[bytes, bytes] = {} self._expires: dict[bytes, float] = {} self.set_exptimes: list[int] = [] self.touch_exptimes: list[int] = [] def _expired(self, key: bytes) -> bool: expires = self._expires.get(key) return expires is not None and self._clock() > expires async def get(self, key: bytes) -> Optional[bytes]: if self._expired(key): self._data.pop(key, None) self._expires.pop(key, None) return None return self._data.get(key) async def set(self, key: bytes, value: bytes, exptime: int = 0) -> bool: self.set_exptimes.append(exptime) self._data[key] = value if exptime > 0: self._expires[key] = self._clock() + exptime else: self._expires.pop(key, None) return True async def delete(self, key: bytes) -> bool: existed = key in self._data self._data.pop(key, None) self._expires.pop(key, None) return existed async def touch(self, key: bytes, exptime: int) -> bool: self.touch_exptimes.append(exptime) if self._expired(key) or key not in self._data: return False self._expires[key] = self._clock() + exptime return True class MemcacheSessionStoreTest(unittest.TestCase): clock: FakeClock client: StubMemcacheClient store: MemcacheSessionStore def setUp(self) -> None: self.clock = FakeClock() self.client = StubMemcacheClient(self.clock) self.store = MemcacheSessionStore(self.client, clock=self.clock) # type: ignore[arg-type] @async_test async def test_save_and_load_round_trip(self) -> None: session = Session('abc', {'foo': 'bar', 'n': 42}) await self.store.save('abc', session) loaded = await self.store.load('abc') self.assertIsNotNone(loaded) assert loaded is not None self.assertEqual('abc', loaded.id) self.assertEqual({'foo': 'bar', 'n': 42}, dict(loaded)) @async_test async def test_load_unknown_session_returns_none(self) -> None: self.assertIsNone(await self.store.load('missing')) @async_test async def test_delete_removes_session(self) -> None: await self.store.save('abc', Session('abc', {'foo': 'bar'})) await self.store.delete('abc') self.assertIsNone(await self.store.load('abc')) @async_test async def test_save_with_max_age_expires(self) -> None: await self.store.save('abc', Session('abc', {'foo': 'bar'}), max_age=60) self.assertIsNotNone(await self.store.load('abc')) self.clock.advance(61) self.assertIsNone(await self.store.load('abc')) @async_test async def test_save_without_max_age_never_expires(self) -> None: await self.store.save('abc', Session('abc', {'foo': 'bar'})) self.assertEqual([0], self.client.set_exptimes) self.clock.advance(10_000_000) self.assertIsNotNone(await self.store.load('abc')) @async_test async def test_load_slides_expiry_when_max_age_given(self) -> None: await self.store.save('abc', Session('abc', {'foo': 'bar'}), max_age=60) self.clock.advance(50) loaded = await self.store.load('abc', max_age=60) self.assertIsNotNone(loaded) self.assertEqual([60], self.client.touch_exptimes) self.clock.advance(50) self.assertIsNotNone(await self.store.load('abc')) @async_test async def test_exptime_over_30_days_converted_to_absolute_timestamp(self) -> None: max_age = 40 * 24 * 60 * 60 await self.store.save('abc', Session('abc', {'foo': 'bar'}), max_age=max_age) self.assertEqual([int(self.clock()) + max_age], self.client.set_exptimes) @async_test async def test_exptime_exactly_30_days_stays_relative(self) -> None: max_age = 30 * 24 * 60 * 60 await self.store.save('abc', Session('abc', {'foo': 'bar'}), max_age=max_age) self.assertEqual([max_age], self.client.set_exptimes) @async_test async def test_touch_over_30_days_converted_to_absolute_timestamp(self) -> None: max_age = 40 * 24 * 60 * 60 await self.store.save('abc', Session('abc', {'foo': 'bar'}), max_age=max_age) self.clock.advance(100) await self.store.load('abc', max_age=max_age) self.assertEqual([int(self.clock()) + max_age], self.client.touch_exptimes) @async_test async def test_pickle_round_trip_of_non_json_values(self) -> None: now = datetime(2026, 7, 20, 12, 0, 0, tzinfo=timezone.utc) session = Session('abc', {'when': now, 'blob': b'\x00\x01', 'items': {1, 2, 3}}) await self.store.save('abc', session) loaded = await self.store.load('abc') self.assertIsNotNone(loaded) assert loaded is not None self.assertEqual(now, loaded['when']) self.assertEqual(b'\x00\x01', loaded['blob']) self.assertEqual({1, 2, 3}, loaded['items']) @async_test async def test_custom_prefix(self) -> None: store = MemcacheSessionStore(self.client, prefix='myapp:sess:', clock=self.clock) # type: ignore[arg-type] await store.save('abc', Session('abc', {'foo': 'bar'})) self.assertIn(b'myapp:sess:abc', self.client._data) self.assertNotIn(b'kaya:session:abc', self.client._data) @async_test async def test_custom_serializer(self) -> None: import json store = MemcacheSessionStore( self.client, # type: ignore[arg-type] dumps=lambda d: json.dumps(d).encode('utf-8'), loads=lambda b: json.loads(b.decode('utf-8')), clock=self.clock, ) await store.save('abc', Session('abc', {'foo': 'bar'})) self.assertEqual(b'{"foo": "bar"}', self.client._data[b'kaya:session:abc']) loaded = await store.load('abc') self.assertIsNotNone(loaded) assert loaded is not None self.assertEqual({'foo': 'bar'}, dict(loaded)) class MemcacheSessionIntegrationTest(unittest.TestCase): app: KayaApp def setUp(self) -> None: store = MemcacheSessionStore(StubMemcacheClient(FakeClock())) # type: ignore[arg-type] self.app = KayaApp(mixins=[SessionMixin(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}') @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)