#!/usr/bin/env python3
"""LineClarity — verify a rating record in Python.

Copy this file into your own project. One dependency (`cryptography`), no network calls except the one you
make yourself to fetch the feed and the public key. Run it directly to check a live line:

    python verify_lineclarity.py --key lc_live_xxx --line ABIL-345-01

What a correct integration checks, in order — all four, every pull:

    1. the Ed25519 signature verifies against our published public key
    2. `valid_until` has not passed, and the record is younger than `failsafe.max_age_s`
    3. `sequence` has not gone backwards. The SAME sequence with the same signature is simply the record you
       already hold, re-served between recomputations - keep using it. Only a lower sequence, or the same
       sequence with a different signature, is a problem.
    4. `status` is OK (DEGRADED is usable but worth an alarm; FAILSAFE means we are already holding static)

If any check fails, use `static_a`. That is the whole safety story: a bad, stale, forged or replayed record
can only ever put you back on the rating you use today.
"""
import argparse
import base64
import hashlib
import json
import ssl
import sys
import urllib.error
import urllib.request
from datetime import datetime, timezone

from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey

BASE = "https://lineclarity.com"


# ---------------------------------------------------------------- canonical form + signature

def canonical_bytes(obj) -> bytes:
    """Exactly what the engine signed: keys sorted, no whitespace, UTF-8, no NaN."""
    return json.dumps(
        obj, sort_keys=True, separators=(",", ":"), ensure_ascii=False, allow_nan=False
    ).encode("utf-8")


def signed_bytes(record: dict) -> bytes:
    return canonical_bytes({k: v for k, v in record.items() if k != "signature"})


def verify_signature(record: dict, public_key_raw_b64: str) -> bool:
    sig = record.get("signature") or {}
    if sig.get("alg") != "Ed25519" or not sig.get("value"):
        return False
    payload = signed_bytes(record)
    if sig.get("digest") and sig["digest"] != "sha256:" + hashlib.sha256(payload).hexdigest():
        return False
    try:
        Ed25519PublicKey.from_public_bytes(base64.b64decode(public_key_raw_b64)).verify(
            base64.b64decode(sig["value"]), payload)
        return True
    except Exception:
        return False


# ---------------------------------------------------------------- the other three checks

def _parse(ts):
    return datetime.strptime(ts, "%Y-%m-%dT%H:%M:%SZ").replace(tzinfo=timezone.utc)


def check_freshness(record: dict, now=None):
    now = now or datetime.now(timezone.utc)
    age = (now - _parse(record["issued_at"])).total_seconds()
    max_age = (record.get("failsafe") or {}).get("max_age_s", 1800)
    if now > _parse(record["valid_until"]):
        return False, f"past valid_until (age {int(age)}s)"
    if age > max_age:
        return False, f"older than max_age_s ({int(age)}s > {max_age}s)"
    return True, f"age {int(age)}s"


def check_sequence(record: dict, last_accepted=None):
    """`last_accepted` is None, an int, or (sequence, digest).

    Receiving the same record again is normal - the feed recomputes on its own cadence, so a poll in between
    returns the record you already hold. That is the current answer, not a replay, and must not send you to
    static. Reject only a sequence that goes backwards, or the same sequence arriving with a different
    signature (two records claiming one position).
    """
    if last_accepted is None:
        return True, "first record"
    last_seq, last_digest = last_accepted if isinstance(last_accepted, tuple) else (last_accepted, None)
    digest = (record.get("signature") or {}).get("digest")
    if record["sequence"] > last_seq:
        return True, ""
    if record["sequence"] == last_seq:
        if last_digest is None or last_digest == digest:
            return True, "unchanged - same record re-served"
        return False, "same sequence, different signature"
    return False, "sequence went backwards (replay or reorder)"


