Files
kefu/deploy/token-usage-20260917/payload/wechat_rpa/test_model_usage_recovery.py
T
2026-09-21 10:34:06 +08:00

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)