"""Lab-only mock services for the Shuffle Full Lab.

Run with: python mock_soar_api.py --host 127.0.0.1 --port 8081
Use only on an isolated lab network. State is held in memory and is cleared
when the process stops.
"""

import argparse
import json
from datetime import datetime, timedelta, timezone
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from urllib.parse import unquote, urlparse


STATE = {"cases": {}, "blocks": {}, "notifications": [], "case_mode": "healthy"}
SAFE_IP = "203.0.113.66"
SAFE_USER = "lab-admin"
MAX_BODY = 64 * 1024
MAX_ENTRIES = 1000


class Handler(BaseHTTPRequestHandler):
    server_version = "CDKShuffleLab/1.0"

    def setup(self):
        super().setup()
        self.connection.settimeout(10)

    def _send(self, status, payload):
        body = json.dumps(payload, indent=2).encode("utf-8")
        self.send_response(status)
        self.send_header("Content-Type", "application/json")
        self.send_header("Content-Length", str(len(body)))
        self.send_header("X-Content-Type-Options", "nosniff")
        self.send_header("Cache-Control", "no-store")
        self.end_headers()
        self.wfile.write(body)

    def _body(self):
        if self.headers.get("Origin") or self.headers.get("Transfer-Encoding"):
            self._send(403, {"error": "browser_and_streamed_requests_not_supported"})
            return None
        if self.headers.get_content_type() != "application/json":
            self._send(415, {"error": "application_json_required"})
            return None
        try:
            lengths = self.headers.get_all("Content-Length", [])
            if len(lengths) != 1 or not lengths[0].isascii() or not lengths[0].isdecimal():
                raise ValueError
            length = int(lengths[0])
        except ValueError:
            self._send(400, {"error": "invalid_content_length"})
            return None
        if length > MAX_BODY:
            self._send(413, {"error": "body_too_large"})
            return None
        try:
            raw = self.rfile.read(length)
            if len(raw) != length:
                raise ValueError
            data = json.loads(raw or b"{}")
            if not isinstance(data, dict):
                raise ValueError
            return data
        except (ValueError, UnicodeDecodeError, RecursionError):
            self._send(400, {"error": "invalid_json"})
            return None
        except TimeoutError:
            self._send(408, {"error": "request_timeout"})
            return None

    def log_message(self, pattern, *args):
        print(f"{self.log_date_time_string()} {pattern % args}")

    def do_GET(self):
        path = unquote(urlparse(self.path).path)
        if path == "/health":
            self._send(200, {"status": "ok", "case_mode": STATE["case_mode"]})
        elif path == f"/identity/{SAFE_USER}":
            self._send(200, {"user": SAFE_USER, "role": "training-admin", "critical": True})
        elif path.startswith("/identity/"):
            self._send(404, {"error": "identity_not_found"})
        elif path.startswith("/cases/"):
            if STATE["case_mode"] == "failed":
                self._send(503, {"error": "case_connector_unavailable"})
                return
            event_id = path.removeprefix("/cases/")
            case = STATE["cases"].get(event_id)
            self._send(200 if case else 404, case or {"error": "case_not_found", "event_id": event_id})
        elif path.startswith("/blocklist/"):
            address = path.removeprefix("/blocklist/")
            block = STATE["blocks"].get(address)
            self._send(200, {"blocked": bool(block), "entry": block})
        elif path == "/state":
            self._send(200, STATE)
        else:
            self._send(404, {"error": "route_not_found"})

    def do_POST(self):
        path = unquote(urlparse(self.path).path)
        data = self._body()
        if data is None:
            return
        if path == "/cases":
            if STATE["case_mode"] in {"failed", "fail_next_write"}:
                if STATE["case_mode"] == "fail_next_write":
                    STATE["case_mode"] = "healthy"
                self._send(503, {"error": "case_connector_unavailable"})
                return
            event_id = data.get("event_id")
            if not isinstance(event_id, str) or not event_id or len(event_id) > 200:
                self._send(400, {"error": "event_id_required"})
                return
            created = event_id not in STATE["cases"]
            if created and len(STATE["cases"]) >= MAX_ENTRIES:
                self._send(409, {"error": "lab_state_full_reset_required"})
                return
            STATE["cases"].setdefault(event_id, {"event_id": event_id, "updates": 0})
            updates = STATE["cases"][event_id]["updates"] + 1
            STATE["cases"][event_id].update(data)
            STATE["cases"][event_id]["updates"] = updates
            self._send(201 if created else 200, {"created": created, "case": STATE["cases"][event_id]})
        elif path == "/notify":
            if len(STATE["notifications"]) >= MAX_ENTRIES:
                self._send(409, {"error": "lab_state_full_reset_required"})
                return
            STATE["notifications"].append(data)
            self._send(202, {"accepted": True, "notification_count": len(STATE["notifications"])})
        elif path == "/blocklist":
            if data.get("source_ip") != SAFE_IP:
                self._send(403, {"error": "lab_address_only", "allowed": SAFE_IP})
                return
            minutes = data.get("expires_minutes", 10)
            if type(minutes) is not int or not 1 <= minutes <= 60:
                self._send(400, {"error": "expires_minutes_must_be_integer_1_to_60"})
                return
            STATE["blocks"][SAFE_IP] = {
                "source_ip": SAFE_IP,
                "expires_at": (datetime.now(timezone.utc) + timedelta(minutes=minutes)).isoformat(),
                "event_id": data.get("event_id"),
            }
            self._send(201, {"blocked": True, "entry": STATE["blocks"][SAFE_IP]})
        elif path == "/control/case-mode":
            mode = data.get("mode")
            if not isinstance(mode, str) or mode not in {"healthy", "failed", "fail_next_write"}:
                self._send(400, {"error": "unsupported_case_mode"})
                return
            STATE["case_mode"] = mode
            self._send(200, {"case_mode": mode})
        elif path == "/reset":
            STATE.update({"cases": {}, "blocks": {}, "notifications": [], "case_mode": "healthy"})
            self._send(200, {"reset": True})
        else:
            self._send(404, {"error": "route_not_found"})

    def do_DELETE(self):
        if self.headers.get("Origin"):
            self._send(403, {"error": "browser_requests_not_supported"})
            return
        path = unquote(urlparse(self.path).path)
        if path.startswith("/blocklist/"):
            address = path.removeprefix("/blocklist/")
            removed = STATE["blocks"].pop(address, None)
            self._send(200, {"removed": bool(removed), "source_ip": address})
        else:
            self._send(404, {"error": "route_not_found"})


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--host", default="127.0.0.1")
    parser.add_argument("--port", type=int, default=8081)
    args = parser.parse_args()
    print(f"Lab API listening on http://{args.host}:{args.port}")
    ThreadingHTTPServer((args.host, args.port), Handler).serve_forever()


if __name__ == "__main__":
    main()