def rating_to_use(record: dict, public_key_raw_b64: str, last_accepted=None):
    """The one call to wrap: returns (amps, source, reason)."""
    if not verify_signature(record, public_key_raw_b64):
        return record["static_a"], "static", "signature did not verify"
    ok, why = check_freshness(record)
    if not ok:
        return record["static_a"], "static", why
    ok, why = check_sequence(record, last_accepted)
    if not ok:
        return record["static_a"], "static", why
    if record["status"] == "FAILSAFE":
        return record["static_a"], "static", "feed is in fail-safe: " + record.get("status_reason", "")
    return record["rating_a"], "lineclarity", record.get("status_reason", "")


# ---------------------------------------------------------------- fetching (only for the demo run)

# Some machines — Windows in particular — carry an EXPIRED ISRG Root X2 cross-sign in their
# certificate store. Python built against OpenSSL 1.1.1 (which is most Pythons before 3.12 on
# Windows) follows that dead path instead of the valid one sitting next to it, and reports
# "certificate has expired". Nothing is wrong with our certificate; it is a path-building bug in
# the client. Rather than make you tidy your trust store to read a rating, this falls back to a
# known-good root set. Verification stays ON either way — this never disables TLS checking.
_ISRG_ROOTS = """
-----BEGIN CERTIFICATE-----
MIIFazCCA1OgAwIBAgIRAIIQz7DSQONZRGPgu2OCiwAwDQYJKoZIhvcNAQELBQAw
TzELMAkGA1UEBhMCVVMxKTAnBgNVBAoTIEludGVybmV0IFNlY3VyaXR5IFJlc2Vh
cmNoIEdyb3VwMRUwEwYDVQQDEwxJU1JHIFJvb3QgWDEwHhcNMTUwNjA0MTEwNDM4
WhcNMzUwNjA0MTEwNDM4WjBPMQswCQYDVQQGEwJVUzEpMCcGA1UEChMgSW50ZXJu
ZXQgU2VjdXJpdHkgUmVzZWFyY2ggR3JvdXAxFTATBgNVBAMTDElTUkcgUm9vdCBY
MTCCAiIwDQYJKoZIhvcNAQEBBQADggIPADCCAgoCggIBAK3oJHP0FDfzm54rVygc
h77ct984kIxuPOZXoHj3dcKi/vVqbvYATyjb3miGbESTtrFj/RQSa78f0uoxmyF+
0TM8ukj13Xnfs7j/EvEhmkvBioZxaUpmZmyPfjxwv60pIgbz5MDmgK7iS4+3mX6U
A5/TR5d8mUgjU+g4rk8Kb4Mu0UlXjIB0ttov0DiNewNwIRt18jA8+o+u3dpjq+sW
T8KOEUt+zwvo/7V3LvSye0rgTBIlDHCNAymg4VMk7BPZ7hm/ELNKjD+Jo2FR3qyH
B5T0Y3HsLuJvW5iB4YlcNHlsdu87kGJ55tukmi8mxdAQ4Q7e2RCOFvu396j3x+UC
B5iPNgiV5+I3lg02dZ77DnKxHZu8A/lJBdiB3QW0KtZB6awBdpUKD9jf1b0SHzUv
KBds0pjBqAlkd25HN7rOrFleaJ1/ctaJxQZBKT5ZPt0m9STJEadao0xAH0ahmbWn
OlFuhjuefXKnEgV4We0+UXgVCwOPjdAvBbI+e0ocS3MFEvzG6uBQE3xDk3SzynTn
jh8BCNAw1FtxNrQHusEwMFxIt4I7mKZ9YIqioymCzLq9gwQbooMDQaHWBfEbwrbw
qHyGO0aoSCqI3Haadr8faqU9GY/rOPNk3sgrDQoo//fb4hVC1CLQJ13hef4Y53CI
rU7m2Ys6xt0nUW7/vGT1M0NPAgMBAAGjQjBAMA4GA1UdDwEB/wQEAwIBBjAPBgNV
HRMBAf8EBTADAQH/MB0GA1UdDgQWBBR5tFnme7bl5AFzgAiIyBpY9umbbjANBgkq
hkiG9w0BAQsFAAOCAgEAVR9YqbyyqFDQDLHYGmkgJykIrGF1XIpu+ILlaS/V9lZL
ubhzEFnTIZd+50xx+7LSYK05qAvqFyFWhfFQDlnrzuBZ6brJFe+GnY+EgPbk6ZGQ
3BebYhtF8GaV0nxvwuo77x/Py9auJ/GpsMiu/X1+mvoiBOv/2X/qkSsisRcOj/KK
NFtY2PwByVS5uCbMiogziUwthDyC3+6WVwW6LLv3xLfHTjuCvjHIInNzktHCgKQ5
ORAzI4JMPJ+GslWYHb4phowim57iaztXOoJwTdwJx4nLCgdNbOhdjsnvzqvHu7Ur
TkXWStAmzOVyyghqpZXjFaH3pO3JLF+l+/+sKAIuvtd7u+Nxe5AW0wdeRlN8NwdC
jNPElpzVmbUq4JUagEiuTDkHzsxHpFKVK7q4+63SM1N95R1NbdWhscdCb+ZAJzVc
oyi3B43njTOQ5yOf+1CceWxG1bQVs5ZufpsMljq4Ui0/1lvh+wjChP4kqKOJ2qxq
4RgqsahDYVvTH9w7jXbyLeiNdd8XM2w9U/t7y0Ff/9yi0GE44Za4rF2LN9d11TPA
mRGunUHBcnWEvgJBQl9nJEiU0Zsnvgc/ubhPgXRR4Xq37Z0j4r7g1SgEEzwxA57d
emyPxgcYxn/eR44/KJ4EBs+lVDR3veyJm+kXQ99b21/+jh5Xos1AnX5iItreGCc=
-----END CERTIFICATE-----
-----BEGIN CERTIFICATE-----
MIICGzCCAaGgAwIBAgIQQdKd0XLq7qeAwSxs6S+HUjAKBggqhkjOPQQDAzBPMQsw
CQYDVQQGEwJVUzEpMCcGA1UEChMgSW50ZXJuZXQgU2VjdXJpdHkgUmVzZWFyY2gg
R3JvdXAxFTATBgNVBAMTDElTUkcgUm9vdCBYMjAeFw0yMDA5MDQwMDAwMDBaFw00
MDA5MTcxNjAwMDBaME8xCzAJBgNVBAYTAlVTMSkwJwYDVQQKEyBJbnRlcm5ldCBT
ZWN1cml0eSBSZXNlYXJjaCBHcm91cDEVMBMGA1UEAxMMSVNSRyBSb290IFgyMHYw
EAYHKoZIzj0CAQYFK4EEACIDYgAEzZvVn4CDCuwJSvMWSj5cz3es3mcFDR0HttwW
+1qLFNvicWDEukWVEYmO6gbf9yoWHKS5xcUy4APgHoIYOIvXRdgKam7mAHf7AlF9
ItgKbppbd9/w+kHsOdx1ymgHDB/qo0IwQDAOBgNVHQ8BAf8EBAMCAQYwDwYDVR0T
AQH/BAUwAwEB/zAdBgNVHQ4EFgQUfEKWrt5LSDv6kviejM9ti6lyN5UwCgYIKoZI
zj0EAwMDaAAwZQIwe3lORlCEwkSHRhtFcP9Ymd70/aTSVaYgLXTWNLxBo1BfASdW
tL4ndQavEi51mI38AjEAi/V3bNTIZargCyzuFJ0nN6T5U6VR5CmD1/iQMVtCnwr1
/q4AaOeMSQ+2b1tbFfLn
-----END CERTIFICATE-----
"""


