209 lines
11 KiB
Python
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)
|