#!/usr/bin/env python3
"""aria-llm-trace — Wrapper für aria_llm_calls Supabase-Inserts (KAR-122).

Library + CLI für Insert + Daily-Aggregation-Query. Wird von Aria selbst
plus Cross-LLM-Eval-Adapter (KAR-121) plus Cost-Cap-Logic (KAR-123) genutzt.

API:
    from aria_llm_trace import log_call, todays_cost, ensure_cap

    log_call(
        task_class="plan", worker_model="claude-opus-4-7", worker_family="anthropic",
        judge_model="gpt-5", judge_family="openai", judge_verdict="pass",
        judge_confidence=0.92, prompt_tokens=2400, completion_tokens=800,
        cost_usd=0.024, latency_ms=3400, input_hash="...", output_hash="...",
        privacy_tier="cloud",
    )

CLI:
    python3 aria-llm-trace.py log --task-class plan --worker-model claude-opus-4-7 ...
    python3 aria-llm-trace.py summary           # today's totals per family
    python3 aria-llm-trace.py summary --days 7  # 7-day rolling
"""
from __future__ import annotations

import argparse
import hashlib
import json
import os
import sys
import time
import urllib.parse
import urllib.request
from pathlib import Path
from typing import Any, Optional

SUPABASE_URL = os.environ.get("SUPABASE_URL", "").rstrip("/")
SUPABASE_SERVICE_ROLE_KEY = os.environ.get("SUPABASE_SERVICE_ROLE_KEY", "")

if not SUPABASE_URL or not SUPABASE_SERVICE_ROLE_KEY:
    # Versuche env-File zu laden wenn Env leer ist (z.B. bei CLI-Aufruf außerhalb systemd)
    env_file = Path("/root/aria/.env")
    if env_file.exists():
        for line in env_file.read_text().splitlines():
            if line.startswith("SUPABASE_URL="):
                SUPABASE_URL = line.split("=", 1)[1].strip().strip('"').strip("'").rstrip("/")
            elif line.startswith("SUPABASE_SERVICE_ROLE_KEY="):
                SUPABASE_SERVICE_ROLE_KEY = line.split("=", 1)[1].strip().strip('"').strip("'")

REST_ENDPOINT = f"{SUPABASE_URL}/rest/v1/aria_llm_calls" if SUPABASE_URL else ""

# Cost-Cap-Default (KAR-123 — Sektion 12.3 Empfehlung $30/Monat)
DEFAULT_MONTHLY_CAP_USD = 50.0
WARN_THRESHOLD_PCT = 0.80


def _headers() -> dict:
    return {
        "apikey": SUPABASE_SERVICE_ROLE_KEY,
        "Authorization": f"Bearer {SUPABASE_SERVICE_ROLE_KEY}",
        "Content-Type": "application/json",
        "Prefer": "return=minimal",
    }


def log_call(**fields) -> bool:
    """Insert eine LLM-Call-Row. Schluckt Errors als Soft-Fail (Aria darf nicht crashen)."""
    if not REST_ENDPOINT:
        sys.stderr.write("[trace] SUPABASE_URL not set — skip log\n")
        return False
    fields.setdefault("ts", None)  # let DB-Default greifen
    if fields["ts"] is None:
        del fields["ts"]
    try:
        req = urllib.request.Request(REST_ENDPOINT, data=json.dumps(fields).encode(),
                                      headers=_headers(), method="POST")
        with urllib.request.urlopen(req, timeout=10) as r:
            return 200 <= r.status < 300
    except Exception as exc:  # noqa: BLE001
        sys.stderr.write(f"[trace] log_call failed: {exc}\n")
        return False


def query(filter_qs: str = "", select: str = "*", limit: int = 50) -> list[dict]:
    """Generic Supabase-REST-Query auf aria_llm_calls."""
    if not REST_ENDPOINT:
        return []
    url = f"{REST_ENDPOINT}?select={select}&limit={limit}"
    if filter_qs:
        url += f"&{filter_qs}"
    req = urllib.request.Request(url, headers=_headers(), method="GET")
    try:
        with urllib.request.urlopen(req, timeout=10) as r:
            return json.loads(r.read())
    except Exception as exc:  # noqa: BLE001
        sys.stderr.write(f"[trace] query failed: {exc}\n")
        return []


