131 lines
5.6 KiB
Python
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.
|