import argparse
import json
import os
import re
import shutil
import subprocess
import sys
from pathlib import Path


AUTOMATION_ID_PATTERN = re.compile(r"^Automation ID:\s*([A-Za-z0-9._-]+)\s*$", re.MULTILINE)
HOOK_SUCCESS_STATUSES = {"archived", "non_automation", "unknown_automation"}


def _payload_field(payload, key):
    if key in payload:
        return payload[key]
    nested = payload.get("payload")
    if isinstance(nested, dict):
        return nested.get(key)
    return None


def _iter_candidate_texts(transcript_path):
    with Path(transcript_path).open("r", encoding="utf-8") as handle:
        for line in handle:
            line = line.strip()
            if not line:
                continue
            try:
                entry = json.loads(line)
            except json.JSONDecodeError:
                continue

            payload = entry.get("payload", {})
            if not isinstance(payload, dict):
                continue

            message = payload.get("message")
            if isinstance(message, str):
                yield message

            content = payload.get("content")
            if isinstance(content, list):
                for item in content:
                    if not isinstance(item, dict):
                        continue
                    text = item.get("text")
                    if isinstance(text, str):
                        yield text


def find_automation_id(transcript_path):
    for text in _iter_candidate_texts(transcript_path):
        match = AUTOMATION_ID_PATTERN.search(text)
        if match:
            return match.group(1)
    return None


def is_known_automation(codex_home, automation_id):
    automation_toml = Path(codex_home) / "automations" / automation_id / "automation.toml"
    if not automation_toml.exists():
        return False

    try:
        contents = automation_toml.read_text(encoding="utf-8")
    except OSError:
        return False

    return f'id = "{automation_id}"' in contents


def default_runner(args):
    completed = subprocess.run(args, capture_output=True, text=True, check=False)
    return completed.returncode, completed.stdout, completed.stderr


def resolve_archive_command(base_command=None, which=shutil.which, is_windows=(os.name == "nt")):
    if base_command is not None:
        return list(base_command)

    if is_windows:
        resolved = which("codex.cmd") or which("codex.exe") or which("codex")
    else:
        resolved = which("codex")

    return [resolved or "codex", "archive"]


def archive_session_from_stop_payload(payload, codex_home=None, runner=None, archive_command=None):
    session_id = _payload_field(payload, "session_id")
    transcript_path = _payload_field(payload, "transcript_path")

    if not session_id:
        return {"status": "missing_session_id"}

    if not transcript_path:
        return {"status": "missing_transcript_path", "session_id": session_id}

    transcript = Path(transcript_path)
    if not transcript.exists():
        return {
            "status": "missing_transcript",
            "session_id": session_id,
            "transcript_path": str(transcript),
        }

    automation_id = find_automation_id(transcript)
    if not automation_id:
        return {"status": "non_automation", "session_id": session_id}

    if codex_home is None:
        codex_home = Path(os.environ.get("CODEX_HOME") or (Path.home() / ".codex"))

    if not is_known_automation(codex_home, automation_id):
        return {
            "status": "unknown_automation",
            "session_id": session_id,
            "automation_id": automation_id,
        }

    if runner is None:
        runner = default_runner

    archive_command = resolve_archive_command(archive_command)

    command = [*archive_command, session_id]
    returncode, stdout, stderr = runner(command)
    if returncode != 0:
        return {
            "status": "archive_failed",
            "session_id": session_id,
            "automation_id": automation_id,
            "command": command,
            "stdout": stdout.strip(),
            "stderr": stderr.strip(),
            "returncode": returncode,
        }

    return {
        "status": "archived",
        "session_id": session_id,
        "automation_id": automation_id,
        "command": command,
    }


def parse_args(argv):
    parser = argparse.ArgumentParser(description="Archive known automation sessions from a Codex Stop hook.")
    parser.add_argument("--codex-home", default=None, help="Override CODEX_HOME for automation lookup.")
    parser.add_argument(
        "--json",
        action="store_true",
        help="Emit diagnostic JSON to stdout for manual runs. Hook mode keeps stdout empty.",
    )
    return parser.parse_args(argv)


def _write_diagnostic(result, *, stdout, stderr, emit_json, success):
    if emit_json:
        print(json.dumps(result, indent=2), file=stdout)
    elif not success:
        print(json.dumps(result, separators=(",", ":")), file=stderr)


def main(argv=None, stdin=None, stdout=None, stderr=None, runner=None):
    args = parse_args(sys.argv[1:] if argv is None else argv)
    stdin = sys.stdin if stdin is None else stdin
    stdout = sys.stdout if stdout is None else stdout
    stderr = sys.stderr if stderr is None else stderr

    try:
        payload = json.load(stdin)
    except json.JSONDecodeError as exc:
        result = {"status": "invalid_hook_payload", "error": str(exc)}
        _write_diagnostic(result, stdout=stdout, stderr=stderr, emit_json=args.json, success=False)
        return 1

    result = archive_session_from_stop_payload(payload, codex_home=args.codex_home, runner=runner)
    success = result["status"] in HOOK_SUCCESS_STATUSES
    _write_diagnostic(result, stdout=stdout, stderr=stderr, emit_json=args.json, success=success)
    return 0 if success else 1


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