"""
RailGuard — Python guardrail SDK.

Wrap every spending action your agent takes:

    from guardrail import Guardrail, SpendBlocked

    guardrail = Guardrail(api_key="rg_...", base_url="https://railguardsecurity.com")
    try:
        payment = guardrail.spend(
            amount=45.00,
            merchant="OpenAI API",
            category="AI & Compute",
            reasoning="I need to process 500 PDF invoices uploaded by the user.",
            prompt_chain=messages,          # [{"role": "user", "content": "..."}, ...]
            rail="stripe",                  # "stripe" | "ramp" | "crypto"
        )
        card = payment["credential"]        # single-use card or crypto tx receipt
    except SpendBlocked as e:
        print("Stop or choose an alternative:", e.reason)

Flagged (out-of-policy) spend pauses: spend() waits until a human decides (wait=True),
or returns immediately with status "flagged" when wait=False (use webhook_url + poll()).
Requires: requests (pip install requests)
"""
from __future__ import annotations

import hashlib
import hmac
import time
from typing import Any, Optional

import requests


class SpendBlocked(Exception):
    def __init__(self, reason: str, code: str, response: dict):
        super().__init__(f"{code}: {reason}")
        self.reason, self.code, self.response = reason, code, response


class SpendPending(Exception):
    """Raised when waiting for a human decision times out."""


class Guardrail:
    def __init__(self, api_key: str, base_url: str, timeout: float = 15.0):
        self.api_key = api_key
        self.base_url = base_url.rstrip("/")
        self.timeout = timeout

    def _headers(self) -> dict:
        return {"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"}

    def spend(
        self,
        amount: float,
        merchant: str,
        category: str = "General",
        reasoning: str = "",
        prompt_chain: Optional[list] = None,
        rail: str = "stripe",
        justification: str = "",
        recipient_address: Optional[str] = None,
        webhook_url: Optional[str] = None,
        wait: bool = True,
        wait_timeout: float = 3600,
        poll_interval: float = 5,
        idempotency_key: Optional[str] = None,
    ) -> dict:
        body: dict[str, Any] = {
            "amount": amount, "merchant": merchant, "category": category, "rail": rail,
            "reasoning": reasoning, "justification": justification or reasoning[:1000],
        }
        if prompt_chain: body["prompt_chain"] = prompt_chain
        if recipient_address: body["recipient_address"] = recipient_address
        if webhook_url: body["webhook_url"] = webhook_url
        # Reuse the same key when retrying: same key + same body never issues a second card.
        if idempotency_key: body["idempotency_key"] = idempotency_key

        r = requests.post(f"{self.base_url}/api/public/transactions", json=body, headers=self._headers(), timeout=self.timeout)
        data = r.json()
        if r.status_code == 409:
            raise SpendBlocked(data.get("instruction", data.get("error", "")), data.get("error_code", "IDEMPOTENCY_MISMATCH"), data)
        if r.status_code == 429:
            raise SpendPending("Rate limited; retry after 60s")
        if r.status_code in (400, 401, 404, 500):
            raise RuntimeError(data.get("error", f"HTTP {r.status_code}"))
        return self._resolve(data, wait, wait_timeout, poll_interval)

    def poll(self, transaction_id: str) -> dict:
        r = requests.get(f"{self.base_url}/api/public/transactions/{transaction_id}", headers=self._headers(), timeout=self.timeout)
        return r.json()

    def report_receipt(self, transaction_id: str, amount: float, merchant: str = "",
                       recipient_address: Optional[str] = None, receipt_ref: Optional[str] = None,
                       receipt_url: Optional[str] = None) -> dict:
        """After paying, report the final charge so the ledger can verify it matches the approval."""
        body = {k: v for k, v in {"amount": amount, "merchant": merchant, "recipient_address": recipient_address,
                "receipt_ref": receipt_ref, "receipt_url": receipt_url}.items() if v is not None}
        r = requests.post(f"{self.base_url}/api/public/transactions/{transaction_id}/receipt",
                          json=body, headers=self._headers(), timeout=self.timeout)
        return r.json()

    def log_payment(self, amount: float, merchant: str, category: str = "General", reasoning: str = "",
                    receipt_ref: Optional[str] = None, recipient_address: Optional[str] = None) -> dict:
        """Log-only mode: record a payment the agent already made."""
        body: dict[str, Any] = {"mode": "report", "amount": amount, "merchant": merchant, "category": category,
                                "reasoning": reasoning, "receipt": {"amount": amount, "merchant": merchant}}
        if receipt_ref: body["receipt"]["receipt_ref"] = receipt_ref
        if recipient_address:
            body["recipient_address"] = recipient_address
            body["receipt"]["recipient_address"] = recipient_address
        r = requests.post(f"{self.base_url}/api/public/transactions", json=body, headers=self._headers(), timeout=self.timeout)
        return r.json()

    def _resolve(self, data: dict, wait: bool, wait_timeout: float, poll_interval: float) -> dict:
        status = data.get("status")
        if status in ("blocked", "rejected"):
            raise SpendBlocked(data.get("reason", ""), data.get("error_code", status.upper()), data)
        if status == "flagged" and wait:
            deadline = time.time() + wait_timeout
            while time.time() < deadline:
                time.sleep(poll_interval)
                data = self.poll(data["transaction_id"])
                if data.get("status") != "flagged":
                    return self._resolve(data, False, 0, 0)
            raise SpendPending(f"No human decision within {wait_timeout}s for {data['transaction_id']}")
        payment = data.get("payment") or {}
        if status == "approved" and payment.get("status") not in (None, "issued"):
            raise RuntimeError(f"Approved but payment {payment.get('status')}: {payment.get('detail')}")
        return {"status": status, "transaction_id": data.get("transaction_id"), **payment}

    def verify_webhook(self, raw_body: bytes, timestamp: str, signature: str, tolerance: int = 300) -> bool:
        """Verify an x-railguard-signature header. Key = sha256(api_key)."""
        if abs(time.time() - int(timestamp)) > tolerance:
            return False
        key = hashlib.sha256(self.api_key.encode()).hexdigest().encode()
        expected = hmac.new(key, f"{timestamp}.".encode() + raw_body, hashlib.sha256).hexdigest()
        return hmac.compare_digest(expected, signature)
