#!/usr/bin/env python3
"""Reference sample: pull DynamoGuard moderation logs from the DynamoAI API as NDJSON.

Not a supported product. Behaviour and limits: "Pull Moderation Logs From The API" in the DynamoAI docs.
"""
import argparse
import json
import os
import sys
import time
import urllib.error
import urllib.parse
import urllib.request
from datetime import datetime, timezone

RETRYABLE_STATUS = {429, 500, 502, 503, 504}


def parse_args(argv):
    parser = argparse.ArgumentParser(
        description="Pull DynamoGuard moderation logs and write them to stdout as NDJSON."
    )
    parser.add_argument("--base-url", required=True, help="API base URL, for example https://api.example.com")
    target = parser.add_mutually_exclusive_group(required=True)
    target.add_argument("--application-id", type=int, help="Read every AI system in this application (project)")
    target.add_argument("--ai-system-id", help="Read one AI system")
    parser.add_argument("--state-file", default="moderation_logs_state.json")
    parser.add_argument("--since", help="First run only: ISO 8601 time or UTC seconds to start from")
    parser.add_argument("--overlap", type=float, default=300.0, help="Seconds re-read on every run (default 300)")
    parser.add_argument("--lag", type=float, default=30.0, help="Stop this many seconds before now (default 30)")
    parser.add_argument("--window", type=float, default=3600.0, help="Largest window per request, seconds (default 3600)")
    parser.add_argument("--max-batch", type=int, default=1000, help="Split windows holding more logs (default 1000)")
    parser.add_argument("--timeout", type=float, default=60.0)
    parser.add_argument("--retries", type=int, default=5)
    return parser.parse_args(argv)


def to_seconds(value):
    try:
        return float(value)
    except ValueError:
        parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
        if parsed.tzinfo is None:
            parsed = parsed.replace(tzinfo=timezone.utc)
        return parsed.timestamp()


class Client:
    def __init__(self, args, token):
        self.base = args.base_url.rstrip("/")
        if args.application_id is not None:
            self.logs_url = f"{self.base}/v1/applications/{args.application_id}/logs"
        else:
            self.logs_url = f"{self.base}/v1/moderation/model/{urllib.parse.quote(args.ai_system_id, safe='')}/logs"
        self.token = token
        self.timeout = args.timeout
        self.retries = args.retries
        self.max_batch = args.max_batch

    def get(self, url, params=None):
        query = f"?{urllib.parse.urlencode(params)}" if params else ""
        request = urllib.request.Request(
            f"{url}{query}",
            headers={"Authorization": f"Bearer {self.token}", "Accept": "application/json"},
        )
        delay = 1.0
        for attempt in range(self.retries + 1):
            try:
                with urllib.request.urlopen(request, timeout=self.timeout) as response:
                    return json.load(response)
            except urllib.error.HTTPError as err:
                if err.code not in RETRYABLE_STATUS or attempt == self.retries:
                    raise
                retry_after = err.headers.get("Retry-After", "")
                wait = float(retry_after) if retry_after.isdigit() else delay
            except urllib.error.URLError:
                if attempt == self.retries:
                    raise
                wait = delay
            time.sleep(wait)
            delay = min(delay * 2, 60.0)
        raise RuntimeError("unreachable")

    def attached_ai_systems(self, application_id):
        return self.get(f"{self.base}/v1/applications/{application_id}").get("aiSystems") or []

    @staticmethod
    def window_params(start, end):
        return {"startTime": f"{start:.3f}", "endTime": f"{end:.3f}", "sortDir": "asc"}

    def count(self, start, end):
        body = self.get(f"{self.logs_url}/count", self.window_params(start, end))
        # The applications endpoint returns a bare number; the AI system endpoint returns {"count": n}.
        return body["count"] if isinstance(body, dict) else int(body)

    def fetch(self, start, end):
        total = self.count(start, end)
        if total == 0:
            return []
        if total > self.max_batch and end - start > 1.0:
            middle = (start + end) / 2
            return self.fetch(start, middle) + self.fetch(middle, end)
        params = dict(self.window_params(start, end), perPage=0)
        return self.get(self.logs_url, params)["logs"]


def load_state(path):
    if not os.path.exists(path):
        return {"position": None, "seen": {}}
    with open(path, encoding="utf-8") as handle:
        return json.load(handle)


def save_state(path, state):
    temporary = f"{path}.tmp"
    with open(temporary, "w", encoding="utf-8") as handle:
        json.dump(state, handle)
    os.replace(temporary, path)


def main(argv):
    args = parse_args(argv)
    token = os.environ.get("DYNAMOAI_API_KEY")
    if not token:
        sys.exit("Set DYNAMOAI_API_KEY to the service user's API key.")
    client = Client(args, token)
    try:
        if args.application_id is not None and not client.attached_ai_systems(args.application_id):
            sys.exit(
                f"Project {args.application_id} has no attached AI systems. On release 3.26 its logs endpoint "
                "then returns logs from every AI system. Attach an AI system, or use --ai-system-id."
            )
    except urllib.error.HTTPError as err:
        sys.exit(f"HTTP {err.code} from {err.url}: {err.read().decode(errors='replace')[:300]}")
    state = load_state(args.state_file)
    seen = state["seen"]

    stop = time.time() - args.lag
    if state["position"] is not None:
        start = state["position"] - args.overlap
    elif args.since:
        start = to_seconds(args.since)
    else:
        start = stop - args.overlap

    written = 0
    while start < stop:
        end = min(start + args.window, stop)
        try:
            logs = client.fetch(start, end)
        except urllib.error.HTTPError as err:
            sys.exit(f"HTTP {err.code} from {err.url}: {err.read().decode(errors='replace')[:300]}")
        for record in sorted(logs, key=lambda item: (item["timestamp"], item["id"])):
            if record["id"] in seen:
                continue
            sys.stdout.write(json.dumps(record, separators=(",", ":")) + "\n")
            seen[record["id"]] = to_seconds(record["timestamp"])
            written += 1
        sys.stdout.flush()
        horizon = end - args.overlap
        state = {"position": end, "seen": {key: ts for key, ts in seen.items() if ts >= horizon}}
        seen = state["seen"]
        save_state(args.state_file, state)
        start = end

    print(f"wrote {written} new logs", file=sys.stderr)


if __name__ == "__main__":
    main(sys.argv[1:])
