import asyncio import hashlib import unittest from datetime import timedelta from typing import Any import httpx from kaya_rbcs.app import create_app, key_from_path from kaya_rbcs.config import Config from kaya_rbcs.store import MemcacheStore from fake_memcached import FakeMemcachedServer class RbcsTest(unittest.IsolatedAsyncioTestCase): async def asyncSetUp(self) -> None: self.memcache = FakeMemcachedServer() await self.memcache.start() self.app = create_app(self._config()) self.store: MemcacheStore = getattr(self.app, 'store') self.client = httpx.AsyncClient( transport=httpx.ASGITransport(app=self.app), base_url='http://testserver', ) async def asyncTearDown(self) -> None: await self.client.aclose() await self.store.close() await self.memcache.stop() def _config(self, **overrides: Any) -> Config: params: dict[str, Any] = { 'host': '127.0.0.1', 'port': 8080, 'path_prefix': '/', 'memcache_host': '127.0.0.1', 'memcache_port': self.memcache.port, } params.update(overrides) return Config(**params) async def test_put_get_roundtrip(self) -> None: key = 'abc123' value = b'hello world' put = await self.client.put( '/' + key, content=value, headers={'content-type': 'application/octet-stream'}, ) self.assertEqual(201, put.status_code) self.assertEqual(key, put.text) self.assertEqual('text/plain', put.headers['content-type']) get = await self.client.get('/' + key) self.assertEqual(200, get.status_code) self.assertEqual(value, get.content) self.assertEqual('application/octet-stream', get.headers['content-type']) async def test_get_missing_key(self) -> None: get = await self.client.get('/does/not/exist') self.assertEqual(404, get.status_code) self.assertEqual(b'', get.content) async def test_nested_key(self) -> None: value = b'nested value' put = await self.client.put('/a/b/c', content=value) self.assertEqual(201, put.status_code) get = await self.client.get('/a/b/c') self.assertEqual(200, get.status_code) self.assertEqual(value, get.content) async def test_default_content_type(self) -> None: await self.client.put('/no-type', content=b'x') get = await self.client.get('/no-type') self.assertEqual('application/octet-stream', get.headers['content-type']) async def test_content_disposition_roundtrip(self) -> None: disposition = 'inline; filename="page.html"' await self.client.put( '/page', content=b'', headers={'content-type': 'text/html', 'content-disposition': disposition}, ) get = await self.client.get('/page') self.assertEqual(200, get.status_code) self.assertEqual('text/html', get.headers['content-type']) self.assertEqual(disposition, get.headers['content-disposition']) async def test_key_prefix_appended(self) -> None: app = create_app(self._config(key_prefix='suffix')) store: MemcacheStore = getattr(app, 'store') client = httpx.AsyncClient( transport=httpx.ASGITransport(app=app), base_url='http://testserver', ) try: put = await client.put('/key1', content=b'x') self.assertEqual(201, put.status_code) self.assertIn(b'key1suffix', self.memcache.data) get = await client.get('/key1') self.assertEqual(200, get.status_code) self.assertEqual(b'x', get.content) finally: await client.aclose() await store.close() async def test_digest_hashes_key(self) -> None: app = create_app(self._config(digest='sha256')) store: MemcacheStore = getattr(app, 'store') client = httpx.AsyncClient( transport=httpx.ASGITransport(app=app), base_url='http://testserver', ) try: put = await client.put('/key1', content=b'x') self.assertEqual(201, put.status_code) expected = hashlib.sha256(b'key1').hexdigest().encode('utf-8') self.assertIn(expected, self.memcache.data) get = await client.get('/key1') self.assertEqual(200, get.status_code) self.assertEqual(b'x', get.content) finally: await client.aclose() await store.close() async def test_path_prefix(self) -> None: app = create_app(self._config(path_prefix='/cache')) store: MemcacheStore = getattr(app, 'store') client = httpx.AsyncClient( transport=httpx.ASGITransport(app=app), base_url='http://testserver', ) try: await client.put('/cache/entry', content=b'value') get = await client.get('/cache/entry') self.assertEqual(200, get.status_code) self.assertEqual(b'value', get.content) self.assertIn(b'entry', self.memcache.data) finally: await client.aclose() await store.close() def test_key_from_path_rejects_escaping(self) -> None: self.assertIsNone(key_from_path('/cache/../../etc/passwd', '/cache')) self.assertIsNone(key_from_path('/cache/../', '/cache')) self.assertIsNone(key_from_path('/', '/')) self.assertEqual('a/b', key_from_path('/cache/a/b', '/cache')) self.assertEqual('foo', key_from_path('/foo', '/')) self.assertEqual('foo', key_from_path('/cache/foo', '/cache')) async def test_escaped_path_is_rejected(self) -> None: # ASGI servers (and httpx) normalize dot segments in the request path # before it reaches the app, so the route is never matched. Either way # the request is rejected; a raw, unnormalized escaping path would be # handled by key_from_path and rejected with 400. app = create_app(self._config(path_prefix='/cache')) store: MemcacheStore = getattr(app, 'store') client = httpx.AsyncClient( transport=httpx.ASGITransport(app=app), base_url='http://testserver', ) try: get = await client.get('/cache/../../etc/passwd') self.assertIn(get.status_code, (400, 404)) put = await client.put('/cache/../../etc/passwd', content=b'x') self.assertIn(put.status_code, (400, 404)) finally: await client.aclose() await store.close() async def test_max_age_expiry(self) -> None: app = create_app(self._config(max_age=timedelta(seconds=1))) store: MemcacheStore = getattr(app, 'store') client = httpx.AsyncClient( transport=httpx.ASGITransport(app=app), base_url='http://testserver', ) try: put = await client.put('/ttl', content=b'x') self.assertEqual(201, put.status_code) get = await client.get('/ttl') self.assertEqual(200, get.status_code) await asyncio.sleep(1.2) get = await client.get('/ttl') self.assertEqual(404, get.status_code) finally: await client.aclose() await store.close() async def test_overwrite_value(self) -> None: await self.client.put('/same', content=b'first') await self.client.put('/same', content=b'second') get = await self.client.get('/same') self.assertEqual(200, get.status_code) self.assertEqual(b'second', get.content) async def test_invalid_method_returns_404(self) -> None: post = await self.client.post('/anything', content=b'x') self.assertEqual(404, post.status_code) if __name__ == '__main__': unittest.main()