306 lines
16 KiB
Python
306 lines
16 KiB
Python
"""Usage accounting from actual provider results, with tenant/RBAC isolation."""
|
|
from __future__ import annotations
|
|
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
from datetime import datetime, timedelta, timezone
|
|
import json
|
|
from pathlib import Path
|
|
import shutil
|
|
import sqlite3
|
|
import tempfile
|
|
import time
|
|
import unittest
|
|
from unittest import mock
|
|
import uuid
|
|
|
|
import admin_backend as backend
|
|
from test_admin_api import _Base
|
|
|
|
|
|
def event(**changes):
|
|
data = {
|
|
'event_id': uuid.uuid4().hex, 'request_id': 'request-1', 'task_id': 'task-1',
|
|
'tenant_id': 'default', 'desktop_account_id': None, 'provider_id': 'provider-a',
|
|
'provider_name': '模型服务A', 'model': 'actual-model', 'kind': 'openai',
|
|
'purpose': 'chat', 'role': 'answer', 'attempt': 1, 'status': 'success',
|
|
'input_tokens': 10, 'output_tokens': 20, 'total_tokens': 30,
|
|
'cached_input_tokens': 2, 'reasoning_tokens': 5, 'latency_ms': 123,
|
|
}
|
|
data.update(changes)
|
|
return data
|
|
|
|
|
|
class ModelUsageLedgerTests(unittest.TestCase):
|
|
def setUp(self):
|
|
self.temp = tempfile.TemporaryDirectory(prefix='model-usage-ledger-test-')
|
|
self.addCleanup(self.temp.cleanup)
|
|
self.db = backend.Database(Path(self.temp.name) / 'usage.db')
|
|
self.db.migrate()
|
|
|
|
def record(self, **changes):
|
|
data = event(**changes)
|
|
self.assertTrue(self.db.record_model_usage(data))
|
|
return data
|
|
|
|
def summary(self, **kwargs):
|
|
return self.db.model_usage_stats(**kwargs)['summary']
|
|
|
|
def legacy(self, task_id='', tenant='default', purpose='chat', created=None):
|
|
self.db.log_model_call({'task_id': task_id, 'tenant_id': tenant, 'purpose': purpose})
|
|
if created:
|
|
with self.db.connect() as db:
|
|
db.execute('UPDATE model_calls SET created_at=? WHERE id=(SELECT MAX(id) FROM model_calls)', (created,))
|
|
|
|
def test_answer_candidates_judge_retry_and_error_are_all_charged(self):
|
|
self.record()
|
|
self.record(provider_id='provider-b', status='error', input_tokens=3, output_tokens=0,
|
|
total_tokens=3, cached_input_tokens=0, reasoning_tokens=0)
|
|
self.record(provider_id='provider-b', attempt=2, input_tokens=4, output_tokens=6,
|
|
total_tokens=10, cached_input_tokens=None, reasoning_tokens=None)
|
|
self.record(role='judge', input_tokens=7, output_tokens=2, total_tokens=9,
|
|
cached_input_tokens=1, reasoning_tokens=1)
|
|
summary = self.summary()
|
|
self.assertEqual((summary['request_count'], summary['known_count'], summary['missing_count']), (4, 4, 0))
|
|
self.assertEqual((summary['input_tokens'], summary['output_tokens'], summary['total_tokens']), (24, 28, 52))
|
|
self.assertEqual((summary['cached_input_tokens'], summary['reasoning_tokens']), (3, 6))
|
|
self.assertEqual(summary['cached_input_known_count'], 3)
|
|
|
|
def test_missing_usage_stays_null_and_is_not_counted_as_known(self):
|
|
self.record(input_tokens=None, output_tokens=None, total_tokens=None,
|
|
cached_input_tokens=None, reasoning_tokens=None, status='error')
|
|
summary = self.summary()
|
|
self.assertEqual((summary['request_count'], summary['known_count'], summary['missing_count']), (1, 0, 1))
|
|
self.assertIsNone(summary['total_tokens'])
|
|
self.assertIsNone(summary['input_tokens'])
|
|
self.assertIsNone(self.db.list_model_usage()['items'][0]['output_tokens'])
|
|
|
|
def test_partial_usage_is_never_filled_or_total_estimated(self):
|
|
self.record(input_tokens=6, output_tokens=2, total_tokens=None,
|
|
cached_input_tokens=None, reasoning_tokens=None)
|
|
summary = self.summary()
|
|
self.assertEqual(summary['input_tokens'], 6)
|
|
self.assertEqual(summary['input_known_count'], 1)
|
|
self.assertEqual(summary['known_count'], 0)
|
|
self.assertIsNone(summary['total_tokens'])
|
|
|
|
def test_zero_reported_tokens_are_known(self):
|
|
self.record(input_tokens=0, output_tokens=0, total_tokens=0,
|
|
cached_input_tokens=0, reasoning_tokens=0)
|
|
summary = self.summary()
|
|
self.assertEqual(summary['known_count'], 1)
|
|
self.assertEqual(summary['missing_count'], 0)
|
|
self.assertEqual(summary['total_tokens'], 0)
|
|
|
|
def test_empty_statistics_are_zero_with_all_requested_calendar_days(self):
|
|
result = self.db.model_usage_stats(days=3)
|
|
self.assertEqual(result['summary']['total_tokens'], 0)
|
|
self.assertEqual(result['summary']['request_count'], 0)
|
|
self.assertEqual(result['by_model'], [])
|
|
self.assertEqual(len(result['daily']), 3)
|
|
self.assertEqual(result['timezone'], 'UTC')
|
|
|
|
def test_event_id_is_idempotent_and_first_record_is_immutable(self):
|
|
data = self.record(event_id='one')
|
|
self.assertFalse(self.db.record_model_usage({**data, 'total_tokens': 999}))
|
|
self.assertEqual(self.summary()['total_tokens'], 30)
|
|
|
|
def test_batch_retry_only_inserts_new_events(self):
|
|
records = [event(event_id='a'), event(event_id='b', role='judge')]
|
|
self.assertEqual(self.db.record_model_usage_batch(records), 2)
|
|
self.assertEqual(self.db.record_model_usage_batch([*records, event(event_id='c')]), 1)
|
|
self.assertEqual(self.db.record_model_usage_batch([]), 0)
|
|
|
|
def test_concurrent_duplicate_writers_insert_once(self):
|
|
data = event(event_id='same-event')
|
|
with ThreadPoolExecutor(max_workers=6) as executor:
|
|
results = list(executor.map(lambda _: self.db.record_model_usage(data), range(12)))
|
|
self.assertEqual(sum(results), 1)
|
|
self.assertEqual(self.summary()['request_count'], 1)
|
|
|
|
def test_validation_or_storage_failure_raises_and_batch_does_not_partially_write(self):
|
|
with self.assertRaises(ValueError):
|
|
self.db.record_model_usage_batch([event(), event(total_tokens=-1)])
|
|
self.assertEqual(self.summary()['request_count'], 0)
|
|
with mock.patch.object(self.db, 'connect', side_effect=sqlite3.OperationalError('fixture locked')):
|
|
with self.assertRaises(sqlite3.OperationalError):
|
|
self.db.record_model_usage(event())
|
|
|
|
def test_usage_write_lock_has_short_budget_without_changing_other_connections(self):
|
|
record = event()
|
|
with self.db.connect() as locker:
|
|
self.assertEqual(locker.execute('PRAGMA busy_timeout').fetchone()[0], 10000)
|
|
locker.execute('BEGIN IMMEDIATE')
|
|
started = time.monotonic()
|
|
with self.assertRaises(sqlite3.OperationalError) as raised:
|
|
self.db.record_model_usage_batch([record])
|
|
elapsed = time.monotonic() - started
|
|
self.assertIn('locked', str(raised.exception))
|
|
self.assertGreaterEqual(elapsed, .15)
|
|
self.assertLess(elapsed, 1.5, 'metering must not inherit the normal 10s lock wait')
|
|
self.assertTrue(self.db.record_model_usage(record))
|
|
with self.db.connect() as other:
|
|
self.assertEqual(other.execute('PRAGMA busy_timeout').fetchone()[0], 10000)
|
|
|
|
def test_invalid_metadata_and_noninteger_tokens_are_rejected(self):
|
|
for changes in ({'event_id': ''}, {'attempt': 0}, {'role': 'system'}, {'purpose': 'secret'},
|
|
{'status': 'maybe'}, {'total_tokens': True}, {'input_tokens': 1.5},
|
|
{'output_tokens': '3'}, {'total_tokens': 2**63},
|
|
{'created_at': '2026-09-17T12:00:00'}):
|
|
with self.subTest(changes=changes), self.assertRaises(ValueError):
|
|
self.db.record_model_usage(event(**changes))
|
|
|
|
def test_only_allowlisted_metadata_is_persisted(self):
|
|
self.record(prompt='private prompt', api_key='private key', headers={'Authorization': 'private bearer'},
|
|
response={'text': 'private reply'})
|
|
item = self.db.list_model_usage()['items'][0]
|
|
serialized = json.dumps(item)
|
|
self.assertNotIn('private', serialized)
|
|
self.assertNotIn('prompt', item)
|
|
with self.db.connect() as db:
|
|
fields = {row[1] for row in db.execute('PRAGMA table_info(model_usage_events)')}
|
|
self.assertEqual(fields, set(item))
|
|
|
|
def test_default_specific_and_multi_account_scopes_are_isolated(self):
|
|
for tenant in ('', 'default', 'tenant-a', 'tenant-b', 'tenant-secret'):
|
|
self.record(tenant_id=tenant)
|
|
self.assertEqual(self.summary()['request_count'], 2)
|
|
self.assertEqual(self.summary(tenant_id='tenant-a')['request_count'], 1)
|
|
self.assertEqual(self.summary(tenant_id=['tenant-a', 'tenant-b'])['request_count'], 2)
|
|
self.assertEqual(self.summary(tenant_id='*')['request_count'], 5)
|
|
self.assertEqual({item['tenant_id'] for item in self.db.list_model_usage(tenant_id=['tenant-a', 'tenant-b'])['items']}, {'tenant-a', 'tenant-b'})
|
|
|
|
def test_all_purposes_default_and_precise_provider_model_filters(self):
|
|
self.record()
|
|
self.record(purpose='guard', model='vision-model')
|
|
self.record(purpose='knowledge', provider_id='provider-b')
|
|
result = self.db.model_usage_stats()
|
|
self.assertEqual(result['summary']['request_count'], 3)
|
|
self.assertEqual({row['purpose'] for row in result['by_purpose']}, {'chat', 'guard', 'knowledge'})
|
|
self.assertEqual(len(result['by_model']), 3)
|
|
self.assertEqual(self.summary(purpose='guard', model='vision-model')['request_count'], 1)
|
|
self.assertEqual(self.summary(provider_id='provider-b')['request_count'], 1)
|
|
self.assertEqual(self.summary(model="' OR 1=1 --")['request_count'], 0)
|
|
|
|
def test_utc_window_includes_today_excludes_tomorrow_and_normalizes_offsets(self):
|
|
class FixedDatetime(datetime):
|
|
@classmethod
|
|
def now(cls, tz=None):
|
|
value = cls(2026, 9, 17, 12, tzinfo=timezone.utc)
|
|
return value if tz is None else value.astimezone(tz)
|
|
self.record(event_id='early', created_at='2026-09-17T00:00:00+08:00')
|
|
self.record(event_id='today', created_at='2026-09-17T00:00:00Z')
|
|
self.record(event_id='tomorrow', created_at='2026-09-18T00:00:00+00:00')
|
|
with mock.patch.object(backend, 'datetime', FixedDatetime):
|
|
result = self.db.model_usage_stats(days=1)
|
|
self.assertEqual(result['since'], '2026-09-17T00:00:00.000+00:00')
|
|
self.assertEqual(result['summary']['request_count'], 1)
|
|
self.assertEqual(result['daily'][0]['date'], '2026-09-17')
|
|
self.assertEqual([item['event_id'] for item in self.db.list_model_usage(days=1)['items']], ['today'])
|
|
self.assertEqual(self.db.model_usage_stats(days=2)['summary']['request_count'], 2)
|
|
|
|
def test_pagination_is_stable_and_database_limit_is_capped(self):
|
|
stamp = datetime.now(timezone.utc).isoformat()
|
|
self.db.record_model_usage_batch([event(event_id=f'event-{n:03}', created_at=stamp) for n in range(205)])
|
|
result = self.db.list_model_usage(limit=999)
|
|
self.assertEqual(result['total'], 205)
|
|
self.assertEqual(len(result['items']), 200)
|
|
page = self.db.list_model_usage(limit=3, offset=2)
|
|
self.assertEqual([item['event_id'] for item in page['items']], ['event-202', 'event-201', 'event-200'])
|
|
self.assertEqual(self.db.model_usage_stats(days=900)['days'], 90)
|
|
|
|
def test_legacy_logs_are_not_fabricated_and_only_unmatched_tasks_are_counted(self):
|
|
self.legacy('has-usage')
|
|
self.legacy('old-task')
|
|
self.legacy('')
|
|
self.legacy('old-guard', purpose='guard')
|
|
self.record(task_id='has-usage')
|
|
result = self.db.model_usage_stats()
|
|
self.assertEqual(result['summary']['request_count'], 1)
|
|
self.assertEqual(result['legacy_unmetered_calls'], 3)
|
|
self.assertEqual(self.db.model_usage_stats(purpose='chat')['legacy_unmetered_calls'], 2)
|
|
self.assertEqual(self.db.model_usage_stats(provider_id='other')['legacy_unmetered_calls'], 3)
|
|
self.assertEqual(result['legacy_unmetered_scope'], 'account_time_purpose')
|
|
|
|
def test_legacy_task_match_does_not_cross_tenants_and_default_alias_matches(self):
|
|
self.legacy('same-task', tenant='tenant-a')
|
|
self.record(task_id='same-task', tenant_id='tenant-b')
|
|
self.assertEqual(self.db.model_usage_stats(tenant_id='tenant-a')['legacy_unmetered_calls'], 1)
|
|
self.legacy('default-task', tenant='')
|
|
self.record(task_id='default-task', tenant_id='default')
|
|
self.assertEqual(self.db.model_usage_stats()['legacy_unmetered_calls'], 0)
|
|
|
|
def test_migration_preserves_old_rows_and_is_repeatable(self):
|
|
self.legacy('old-task')
|
|
with self.db.connect() as db:
|
|
db.execute('DROP TABLE model_usage_events')
|
|
self.db.migrate()
|
|
self.db.migrate()
|
|
self.assertEqual(self.summary()['request_count'], 0)
|
|
self.assertEqual(self.db.model_usage_stats()['legacy_unmetered_calls'], 1)
|
|
self.record()
|
|
with self.db.connect() as db:
|
|
indexes = {row[1] for row in db.execute('PRAGMA index_list(model_usage_events)')}
|
|
self.assertIn('idx_model_usage_tenant_created', indexes)
|
|
self.assertIn('idx_model_usage_task_tenant', indexes)
|
|
|
|
|
|
class ModelUsageApiTests(_Base):
|
|
def tearDown(self):
|
|
self.client.close()
|
|
shutil.rmtree(self.root)
|
|
|
|
def test_stats_and_log_require_login_and_stats_read(self):
|
|
self.db.save_role('usage-no-stats', '无统计', ['model:read'], 1, '127.0.0.1')
|
|
denied = self.make_user('usage-denied', 'usage-no-stats')
|
|
for path in ('/api/v2/stats/model-usage', '/api/v2/stats/model-usage/log'):
|
|
self.assertEqual(self.client.get(path).status_code, 401)
|
|
self.assertEqual(self.client.get(path, headers=denied).status_code, 403)
|
|
|
|
def test_admin_can_query_all_usage_and_filter_model(self):
|
|
self.db.record_model_usage(event(tenant_id='a'))
|
|
self.db.record_model_usage(event(tenant_id='b', purpose='knowledge', model='other'))
|
|
headers = self.login()
|
|
response = self.client.get('/api/v2/stats/model-usage', headers=headers, params={'account_id': -1})
|
|
self.assertEqual(response.status_code, 200, response.text)
|
|
self.assertEqual(response.json()['summary']['request_count'], 2)
|
|
rows = self.client.get('/api/v2/stats/model-usage/log', headers=headers,
|
|
params={'account_id': -1, 'purpose': 'knowledge', 'model': 'other'}).json()
|
|
self.assertEqual(rows['total'], 1)
|
|
self.assertEqual(rows['items'][0]['tenant_id'], 'b')
|
|
|
|
def test_scoped_admin_all_means_only_authorized_accounts(self):
|
|
first = self.desktop_login('zyt-user-1', 'usage-device-one')
|
|
second = self.desktop_login('zyt-user-2', 'usage-device-two')
|
|
a = self.client.get('/api/v2/desktop/me', headers=first).json()
|
|
b = self.client.get('/api/v2/desktop/me', headers=second).json()
|
|
for account in (a, b):
|
|
self.db.record_model_usage(event(tenant_id=account['tenant_id'], desktop_account_id=account['id']))
|
|
self.db.save_role('usage-scoped', '统计人员', ['stats:read'], 1, '127.0.0.1')
|
|
viewer = self.make_user('usage-viewer', 'usage-scoped')
|
|
uid = self.client.get('/api/v2/me', headers=viewer).json()['id']
|
|
self.db.set_desktop_account_admins(a['id'], [uid], 1, '127.0.0.1')
|
|
for path in ('/api/v2/stats/model-usage', '/api/v2/stats/model-usage/log'):
|
|
denied = self.client.get(path, headers=viewer, params={'account_id': b['id']})
|
|
self.assertEqual(denied.status_code, 403)
|
|
self.assertEqual(self.client.get(path, headers=viewer).status_code, 403)
|
|
allowed = self.client.get(path, headers=viewer, params={'account_id': -1})
|
|
self.assertEqual(allowed.status_code, 200, allowed.text)
|
|
data = allowed.json()
|
|
self.assertEqual(data.get('total', data.get('summary', {}).get('request_count')), 1)
|
|
|
|
def test_invalid_query_ranges_are_rejected(self):
|
|
headers = self.login()
|
|
for path, params in (
|
|
('/api/v2/stats/model-usage', {'days': 0}),
|
|
('/api/v2/stats/model-usage', {'days': 91}),
|
|
('/api/v2/stats/model-usage', {'purpose': 'not-a-purpose'}),
|
|
('/api/v2/stats/model-usage/log', {'limit': 201}),
|
|
('/api/v2/stats/model-usage/log', {'offset': -1}),
|
|
):
|
|
with self.subTest(params=params):
|
|
self.assertEqual(self.client.get(path, headers=headers, params=params).status_code, 422)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|