Files
2026-06-27 22:22:42 +08:00

252 lines
11 KiB
Python

import json
import os
import threading
from typing import Any
try:
import psycopg
except Exception: # pragma: no cover
psycopg = None
class BaseAuditStore:
def record_api_request(self, payload: dict[str, Any]) -> None:
return None
def record_llm_call(self, payload: dict[str, Any]) -> None:
return None
def upsert_daily_usage(self, payload: dict[str, Any]) -> None:
return None
class NullAuditStore(BaseAuditStore):
pass
class PostgresAuditStore(BaseAuditStore):
def __init__(self, database_url: str) -> None:
if psycopg is None:
raise RuntimeError("psycopg 未安装,无法使用 PostgreSQL 审计存储")
self.database_url = database_url
self._init_lock = threading.Lock()
self._initialized = False
def _connect(self):
return psycopg.connect(self.database_url, autocommit=True)
def _ensure_initialized(self) -> None:
if self._initialized:
return
with self._init_lock:
if self._initialized:
return
with self._connect() as conn:
with conn.cursor() as cur:
cur.execute(
"""
CREATE TABLE IF NOT EXISTS api_request_audit (
id BIGSERIAL PRIMARY KEY,
request_id TEXT NOT NULL,
session_hash TEXT NOT NULL,
ip_hash TEXT NOT NULL,
route TEXT NOT NULL,
method TEXT NOT NULL,
status_code INTEGER NOT NULL,
decision TEXT NOT NULL,
delay_ms INTEGER NOT NULL DEFAULT 0,
queue_ms INTEGER NOT NULL DEFAULT 0,
error_code TEXT NOT NULL DEFAULT '',
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
metadata_json JSONB NOT NULL DEFAULT '{}'::jsonb
)
"""
)
cur.execute(
"""
CREATE TABLE IF NOT EXISTS llm_call_audit (
id BIGSERIAL PRIMARY KEY,
request_id TEXT NOT NULL,
session_hash TEXT NOT NULL,
ip_hash TEXT NOT NULL,
job_type TEXT NOT NULL,
model TEXT NOT NULL,
estimated_input_tokens INTEGER NOT NULL DEFAULT 0,
max_output_tokens INTEGER NOT NULL DEFAULT 0,
estimated_cost NUMERIC(18, 8) NOT NULL DEFAULT 0,
actual_output_chars INTEGER NOT NULL DEFAULT 0,
actual_cost NUMERIC(18, 8) NOT NULL DEFAULT 0,
queue_ms INTEGER NOT NULL DEFAULT 0,
run_ms INTEGER NOT NULL DEFAULT 0,
total_ms INTEGER NOT NULL DEFAULT 0,
status TEXT NOT NULL,
error_code TEXT NOT NULL DEFAULT '',
started_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
finished_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
metadata_json JSONB NOT NULL DEFAULT '{}'::jsonb
)
"""
)
cur.execute(
"ALTER TABLE llm_call_audit ADD COLUMN IF NOT EXISTS queue_ms INTEGER NOT NULL DEFAULT 0"
)
cur.execute(
"ALTER TABLE llm_call_audit ADD COLUMN IF NOT EXISTS run_ms INTEGER NOT NULL DEFAULT 0"
)
cur.execute(
"ALTER TABLE llm_call_audit ADD COLUMN IF NOT EXISTS total_ms INTEGER NOT NULL DEFAULT 0"
)
cur.execute(
"""
CREATE TABLE IF NOT EXISTS risk_events (
id BIGSERIAL PRIMARY KEY,
session_hash TEXT NOT NULL,
ip_hash TEXT NOT NULL,
event_type TEXT NOT NULL,
severity TEXT NOT NULL,
reason TEXT NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
metadata_json JSONB NOT NULL DEFAULT '{}'::jsonb
)
"""
)
cur.execute(
"""
CREATE TABLE IF NOT EXISTS daily_budget_usage (
usage_day DATE NOT NULL,
scope TEXT NOT NULL,
scope_hash TEXT NOT NULL,
estimated_cost NUMERIC(18, 8) NOT NULL DEFAULT 0,
actual_cost NUMERIC(18, 8) NOT NULL DEFAULT 0,
request_count INTEGER NOT NULL DEFAULT 0,
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
PRIMARY KEY (usage_day, scope, scope_hash)
)
"""
)
cur.execute(
"CREATE INDEX IF NOT EXISTS idx_api_request_audit_request_id ON api_request_audit (request_id)"
)
cur.execute(
"CREATE INDEX IF NOT EXISTS idx_api_request_audit_route_created_at ON api_request_audit (route, created_at DESC)"
)
cur.execute(
"CREATE INDEX IF NOT EXISTS idx_llm_call_audit_request_id ON llm_call_audit (request_id)"
)
cur.execute(
"CREATE INDEX IF NOT EXISTS idx_llm_call_audit_job_type_started_at ON llm_call_audit (job_type, started_at DESC)"
)
cur.execute(
"CREATE INDEX IF NOT EXISTS idx_llm_call_audit_model_started_at ON llm_call_audit (model, started_at DESC)"
)
self._initialized = True
def record_api_request(self, payload: dict[str, Any]) -> None:
self._ensure_initialized()
metadata = payload.get("metadata") or {}
with self._connect() as conn:
with conn.cursor() as cur:
cur.execute(
"""
INSERT INTO api_request_audit (
request_id, session_hash, ip_hash, route, method, status_code,
decision, delay_ms, queue_ms, error_code, metadata_json
)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s::jsonb)
""",
(
payload["request_id"],
payload["session_hash"],
payload["ip_hash"],
payload["route"],
payload["method"],
int(payload["status_code"]),
payload["decision"],
int(payload.get("delay_ms", 0)),
int(payload.get("queue_ms", 0)),
payload.get("error_code", ""),
json.dumps(metadata, ensure_ascii=False),
),
)
def record_llm_call(self, payload: dict[str, Any]) -> None:
self._ensure_initialized()
metadata = payload.get("metadata") or {}
with self._connect() as conn:
with conn.cursor() as cur:
cur.execute(
"""
INSERT INTO llm_call_audit (
request_id, session_hash, ip_hash, job_type, model,
estimated_input_tokens, max_output_tokens, estimated_cost,
actual_output_chars, actual_cost, queue_ms, run_ms, total_ms,
status, error_code, metadata_json
)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s::jsonb)
""",
(
payload["request_id"],
payload["session_hash"],
payload["ip_hash"],
payload["job_type"],
payload["model"],
int(payload.get("estimated_input_tokens", 0)),
int(payload.get("max_output_tokens", 0)),
float(payload.get("estimated_cost", 0.0)),
int(payload.get("actual_output_chars", 0)),
float(payload.get("actual_cost", 0.0)),
int(payload.get("queue_ms", 0)),
int(payload.get("run_ms", 0)),
int(payload.get("total_ms", 0)),
payload["status"],
payload.get("error_code", ""),
json.dumps(metadata, ensure_ascii=False),
),
)
def upsert_daily_usage(self, payload: dict[str, Any]) -> None:
self._ensure_initialized()
with self._connect() as conn:
with conn.cursor() as cur:
cur.execute(
"""
INSERT INTO daily_budget_usage (
usage_day, scope, scope_hash, estimated_cost, actual_cost, request_count
)
VALUES (%s, %s, %s, %s, %s, %s)
ON CONFLICT (usage_day, scope, scope_hash)
DO UPDATE SET
estimated_cost = daily_budget_usage.estimated_cost + EXCLUDED.estimated_cost,
actual_cost = daily_budget_usage.actual_cost + EXCLUDED.actual_cost,
request_count = daily_budget_usage.request_count + EXCLUDED.request_count,
updated_at = CURRENT_TIMESTAMP
""",
(
payload["usage_day"],
payload["scope"],
payload["scope_hash"],
float(payload.get("estimated_cost", 0.0)),
float(payload.get("actual_cost", 0.0)),
int(payload.get("request_count", 0)),
),
)
_audit_store: BaseAuditStore | None = None
def get_audit_store(database_url: str | None = None) -> BaseAuditStore:
global _audit_store
if _audit_store is not None:
return _audit_store
if database_url or os.getenv("DATABASE_URL"):
_audit_store = PostgresAuditStore(database_url or os.getenv("DATABASE_URL", ""))
else:
_audit_store = NullAuditStore()
return _audit_store
def reset_audit_store() -> None:
global _audit_store
_audit_store = None