def todays_cost() -> float:
    """Sum cost_usd für heute (UTC). Nutzt aria_llm_calls_daily View."""
    if not SUPABASE_URL:
        return 0.0
    today = time.strftime("%Y-%m-%d", time.gmtime())
    rows = query(filter_qs=f"day=eq.{today}",
                  select="cost_usd_sum",
                  limit=200)
    # daily-View aggregiert per family — wir summieren clientseitig
    return sum(float(r.get("cost_usd_sum") or 0) for r in rows) if rows else 0.0


def month_to_date_cost() -> float:
    """Sum cost_usd für laufenden Monat."""
    if not REST_ENDPOINT:
        return 0.0
    month_start = time.strftime("%Y-%m-01T00:00:00", time.gmtime())
    rows = query(filter_qs=f"ts=gte.{month_start}",
                  select="cost_usd",
                  limit=10000)
    return sum(float(r.get("cost_usd") or 0) for r in rows) if rows else 0.0


def ensure_cap(cap_usd: float = DEFAULT_MONTHLY_CAP_USD) -> dict:
    """KAR-123 Cost-Cap-Check. Returns Status-Dict:
       {used_usd, cap_usd, used_pct, warn, downgrade}
    """
    used = month_to_date_cost()
    pct = used / cap_usd if cap_usd > 0 else 0
    return {
        "used_usd": round(used, 6),
        "cap_usd": cap_usd,
        "used_pct": round(pct, 3),
        "warn": pct >= WARN_THRESHOLD_PCT,
        "downgrade": pct >= 1.0,
    }


def make_hash(text: str) -> str:
    return hashlib.sha256(text.encode("utf-8")).hexdigest()[:16]


# ---- CLI --------------------------------------------------------------------

def _cli_log(args) -> int:
    payload = {
        "task_id": args.task_id,
        "task_class": args.task_class,
        "worker_model": args.worker_model,
        "worker_family": args.worker_family,
        "judge_model": args.judge_model,
        "judge_family": args.judge_family,
        "judge_verdict": args.judge_verdict,
        "judge_confidence": args.judge_confidence,
        "prompt_tokens": args.prompt_tokens,
        "completion_tokens": args.completion_tokens,
        "cost_usd": args.cost_usd,
        "latency_ms": args.latency_ms,
        "input_hash": args.input_hash,
        "output_hash": args.output_hash,
        "privacy_tier": args.privacy_tier,
    }
    payload = {k: v for k, v in payload.items() if v is not None}
    ok = log_call(**payload)
    print("logged" if ok else "FAIL")
    return 0 if ok else 1


def _cli_summary(args) -> int:
    cap = ensure_cap()
    print(json.dumps({
        "month_to_date": cap,
        "todays_cost_usd": round(todays_cost(), 6),
        "note": "KAR-122 trace + KAR-123 cap-check",
    }, indent=2))
    return 0


def _cli_cap(args) -> int:
    res = ensure_cap(args.cap)
    print(json.dumps(res, indent=2))
    return 0


def main() -> int:
    parser = argparse.ArgumentParser(description="Aria LLM-Call Trace (KAR-122)")
    sub = parser.add_subparsers(dest="cmd", required=True)

    p_log = sub.add_parser("log", help="insert a llm-call trace row")
    p_log.add_argument("--task-id")
    p_log.add_argument("--task-class")
    p_log.add_argument("--worker-model", required=True)
    p_log.add_argument("--worker-family")
    p_log.add_argument("--judge-model")
    p_log.add_argument("--judge-family")
    p_log.add_argument("--judge-verdict", choices=["pass", "fail", None])
    p_log.add_argument("--judge-confidence", type=float)
    p_log.add_argument("--prompt-tokens", type=int)
    p_log.add_argument("--completion-tokens", type=int)
    p_log.add_argument("--cost-usd", type=float)
    p_log.add_argument("--latency-ms", type=int)
    p_log.add_argument("--input-hash")
    p_log.add_argument("--output-hash")
    p_log.add_argument("--privacy-tier", default="unknown")
    p_log.set_defaults(func=_cli_log)

    p_sum = sub.add_parser("summary", help="cost summary + cap-check (KAR-123)")
    p_sum.set_defaults(func=_cli_summary)

    p_cap = sub.add_parser("cap", help="cost-cap-status only")
    p_cap.add_argument("--cap", type=float, default=DEFAULT_MONTHLY_CAP_USD)
    p_cap.set_defaults(func=_cli_cap)

    args = parser.parse_args()
    return args.func(args)


if __name__ == "__main__":
    sys.exit(main())
