"""Stress harness for the AI Intake path (milestone 8).

Drives the REAL CmsClient against the REAL staging CMS over signed HTTP. No
stubs: the point is to exercise the seam that unit tests on both sides missed.

Usage:
    python stress.py register <n> [<tag>]
    python stress.py lifecycle <n> [<tag>]
    python stress.py duplicates <n>
    python stress.py audit <tag>
"""

from __future__ import annotations

import base64
import os
import sys
import time
import uuid
from pathlib import Path

sys.path.insert(0, r"C:\taxpilot-ai-work")

ENV = Path(r"C:\taxpilot-ai-staging\shared\.env")

for line in ENV.read_text(encoding="utf-8").splitlines():
    line = line.strip()
    if not line or line.startswith("#") or "=" not in line:
        continue
    key, value = line.split("=", 1)
    os.environ.setdefault(key.strip(), value.strip().strip('"').strip("'"))

from app.api.client import CmsClient  # noqa: E402
from app.config.settings import Settings  # noqa: E402


def client() -> CmsClient:
    return CmsClient(Settings.from_env())


PDF = b"%PDF-1.4\n% stress harness\n1 0 obj<</Type/Catalog>>endobj\ntrailer<</Root 1 0 R>>"


def _rss_mb() -> float:
    """Resident memory of this process, in MB."""
    try:
        import psutil

        return round(psutil.Process().memory_info().rss / 1048576, 1)
    except Exception:
        return 0.0


def register(n: int, tag: str) -> dict:
    cms = client()
    started = time.monotonic()
    rss_before = _rss_mb()
    refs, failures = [], []

    for i in range(n):
        body = {
            "source": "whatsapp",
            "source_account": "923049637232",
            "provider_message_id": f"{tag}-{i:05d}",
            "original_filename": f"{tag}-{i:05d}.pdf",
            "file_base64": base64.b64encode(PDF + str(i).encode()).decode(),
        }
        try:
            refs.append(cms.register_intake(body)["reference"])
        except Exception as error:  # noqa: BLE001
            failures.append(f"{i}: {error}")

    elapsed = time.monotonic() - started

    return {
        "sent": n,
        "rss_start_mb": rss_before,
        "rss_end_mb": _rss_mb(),
        "registered": len(refs),
        "unique": len(set(refs)),
        "failures": failures[:5],
        "failed": len(failures),
        "seconds": round(elapsed, 1),
        "per_second": round(n / elapsed, 1) if elapsed else 0,
        "refs": refs,
    }


def lifecycle(n: int, tag: str) -> dict:
    """Register, then walk each document through to a terminal state."""
    result = register(n, tag)
    cms = client()
    started = time.monotonic()
    moved, stuck = 0, []

    for reference in result["refs"]:
        try:
            cms.update_intake(reference, status="processing", current_step="read")
            cms.update_intake(reference, status="approval", reason="stress")
            cms.update_intake(reference, status="completed", reason="stress")
            moved += 1
        except Exception as error:  # noqa: BLE001
            stuck.append(f"{reference}: {error}")

    result["transitions_seconds"] = round(time.monotonic() - started, 1)
    result["completed"] = moved
    result["stuck"] = stuck[:5]
    result.pop("refs", None)

    return result


def duplicates(n: int) -> dict:
    """The same message id, n times. Expect one record."""
    cms = client()
    message_id = f"DUP-{uuid.uuid4().hex[:8].upper()}"
    refs = []

    for _ in range(n):
        refs.append(cms.register_intake({
            "source": "whatsapp",
            "provider_message_id": message_id,
            "original_filename": "duplicate.pdf",
            "file_base64": base64.b64encode(PDF).decode(),
        })["reference"])

    return {"sent": n, "distinct_records": len(set(refs)), "reference": refs[0]}


def audit(tag: str) -> dict:
    cms = client()
    pending = cms.pending_intake()

    return {"cms_reports_unfinished": len(pending)}


if __name__ == "__main__":
    command = sys.argv[1]
    import json

    if command == "register":
        out = register(int(sys.argv[2]), sys.argv[3] if len(sys.argv) > 3 else "S")
        out.pop("refs", None)
    elif command == "lifecycle":
        out = lifecycle(int(sys.argv[2]), sys.argv[3] if len(sys.argv) > 3 else "L")
    elif command == "duplicates":
        out = duplicates(int(sys.argv[2]))
    else:
        out = audit(sys.argv[2] if len(sys.argv) > 2 else "")

    print(json.dumps(out, indent=2))
