Files
kefu/wechat_rpa/test_model_usage_collection.py
T
2026-09-21 10:34:06 +08:00

209 lines
11 KiB
Python

"""Token accounting through real HTTP adapters with synthetic upstream responses."""
import asyncio
import json
import shutil
import sqlite3
import tempfile
from pathlib import Path
from unittest import TestCase, mock
import httpx
import admin_backend
import model_gateway as gw
import model_protocol as mp
from model_usage import UsageRecorder
import test_model_gateway as fixtures
class UsageParsingTest(TestCase):
def test_openai_details_are_subsets_not_extra_tokens(self):
parsed = mp.parse_usage('openai', {'usage': {'prompt_tokens': 100, 'completion_tokens': 30,
'total_tokens': 130, 'prompt_tokens_details': {'cached_tokens': 60},
'completion_tokens_details': {'reasoning_tokens': 20}}})
self.assertEqual(parsed, dict(input_tokens=100, output_tokens=30, total_tokens=130,
cached_input_tokens=60, reasoning_tokens=20))
def test_claude_includes_cache_reads_and_writes_exactly_once(self):
parsed = mp.parse_usage('claude', {'usage': {'input_tokens': 10, 'output_tokens': 20,
'cache_creation_input_tokens': 40, 'cache_read_input_tokens': 100,
'output_tokens_details': {'thinking_tokens': 5}}})
self.assertEqual(parsed, dict(input_tokens=150, output_tokens=20, total_tokens=170,
cached_input_tokens=100, reasoning_tokens=5))
def test_claude_without_cache_fields_still_counts_reported_tokens(self):
parsed = mp.parse_usage('claude', {'usage': {'input_tokens': 10, 'output_tokens': 20}})
self.assertEqual(parsed['total_tokens'], 30)
self.assertIsNone(parsed['cached_input_tokens'])
def test_dify_usage_in_metadata_and_compatible_input_output_names(self):
parsed = mp.parse_usage('dify', {'metadata': {'usage': {
'prompt_tokens': 12, 'completion_tokens': 8, 'total_tokens': 20}}})
self.assertEqual(parsed['total_tokens'], 20)
other = mp.parse_usage('openai', {'usage': {'input_tokens': 12, 'output_tokens': 8}})
self.assertEqual(parsed, other)
def test_missing_partial_and_zero_remain_distinct(self):
for response in (None, [], 'bad', {}, {'usage': []}, {'usage': {}}):
self.assertTrue(all(v is None for v in mp.parse_usage('openai', response).values()))
parsed = mp.parse_usage('openai', {'usage': {'total_tokens': 50}})
self.assertEqual(parsed['total_tokens'], 50)
self.assertIsNone(parsed['input_tokens'])
self.assertIsNone(parsed['output_tokens'])
self.assertIsNone(mp.parse_usage('openai', {'usage': {'prompt_tokens': 20}})['total_tokens'])
self.assertEqual(mp.parse_usage('openai', {'usage': {'prompt_tokens': 0, 'completion_tokens': 0}})['total_tokens'], 0)
def test_invalid_values_and_details_cannot_break_reply_processing(self):
for value in (-1, True, 2.3, 'NaN', {}, [], 10**30):
with self.subTest(value=value):
parsed = mp.parse_usage('openai', {'usage': {'prompt_tokens': value,
'completion_tokens': None, 'prompt_tokens_details': []}})
self.assertIsNone(parsed['input_tokens'])
self.assertIsNone(parsed['total_tokens'])
self.assertEqual(mp.parse_usage('openai', {'usage': {'prompt_tokens': '12', 'completion_tokens': '4'}})['total_tokens'], 16)
class AttemptMeteringTest(TestCase):
def call(self, handler, *, outlet=None, **kwargs):
database = mock.Mock()
temporary = tempfile.TemporaryDirectory()
self.addCleanup(temporary.cleanup)
database.path = Path(temporary.name) / "usage.db"
saved = []
database.record_model_usage_batch.side_effect = lambda events: saved.extend(dict(e) for e in events)
async def run():
async with UsageRecorder(database, tenant_id='tenant-test', task_id='task-test') as recorder:
async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
return await gw.call_outlet(client, outlet or fixtures._outlet(),
[{'role': 'user', 'content': 'private-prompt'}], usage_recorder=recorder, **kwargs)
result = asyncio.run(run())
return result, saved
def test_retry_counts_each_actual_request_even_failed_response_usage(self):
requests = []
def handler(request):
requests.append(request)
if len(requests) == 1:
return httpx.Response(503, json={'usage': {'prompt_tokens': 2, 'completion_tokens': 1}})
return httpx.Response(200, json={**fixtures._openai_body('ok'),
'usage': {'prompt_tokens': 10, 'completion_tokens': 3}})
result, events = self.call(handler)
self.assertEqual(result['text'], 'ok')
self.assertEqual([e['attempt'] for e in events], [1, 2])
self.assertEqual([e['status'] for e in events], ['error', 'success'])
self.assertEqual(sum(e['total_tokens'] for e in events), 16)
self.assertNotIn('private-prompt', json.dumps(events))
self.assertEqual(len({e['event_id'] for e in events}), 2)
def test_invalid_reply_retains_usage_even_when_answer_parser_fails(self):
result, events = self.call(lambda _: httpx.Response(200, json={'usage': {'prompt_tokens': 9, 'completion_tokens': 7}}))
self.assertTrue(result['error'])
self.assertEqual(events[0]['total_tokens'], 16)
self.assertEqual(events[0]['status'], 'error')
def test_unknown_usage_http_errors_and_timeouts_do_not_become_zero(self):
result, events = self.call(lambda _: httpx.Response(401, json={'error': 'no'}))
self.assertEqual(len(events), 1)
self.assertIsNone(events[0]['total_tokens'])
async def slow(_):
await asyncio.sleep(2)
return httpx.Response(200, json=fixtures._openai_body('late'))
result, events = self.call(slow, deadline=0.01)
self.assertTrue(result['error'])
self.assertEqual(len(events), 1)
self.assertEqual(events[0]['status'], 'error')
self.assertIsNone(events[0]['total_tokens'])
def test_breaker_or_bad_config_does_not_count_requests_never_sent(self):
outlet = fixtures._outlet()
for _ in range(gw.BREAKER_THRESHOLD):
outlet.breaker.record(False)
handler = mock.Mock(side_effect=AssertionError('must not call provider'))
_, events = self.call(handler, outlet=outlet)
self.assertEqual(events, [])
handler.assert_not_called()
def test_meter_write_failure_is_logged_and_does_not_change_answer(self):
database = mock.Mock()
temporary = tempfile.TemporaryDirectory()
self.addCleanup(temporary.cleanup)
database.path = Path(temporary.name) / "usage.db"
database.record_model_usage_batch.side_effect = sqlite3.OperationalError('test')
async def run():
async with UsageRecorder(database, tenant_id='test') as recorder:
recorder.capture(fixtures._outlet(), {}, attempt=1, role='answer', status='success', latency_ms=1)
return 'answer'
with self.assertLogs('model_usage', level='WARNING') as logs:
self.assertEqual(asyncio.run(run()), 'answer')
self.assertEqual(database.record_model_usage_batch.call_count, 2)
self.assertEqual(len(logs.output), 2)
class GatewayUsageTest(TestCase):
setUp = fixtures.GatewayEndpointTest.setUp
_post = fixtures.GatewayEndpointTest._post
def tearDown(self):
shutil.rmtree(self.root)
def usage(self):
return admin_backend.Database(self.db_path).list_model_usage(7, tenant_id=self.desktop_account['tenant_id'])
def upstream(self, request):
text = json.loads(request.content)['messages'][-1]['content']
reply = '{"winner":"B","score":0.9,"risk":"low"}' if '评审' in text else 'synthetic reply'
return httpx.Response(200, json={**fixtures._openai_body(reply),
'usage': {'prompt_tokens': 11, 'completion_tokens': 7, 'total_tokens': 18}})
def test_all_candidates_and_judge_count_but_idempotent_replay_does_not(self):
headers = {'X-Idempotency-Key': 'usage-e2e-1'}
self.assertEqual(self._post(self.upstream, headers=headers).status_code, 200)
events = self.usage()
self.assertEqual(events['total'], 3)
self.assertEqual({e['role'] for e in events['items']}, {'answer', 'judge'})
self.assertEqual({e['provider_id'] for e in events['items']}, {'a', 'b', 'j'})
self.assertEqual(sum(e['total_tokens'] for e in events['items']), 54)
self.assertEqual(self._post(self.upstream, headers=headers).status_code, 200)
self.assertEqual(self.usage()['total'], 3)
def test_all_failed_requests_are_metered_before_502(self):
response = self._post(lambda _: httpx.Response(400, json={'usage': {'total_tokens': 15}}))
self.assertEqual(response.status_code, 502)
self.assertEqual(self.usage()['total'], 2)
self.assertTrue(all(e['total_tokens'] == 15 for e in self.usage()['items']))
def test_guard_and_tool_round_do_not_add_judge_usage(self):
self.assertEqual(self._post(self.upstream, body={'customer_text': '界面检测', 'purpose': 'guard'}).status_code, 200)
guard = self.usage()['items']
self.assertEqual(len(guard), 1)
self.assertEqual((guard[0]['purpose'], guard[0]['role']), ('guard', 'answer'))
response = self._post(self.upstream, body={'customer_text': 'tool', 'tools': [
{'type': 'function', 'function': {'name': 'test', 'parameters': {'type': 'object'}}}]})
self.assertEqual(response.status_code, 200)
self.assertEqual(self.usage()['total'], 2)
class KnowledgeUsageTest(TestCase):
def test_extraction_counts_usage_before_draft_validation_failure(self):
import knowledge_worker
from knowledge_model_output import ModelOutputError
from test_knowledge_model_output import candidate
database = mock.Mock()
temporary = tempfile.TemporaryDirectory()
self.addCleanup(temporary.cleanup)
database.path = Path(temporary.name) / "usage.db"
saved = []
database.record_model_usage_batch.side_effect = lambda events: saved.extend(dict(e) for e in events)
client_type = httpx.AsyncClient
transport = httpx.MockTransport(lambda _: httpx.Response(200, json={
**fixtures._openai_body('{"answer":null}'), 'usage': {'prompt_tokens': 200, 'completion_tokens': 20}}))
with mock.patch('knowledge_models.resolve_model', return_value=fixtures._outlet()), \
mock.patch('httpx.AsyncClient', side_effect=lambda **kw: client_type(transport=transport, **kw)):
with self.assertRaises(ModelOutputError):
knowledge_worker.enhance(candidate(), database, {'model_timeout_seconds': 90},
tenant_id='knowledge-tenant', task_id='job-id')
self.assertEqual(len(saved), 1)
self.assertEqual((saved[0]['tenant_id'], saved[0]['task_id'], saved[0]['purpose']),
('knowledge-tenant', 'job-id', 'knowledge'))
self.assertEqual(saved[0]['total_tokens'], 220)