"""Accounting must survive locks and malformed/cancelled concurrent requests.""" import asyncio import json import tempfile import time from pathlib import Path from unittest import TestCase import httpx import admin_backend import model_gateway as gw from model_usage import UsageRecorder, recover_pending_usage import test_model_gateway as fixtures import test_model_usage_collection as collection class UsageRecoveryTest(TestCase): def test_db_lock_spools_quickly_and_replays_exactly_once_after_unlock(self): with tempfile.TemporaryDirectory() as directory: database = admin_backend.Database(Path(directory) / 'usage.db') database.migrate() lock = database.connect() lock.execute('BEGIN IMMEDIATE') async def run(): async with UsageRecorder(database, tenant_id='scope') as recorder: recorder.capture(fixtures._outlet(), {'usage': {'total_tokens': 23}}, attempt=1, role='answer', status='success', latency_ms=10) return 'normal reply' started = time.monotonic() try: with self.assertLogs('model_usage', level='WARNING'): self.assertEqual(asyncio.run(run()), 'normal reply') self.assertLess(time.monotonic() - started, 1.8) files = list((Path(directory) / 'model_usage_pending').glob('*.json')) self.assertEqual(len(files), 1) raw = files[0].read_bytes() finally: lock.rollback() lock.close() recover_pending_usage(database) self.assertFalse(files[0].exists()) self.assertEqual(database.model_usage_stats(1, 'scope')['summary']['total_tokens'], 23) # A crash after DB commit but before deleting the file is safe. files[0].write_bytes(raw) recover_pending_usage(database) self.assertEqual(database.model_usage_stats(1, 'scope')['summary']['request_count'], 1) def test_corrupt_spool_is_quarantined_without_blocking_valid_events(self): with tempfile.TemporaryDirectory() as directory: database = admin_backend.Database(Path(directory) / 'usage.db') database.migrate() folder = Path(directory) / 'model_usage_pending' folder.mkdir() bad = folder / '001-bad.json' bad.write_text('{incomplete', encoding='utf-8') recorder = UsageRecorder(database, tenant_id='scope') recorder.capture(fixtures._outlet(), {'usage': {'total_tokens': 27}}, attempt=1, role='answer', status='success', latency_ms=1) good = folder / '002-good.json' good.write_text(json.dumps({'version': 1, 'events': recorder.events}), encoding='utf-8') with self.assertLogs('model_usage', level='WARNING'): recover_pending_usage(database) self.assertFalse(bad.exists()) self.assertFalse(good.exists()) invalid = list(folder.glob('*.invalid-*')) self.assertEqual(len(invalid), 1) self.assertEqual(invalid[0].read_text(), '{incomplete') self.assertEqual(database.model_usage_stats(1, 'scope')['summary']['total_tokens'], 27) def test_cancelled_upstream_is_recorded_before_request_exit(self): with tempfile.TemporaryDirectory() as directory: database = admin_backend.Database(Path(directory) / 'usage.db') database.migrate() async def run(): entered = asyncio.Event() async def handler(request): entered.set() await asyncio.sleep(10) return httpx.Response(200, json={}) async def request(): async with UsageRecorder(database, tenant_id='scope') as recorder: async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: await gw.call_outlet(client, fixtures._outlet(), [], usage_recorder=recorder) task = asyncio.create_task(request()) await entered.wait() task.cancel() with self.assertRaises(asyncio.CancelledError): await task asyncio.run(run()) row = database.list_model_usage(1, tenant_id='scope')['items'][0] self.assertEqual(row['status'], 'error') self.assertIsNone(row['total_tokens']) class MalformedCandidateUsageTest(TestCase): setUp = collection.GatewayUsageTest.setUp tearDown = collection.GatewayUsageTest.tearDown _post = collection.GatewayUsageTest._post usage = collection.GatewayUsageTest.usage upstream = collection.GatewayUsageTest.upstream def test_bad_candidate_does_not_interrupt_slow_healthy_candidate_accounting(self): calls = 0 async def handler(request): nonlocal calls calls += 1 if calls == 1: return httpx.Response(200, json={'choices': [{'message': {'content': '', 'tool_calls': 1}}], 'usage': {'total_tokens': 10}}) await asyncio.sleep(0.1) return self.upstream(request) response = self._post(handler) self.assertEqual(response.status_code, 200) self.assertEqual(response.json()['reply'], 'synthetic reply') events = self.usage() self.assertEqual(events['total'], 3) self.assertEqual(sum(e['total_tokens'] for e in events['items']), 46) self.assertEqual(sum(e['status'] == 'error' for e in events['items']), 1)