"""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()