119 lines
5.6 KiB
Python
119 lines
5.6 KiB
Python
"""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)
|