246 lines
9.5 KiB
Python
246 lines
9.5 KiB
Python
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, SessionMiddleware, SessionStore
|
|
from kaya.session._cookie import format_set_cookie, parse_cookie_value
|
|
|
|
|
|
class SessionTest(unittest.TestCase):
|
|
app: KayaApp
|
|
store: InMemorySessionStore
|
|
session_app: SessionMiddleware
|
|
|
|
def setUp(self) -> None:
|
|
self.app = KayaApp()
|
|
self.store = InMemorySessionStore()
|
|
self.session_app = SessionMiddleware(self.app, self.store)
|
|
|
|
@self.session_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.session_app.GET('/read')
|
|
async def read(ctx: HttpContext) -> None:
|
|
n = ctx.session.get('visits', 0)
|
|
await ctx.send_str(200, f'visits: {n}')
|
|
|
|
@self.session_app.GET('/write')
|
|
async def write(ctx: HttpContext) -> None:
|
|
ctx.session['foo'] = 'bar'
|
|
await ctx.send_str(200, 'ok')
|
|
|
|
@self.session_app.GET('/clear')
|
|
async def clear(ctx: HttpContext) -> None:
|
|
ctx.session.invalidate()
|
|
await ctx.send_str(200, 'cleared')
|
|
|
|
@self.session_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.session_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.session_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.session_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.session_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.session_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.session_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.session_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._sessions)
|
|
|
|
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._sessions)
|
|
self.assertIn(new_id, self.store._sessions)
|
|
|
|
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.session_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)
|
|
|
|
|
|
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._sessions.get(session_id))
|
|
|
|
import asyncio
|
|
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)))
|
|
|
|
|
|
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()
|