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

131 lines
5.6 KiB
Python

"""Per-upstream-attempt metering, separate from deduplicated chat/review logs."""
from __future__ import annotations
import asyncio
import json
import os
import logging
from datetime import datetime, timezone
from itertools import islice
from pathlib import Path
from uuid import uuid4
import model_protocol
logger = logging.getLogger(__name__)
class UsageRecorder:
"""Accumulate bounded request metadata; never persist prompts or provider bodies.
Flush outside model deadlines, on both successful and failed gateway requests.
A recorder belongs to one authenticated request (or one knowledge extraction),
so concurrent candidate calls share attribution without any global context.
"""
def __init__(self, database, *, tenant_id, task_id="", desktop_account_id=None,
purpose="chat"):
self.database = database
self.request_id = uuid4().hex
self.context = {
"request_id": self.request_id,
"task_id": str(task_id or f"usage-{self.request_id}"),
"tenant_id": str(tenant_id),
"desktop_account_id": desktop_account_id,
"purpose": purpose if purpose in {"chat", "guard", "knowledge"} else "chat",
}
self.events = []
def capture(self, outlet, data, *, attempt, role, status, latency_ms):
self.events.append({
**self.context,
"event_id": uuid4().hex,
"provider_id": outlet.id,
"provider_name": outlet.name,
# Keep the configured model name: provider echoes may contain arbitrary data.
"model": str(outlet.config.get("model") or ""),
"kind": outlet.kind,
"role": role,
"attempt": attempt,
"status": status,
"latency_ms": latency_ms,
"created_at": datetime.now(timezone.utc).isoformat(timespec="milliseconds"),
**model_protocol.parse_usage(outlet.kind, data, outlet.config.get("base_url") or ""),
})
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc, traceback):
if not self.events:
return
# Shield the DB write from client disconnects. Retain/await the task so
# asyncio.run (knowledge worker) cannot discard it while closing its loop.
write = asyncio.create_task(asyncio.to_thread(self._write))
try:
await asyncio.shield(write)
except asyncio.CancelledError:
await write
raise
def _write(self):
# Event IDs make retry safe even if the original commit outcome is unknown.
for attempt in range(2):
try:
self.database.record_model_usage_batch(self.events)
self.events.clear()
try:
recover_pending_usage(self.database, limit=4)
except Exception as exc:
logger.warning("Token usage recovery unavailable: error=%s", type(exc).__name__)
return
except Exception as exc:
logger.warning("Token usage persistence failed: request=%s events=%d attempt=%d error=%s",
self.request_id, len(self.events), attempt + 1, type(exc).__name__)
# A locked accounting table must not hold up a customer reply. Persist
# metadata atomically, then replay the same event IDs after recovery.
try:
folder = Path(self.database.path).parent / 'model_usage_pending'
folder.mkdir(mode=0o700, exist_ok=True)
target = folder / (self.request_id + '.json')
temporary = folder / (self.request_id + '.tmp')
descriptor = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
with os.fdopen(descriptor, 'w', encoding='utf-8') as output:
json.dump({'version': 1, 'events': self.events}, output, ensure_ascii=False)
output.flush()
os.fsync(output.fileno())
os.replace(temporary, target)
self.events.clear()
except Exception as exc:
logger.error('Token usage spool failed: request=%s error=%s', self.request_id, type(exc).__name__)
def recover_pending_usage(database, *, limit=20):
"""Replay bounded metadata-only files; concurrent recovery is idempotent."""
folder = Path(database.path).parent / 'model_usage_pending'
if not folder.is_dir():
return
for path in islice(folder.glob('*.json'), limit):
try:
if path.stat().st_size > 1_000_000:
raise ValueError('oversized usage spool')
payload = json.loads(path.read_text(encoding='utf-8'))
if not isinstance(payload, dict) or payload.get('version') != 1 or not isinstance(payload.get('events'), list):
raise ValueError('invalid usage spool')
database.record_model_usage_batch(payload['events'])
path.unlink(missing_ok=True)
except FileNotFoundError:
continue
except ValueError as exc:
# Preserve malformed metadata for inspection without blocking every
# later event forever. Quarantine never deletes the original bytes.
try:
os.replace(path, path.with_suffix('.invalid-' + uuid4().hex))
except FileNotFoundError:
continue
logger.warning('Invalid token usage spool quarantined: error=%s', type(exc).__name__)
except Exception as exc:
logger.warning('Token usage replay pending: error=%s', type(exc).__name__)
break # At most one short DB lock timeout per sweep.