_WARNED = []


def _context():
    return ssl.create_default_context()


def _fallback_context():
    ctx = ssl.create_default_context()
    ctx.load_verify_locations(cadata=_ISRG_ROOTS)
    return ctx


def _get(url, key=None):
    req = urllib.request.Request(url, headers={"X-API-Key": key} if key else {})
    try:
        with urllib.request.urlopen(req, timeout=30, context=_context()) as r:
            return json.loads(r.read().decode())
    except urllib.error.URLError as e:
        if not isinstance(getattr(e, "reason", None), ssl.SSLCertVerificationError):
            raise
        if not _WARNED:
            print("  note: your system certificate store rejected this chain. That is a known\n"
                  "        OpenSSL 1.1.1 path-building issue on some machines, not a problem with\n"
                  "        the server. Retrying against the published ISRG roots — TLS\n"
                  "        verification is still enforced.", file=sys.stderr)
            _WARNED.append(1)
        with urllib.request.urlopen(req, timeout=30, context=_fallback_context()) as r:
            return json.loads(r.read().decode())


def public_key_for(record, base=BASE):
    keys = _get(f"{base}/api/v1/pubkey")["keys"]
    want = (record.get("signature") or {}).get("key_id")
    for k in keys:
        if k["key_id"] == want:
            return k["public_key_raw_b64"]
    raise SystemExit(f"the record was signed with {want}, which is not a published key — do not trust it")


