import asyncio import time import unittest from unittest.mock import AsyncMock, patch from aiohttp.test_utils import TestClient, TestServer from decision_server.server import PREFIX from decision_server.tests.test_service import plan_value, request_value from decision_server.tests.test_web_config import config from decision_server.web_config import COOKIE from decision_server.web_server import CATALOG, MANAGER, create_website_app from decision_server.web_sessions import Limits class WebsiteTests(unittest.IsolatedAsyncioTestCase): async def asyncSetUp(self): self.origin = "https://site.test" self.app = create_website_app(self.origin) self.client = TestClient(TestServer(self.app)) await self.client.start_server() self.manager = self.app[MANAGER] self.app[CATALOG].refresh = AsyncMock() self.headers = {"Host": "site.test", "Origin": self.origin} async def asyncTearDown(self): await self.client.close() async def visitor(self): response = await self.client.post(PREFIX + "/session", json={}, headers=self.headers) self.assertEqual(response.status, 200) value = await response.json() cookie = response.cookies[COOKIE] self.assertTrue(cookie["httponly"]) self.assertTrue(cookie["secure"]) self.assertEqual(cookie["samesite"], "Strict") self.assertEqual(cookie["domain"], "") return { **self.headers, "Cookie": COOKIE + "=" + cookie.value, "X-CSRF-Token": value["csrfToken"], "X-Config-Version": "0", }, self.manager.values[cookie.value] async def save(self, headers): response = await self.client.put(PREFIX + "/configuration", json=config(), headers=headers) self.assertEqual(response.status, 200, await response.text()) headers["X-Config-Version"] = response.headers["X-Config-Version"] return await response.json() async def test_boundary_no_cookie_csrf_origin_host_or_cross_site(self): response = await self.client.get(PREFIX + "/status", headers=self.headers) self.assertEqual(response.status, 401) headers, _ = await self.visitor() for patch_headers in ( {"Origin": "https://evil.test"}, {"Origin": ""}, {"Host": "evil.test"}, {"X-CSRF-Token": "bad"}, {"Sec-Fetch-Site": "same-site"}, ): response = await self.client.put( PREFIX + "/configuration", json=config(), headers={**headers, **patch_headers} ) self.assertEqual(response.status, 403) self.assertNotIn("Access-Control-Allow-Origin", response.headers) self.assertEqual(response.headers["Cache-Control"], "no-store") response = await self.client.post( PREFIX + "/session", json={}, headers={"Host": "site.test"} ) self.assertEqual(response.status, 403) async def test_atomic_credentials_no_metadata_files_and_stale_tab(self): a, av = await self.visitor() b, bv = await self.visitor() result = await self.save(a) self.assertNotIn("secret-fixture", str(result)) self.assertFalse(bv.service.connections.values) self.assertFalse(hasattr(av.service.connections, "path")) response = await self.client.put( PREFIX + "/configuration", json=config(), headers={**a, "X-Config-Version": "0"} ) self.assertEqual(response.status, 409) bad = config() bad["jev"]["apiKey"] = "invalid key" response = await self.client.put(PREFIX + "/configuration", json=bad, headers=a) self.assertEqual(response.status, 400) self.assertEqual(av.service.connections.values["jev"].key, "jev-secret-fixture") response = await self.client.get(PREFIX + "/status", headers=b) self.assertFalse((await response.json())["ready"]) async def test_two_visitors_identical_request_ids_and_cancel_isolation(self): a, av = await self.visitor() b, bv = await self.visitor() await self.save(a) await self.save(b) started = asyncio.Event() release = asyncio.Event() async def provider(*_): started.set() await release.wait() return plan_value(), {} with patch("decision_server.providers.openai.plan", side_effect=provider): pending = asyncio.create_task( self.client.post(PREFIX + "/plan", json=request_value(), headers=b) ) await asyncio.wait_for(started.wait(), 2) response = await self.client.post( PREFIX + "/cancel", json={"runId": "run1", "requestId": "r1"}, headers=a ) self.assertFalse((await response.json())["cancelled"]) self.assertEqual(len(bv.service.active), 1) await self.client.delete(PREFIX + "/session", headers=a) self.assertTrue(av.closed) self.assertTrue(bv.service.connections.values) release.set() self.assertEqual((await pending).status, 200) self.assertEqual(self.manager.inference, 0) async def test_ttl_status_does_not_refresh_and_credentials_destroyed(self): headers, visitor = await self.visitor() await self.save(headers) touched = visitor.touched await self.client.get(PREFIX + "/status", headers=headers) await self.client.post(PREFIX + "/session", json={}, headers=headers) self.assertEqual(visitor.touched, touched) visitor.touched = time.monotonic() - 1801 response = await self.client.get(PREFIX + "/status", headers=headers) self.assertEqual(response.status, 401) self.assertFalse(visitor.service.connections.values) self.assertTrue(visitor.closed) async def test_limits_ip_spoof_does_not_bypass_and_no_implicit_cli(self): self.manager.limits = Limits(ip_sessions=2, codex=1) a, av = await self.visitor() b, bv = await self.visitor() response = await self.client.post( PREFIX + "/session", json={}, headers={**self.headers, "X-Real-IP": "1.2.3.4"} ) self.assertEqual(response.status, 429) self.manager.reserve_codex(av) with self.assertRaisesRegex(Exception, "subscription_capacity"): self.manager.reserve_codex(bv) av.codex_reserved = False with patch.object(av.service.codex, "start", new_callable=AsyncMock) as start: response = await self.client.get(PREFIX + "/codex/status", headers=a) self.assertFalse((await response.json())["loggedIn"]) start.assert_not_called() await self.save(b) self.manager.inference = self.manager.limits.inference response = await self.client.post(PREFIX + "/plan", json=request_value(), headers=b) self.assertEqual(response.status, 429) self.manager.inference = 0 async def test_no_arbitrary_rpc_or_local_connection_or_queries(self): headers, _ = await self.visitor() for path in ("/connections", "/codex/exec", "/codex/rpc"): response = await self.client.post(PREFIX + path, json={}, headers=headers) self.assertIn(response.status, (404, 405)) response = await self.client.get(PREFIX + "/status?key=fixture", headers=headers) self.assertEqual(response.status, 400)