"""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.