def main():
    p = argparse.ArgumentParser(description="Verify a LineClarity rating record")
    p.add_argument("--key", help="your API key (omit to use the public demo feed)")
    p.add_argument("--line", help="your line_id")
    p.add_argument("--base", default=BASE)
    p.add_argument("--file", help="verify a record already saved to disk instead of fetching one")
    p.add_argument("--last-sequence", type=int, default=None)
    a = p.parse_args()

    if a.file:
        record = json.load(open(a.file, encoding="utf-8"))
    elif a.key and a.line:
        record = _get(f"{a.base}/api/v1/rating?line={a.line}&format=json", a.key)
    else:
        # Penwortham-Kirkby, the same GB 400 kV circuit the site demo uses. Both ends, because a
        # circuit with no route is rated at one point on the worst-case wind angle and comes back
        # DEGRADED - a poor first impression for something meant to be tried in one command.
        record = _get(f"{a.base}/api/v1/demo/current?lat=53.75&lon=-2.72"
                      f"&lat2=53.48&lon2=-2.89&kv=400&conductor=Zebra%20ACSR%20%28quad%29")

    pub = public_key_for(record, a.base)
    amps, source, reason = rating_to_use(record, pub, a.last_sequence)

    print(f"line        {record['line']['line_id']}  ({record['line'].get('ext_ref') or 'no ext_ref'})")
    print(f"issued      {record['issued_at']}   valid until {record['valid_until']}")
    print(f"sequence    {record['sequence']}   status {record['status']} {record.get('status_reason','')}")
    print(f"signature   {'VERIFIED' if verify_signature(record, pub) else 'FAILED'}"
          f"   key {record['signature']['key_id']}")
    print(f"static      {record['static_a']} A  ({record.get('static_basis','')})")
    print(f"USE         {amps} A   (source: {source}{'; ' + reason if reason else ''})")
    sm = record.get("summary") or {}
    if "avg_gain_pct" in sm:
        # 0% is a real answer: on a mild, still hour the safe limit sits on the static floor. Printing it
        # either way is the point - a rating you can only see when it flatters us is not worth verifying.
        print(f"uplift      {sm['avg_gain_pct']}%  (+{sm.get('avg_extra_mw', 0)} MW) averaged over the next "
              f"{record.get('horizon_hours', len(record.get('horizon', [])))}h")
    rt = record.get("route") or {}
    if rt.get("basis") == "route":
        print(f"spans       worst of {rt['spans_modelled']} sampled over {rt['length_km']} km")
    return 0 if source == "lineclarity" else 1


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