#!/usr/bin/env python3
"""Pull captured webhooks into localhost. Python 3.10+, no dependencies."""

import argparse
import base64
import json
import os
import sys
import time
import urllib.error
import urllib.parse
import urllib.request

HOP = {
    "host",
    "content-length",
    "connection",
    "transfer-encoding",
    "keep-alive",
    "upgrade",
    "te",
    "trailer",
    "proxy-authenticate",
    "proxy-authorization",
}


class NoRedirect(urllib.request.HTTPRedirectHandler):
    def redirect_request(self, *args, **kwargs):
        return None


OPENER = urllib.request.build_opener(urllib.request.ProxyHandler({}), NoRedirect())


class APIError(Exception):
    def __init__(self, status, payload):
        self.status, self.payload = status, payload
        super().__init__(
            f"API returned {status}: {payload.get('error',{}).get('code','error')}"
        )


def api(args, path, data=None):
    body = json.dumps(data).encode() if data is not None else None
    request = urllib.request.Request(
        args.api.rstrip("/") + path,
        data=body,
        headers={
            "Authorization": "Bearer " + args.key,
            "Content-Type": "application/json",
        },
    )
    try:
        with OPENER.open(request, timeout=40) as response:
            return json.load(response)
    except urllib.error.HTTPError as exc:
        try:
            payload = json.load(exc)
        except ValueError:
            payload = {}
        raise APIError(exc.code, payload) from None


def latest(args):
    result = api(args, f"/endpoints/{args.endpoint}/requests?limit=1")
    return result["items"][0]["id"] if result["items"] else None


def batches(args, cursor):
    query = {"limit": 100}
    if cursor:
        query["after"] = cursor
    result = api(
        args, f"/endpoints/{args.endpoint}/requests?" + urllib.parse.urlencode(query)
    )
    # Fetch the entire bounded interval before advancing the live cursor.
    newest = result["items"][0]["id"] if result["items"] else cursor
    rows = result["items"]
    while result.get("next"):
        result = api(
            args,
            f"/endpoints/{args.endpoint}/requests?"
            + urllib.parse.urlencode({"before": result["next"], "limit": 100}),
        )
        stop = False
        for row in result["items"]:
            if row["id"] == cursor:
                stop = True
                break
            rows.append(row)
        if stop:
            break
    return list(reversed(rows)), newest


def local_url(target, path, query):
    parsed = urllib.parse.urlsplit(target)
    return urllib.parse.urlunsplit(
        (
            parsed.scheme,
            parsed.netloc,
            parsed.path.rstrip("/") + "/" + path.lstrip("/"),
            query,
            "",
        )
    )


def deliver(args, row):
    obj = api(args, "/requests/" + row["id"])
    body = (
        base64.b64decode(obj["body_base64"], validate=True)
        if "body_base64" in obj
        else obj.get("body", "").encode()
    )
    url = local_url(args.to, obj["path"], obj.get("query_string", ""))
    tokens = {
        p.strip().lower()
        for k, v in obj["headers"]
        if k.lower() == "connection"
        for p in v.split(",")
    }
    headers = {
        k: v
        for k, v in obj["headers"]
        if k.lower() not in HOP | tokens
        and not k.lower().startswith(("proxy-", "x-forwarded-", "x-ammo-"))
    }
    headers.update({"X-Ammo-Request-Id": obj["id"], "X-Ammo-Delivery": "cli"})
    report = {
        "via": "cli",
        "sent_url": url,
        "response_status": None,
        "response_headers": [],
        "response_body_base64": "",
        "duration_ms": 0,
        "error": "",
    }
    start = time.monotonic()
    request = urllib.request.Request(
        url, data=body, headers=headers, method=obj["method"]
    )
    try:
        try:
            response = OPENER.open(request, timeout=10)
        except urllib.error.HTTPError as exc:
            response = exc
        with response:
            report.update(
                response_status=response.code,
                response_headers=list(response.headers.items())[:100],
                response_body_base64=base64.b64encode(response.read(65536)).decode(),
            )
    except (OSError, urllib.error.URLError) as exc:
        report["error"] = type(exc).__name__
    report["duration_ms"] = round((time.monotonic() - start) * 1000)
    try:
        api(args, "/requests/" + obj["id"] + "/deliveries", report)
    except (APIError, OSError, urllib.error.URLError) as exc:
        print(
            f"{obj['id']}: delivered locally; result report failed ({type(exc).__name__}). Not replaying.",
            file=sys.stderr,
        )
    print(
        f"{obj['id']} {obj['method']} {obj['path']} -> {report['response_status'] or report['error']} ({report['duration_ms']}ms)",
        flush=True,
    )


def main(argv=None):
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--key", default=os.environ.get("AMMO_KEY"))
    parser.add_argument("--endpoint", required=True)
    parser.add_argument("--to", required=True)
    parser.add_argument("--api", default="https://ammo.tools/api/v1")
    parser.add_argument("--interval", type=int, default=2)
    parser.add_argument("--replay-last", type=int, default=0)
    args = parser.parse_args(argv)
    target = urllib.parse.urlsplit(args.to)
    if not args.key or not 1 <= args.interval <= 60 or args.replay_last < 0:
        parser.error(
            "Set --key or AMMO_KEY; interval must be 1–60 and replay-last must be nonnegative."
        )
    if (
        target.scheme != "http"
        or target.hostname not in {"localhost", "127.0.0.1", "::1"}
        or not target.port
        or target.username
        or target.password
        or target.fragment
        or target.query
    ):
        parser.error(
            "--to must be an HTTP localhost URL with an explicit port and no credentials, query or fragment."
        )
    try:
        cursor = latest(args)
        if args.replay_last:
            rows = []
            before = None
            while len(rows) < args.replay_last:
                query = {"limit": min(100, args.replay_last - len(rows))}
                if before:
                    query["before"] = before
                result = api(
                    args,
                    f"/endpoints/{args.endpoint}/requests?"
                    + urllib.parse.urlencode(query),
                )
                rows.extend(result["items"])
                if not result.get("next"):
                    break
                before = result["next"]
            for row in reversed(rows):
                deliver(args, row)
        print("Listening for new captures. Press Ctrl+C to stop.", flush=True)
        while True:
            try:
                rows, newest = batches(args, cursor)
                for row in rows:
                    deliver(args, row)
                    cursor = row["id"]
                cursor = newest
            except APIError as exc:
                if exc.status in {401, 403}:
                    print(str(exc), file=sys.stderr)
                    return 1
                if exc.payload.get("error", {}).get("code") == "invalid_cursor":
                    cursor = latest(args)
                    print(
                        "Cursor expired; resumed from newest capture.", file=sys.stderr
                    )
                else:
                    print(str(exc), file=sys.stderr)
            except (OSError, urllib.error.URLError) as exc:
                print(
                    f"Connection unavailable ({type(exc).__name__}); polling again shortly.",
                    file=sys.stderr,
                )
            time.sleep(args.interval)
    except KeyboardInterrupt:
        print("Stopped.")
    except (APIError, OSError, urllib.error.URLError) as exc:
        print(str(exc), file=sys.stderr)
        return 1
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
