#!/usr/bin/env python3
"""
sebbi_spend_guard.py  -  a hard budget for any AI agent
=======================================================

One file, standard library only. Put a ceiling on what an agent can spend,
and stop it dead when the ceiling is reached.

    from sebbi_spend_guard import Guard, BudgetSpent

    guard = Guard(daily_limit_gbp=5.00, price_per_million=3.00)

    @guard.meter
    def ask_model(prompt):
        ...                      # your existing call, unchanged

    try:
        ask_model("summarise this")
    except BudgetSpent as e:
        print("stopped:", e)     # the call never happened

What it does
  - Estimates the cost of every call BEFORE it is made and refuses the call
    that would break the budget, rather than reporting it afterwards.
  - Catches runaway loops: the same prompt repeated is served from a local
    cache, and a burst of calls in a few seconds trips a circuit breaker.
  - Keeps a local, hash-chained record of every decision in spend_guard.log,
    so the spend record cannot be quietly edited after an incident.
  - Sends nothing anywhere. Prompts never leave the machine; only you read
    the log.

Run it directly for a demonstration:  python3 sebbi_spend_guard.py
"""

import functools
import hashlib
import json
import os
import threading
import time

__version__ = "1.0.0"
LOG = os.environ.get("SEBBI_GUARD_LOG", "spend_guard.log")


class BudgetSpent(Exception):
    """The call was refused. It did not happen."""


class Guard:
    def __init__(self, daily_limit_gbp=5.0, price_per_million=3.00,
                 burst=25, burst_seconds=10, cache=True, log=LOG):
        self.limit = float(daily_limit_gbp)
        self.price = float(price_per_million)
        self.burst, self.window = int(burst), float(burst_seconds)
        self.cache_on, self.log_path = bool(cache), log
        self._lock = threading.RLock()
        self._spent, self._day = 0.0, time.strftime("%Y-%m-%d")
        self._recent, self._cache, self._tip = [], {}, "GENESIS"

    # ---------------------------------------------------------------- record
    def _seal(self, event):
        line = {"ts": round(time.time(), 3), "prev": self._tip}
        line.update(event)
        body = json.dumps(line, sort_keys=True, separators=(",", ":"))
        line["hash"] = self._tip = hashlib.sha256(body.encode()).hexdigest()
        try:
            with open(self.log_path, "a", encoding="utf-8") as f:
                f.write(json.dumps(line, sort_keys=True, separators=(",", ":")) + "\n")
        except Exception:
            pass
        return line["hash"]

    @staticmethod
    def tokens(text):
        """Estimate: about four characters per token in English."""
        return max(1, round(len(text or "") / 4))

    def cost(self, text):
        return self.tokens(text) * self.price / 1_000_000

    def spent_today(self):
        with self._lock:
            self._roll()
            return round(self._spent, 6)

    def remaining(self):
        return round(max(0.0, self.limit - self.spent_today()), 6)

    def _roll(self):
        today = time.strftime("%Y-%m-%d")
        if today != self._day:
            self._day, self._spent, self._recent = today, 0.0, []

    # ---------------------------------------------------------------- gate
    def check(self, prompt):
        """Decide before the call. Returns a cached answer, or None to proceed."""
        with self._lock:
            self._roll()
            now = time.time()
            self._recent = [t for t in self._recent if now - t < self.window]
            key = hashlib.sha256((prompt or "").encode()).hexdigest()

            if self.cache_on and key in self._cache:
                self._seal({"event": "cache_hit", "key": key[:16], "saved": self.cost(prompt)})
                return self._cache[key]

            if len(self._recent) >= self.burst:
                self._seal({"event": "refused", "why": "runaway", "calls": len(self._recent)})
                raise BudgetSpent("%d calls in %.0f seconds looks like a runaway loop"
                                  % (len(self._recent), self.window))

            due = self.cost(prompt)
            if self._spent + due > self.limit:
                self._seal({"event": "refused", "why": "budget", "spent": round(self._spent, 6),
                            "limit": self.limit})
                raise BudgetSpent("this call needs £%.4f and only £%.4f is left today"
                                  % (due, self.remaining()))

            self._recent.append(now)
            self._spent += due
            self._seal({"event": "allowed", "key": key[:16], "tokens": self.tokens(prompt),
                        "cost": round(due, 6), "spent": round(self._spent, 6)})
            return None

    def record(self, prompt, answer):
        if self.cache_on:
            with self._lock:
                self._cache[hashlib.sha256((prompt or "").encode()).hexdigest()] = answer

    def meter(self, fn):
        """Decorator. The first string argument is treated as the prompt."""
        @functools.wraps(fn)
        def inner(*args, **kwargs):
            prompt = next((a for a in args if isinstance(a, str)),
                          next((v for v in kwargs.values() if isinstance(v, str)), ""))
            cached = self.check(prompt)
            if cached is not None:
                return cached
            out = fn(*args, **kwargs)
            self.record(prompt, out)
            return out
        return inner

    # ---------------------------------------------------------------- audit
    def verify_log(self):
        """Re-walk the local log. Returns (ok, lines_checked, first_bad_line)."""
        tip, n = "GENESIS", 0
        try:
            with open(self.log_path, "r", encoding="utf-8") as f:
                for n, raw in enumerate(f, 1):
                    row = json.loads(raw)
                    claimed = row.pop("hash", None)
                    row["prev"] = tip
                    body = json.dumps(row, sort_keys=True, separators=(",", ":"))
                    tip = hashlib.sha256(body.encode()).hexdigest()
                    if claimed != tip:
                        return False, n, n
        except FileNotFoundError:
            return True, 0, None
        return True, n, None


if __name__ == "__main__":
    g = Guard(daily_limit_gbp=0.01, price_per_million=3.00)

    @g.meter
    def ask(prompt):
        return "answer to: " + prompt[:30]

    print("sebbi.pro Spend Guard", __version__)
    print(" first call     :", ask("tell me about AI governance " * 20)[:40])
    print(" same again     :", ask("tell me about AI governance " * 20)[:40], "(served from cache, cost nothing)")
    print(" spent today    : £%.5f of £%.2f" % (g.spent_today(), g.limit))
    try:
        ask("a very expensive new question " * 2000)
    except BudgetSpent as e:
        print(" refused        :", e)
    ok, n, bad = g.verify_log()
    print(" local log      :", "intact" if ok else "BROKEN at line %s" % bad, "(%d lines)" % n)
    print(" more tools     : https://sebbi.pro/tools")
