"""Dependency-free client for the owner-scoped MMAT Studio v1 API.

Pass stable idempotency keys explicitly for creates and runs. This client never
retries a request automatically, follows redirects, or includes credentials and
server response bodies in exception messages.
"""
from __future__ import annotations

import json
import math
import re
import time
import urllib.error
import urllib.request
from urllib.parse import quote, urlencode, urlsplit
from uuid import uuid4

__all__ = ["MMATStudio", "APIError", "APIProtocolError", "new_idempotency_key"]

_TERMINAL = frozenset({"completed", "failed", "cancelled", "interrupted"})
_KEY = re.compile(r"^[A-Za-z0-9._:-]{1,200}$")
_MAX_RESPONSE = 8 * 1024 * 1024
_SAFE_CODES = frozenset({"invalid_response", "invalid_request", "authentication_failed", "permission_denied", "not_found", "version_conflict", "idempotency_conflict", "validation_failed", "rate_limited", "internal_error"})


class APIError(Exception):
    """A safe failure message, HTTP status and optional machine error code."""

    def __init__(self, status: int | None = None, *, code: str | None = None):
        self.status = status
        self.code = code if isinstance(code, str) and code in _SAFE_CODES else None
        messages = {
            400: "Invalid API request", 401: "API authentication failed",
            403: "API permission denied", 404: "API resource not found",
            409: "API version or idempotency conflict", 413: "API request is too large",
            422: "API request validation failed", 429: "API concurrency or rate limit reached",
        }
        message = messages.get(status, "API request failed")
        if status is not None:
            message += f" (HTTP {status})"
        super().__init__(message)


class APIProtocolError(APIError):
    """An invalid response or a completed run without a usable output."""

    def __init__(self):
        super().__init__(code="invalid_response")
        self.args = ("API returned an invalid or incomplete response",)


class _NoRedirect(urllib.request.HTTPRedirectHandler):
    def redirect_request(self, req, fp, code, msg, headers, newurl):
        return None


def new_idempotency_key() -> str:
    """Create one key for one logical operation; persist it before sending."""
    return str(uuid4())


def _positive_number(value, name):
    if isinstance(value, bool) or not isinstance(value, (int, float)) or not math.isfinite(value) or value <= 0:
        raise ValueError(f"{name} must be a positive finite number")
    return value


def _version(value):
    if type(value) is not int or value < 1:
        raise ValueError("expected_version must be a positive integer")
    return value


def _part(value):
    if not isinstance(value, str) or not value or value in {".", ".."} or any(ord(c) < 32 for c in value):
        raise ValueError("A valid resource identifier is required")
    return quote(value, safe="")


def _file_path(value):
    if not isinstance(value, str) or len(value) > 160 or "\\" in value:
        raise ValueError("A relative bucket file path is required")
    parts = value.split("/")
    if len(parts) > 5 or any(not p or p.startswith(".") or any(ord(c) < 32 for c in p) for p in parts) or value == "MEMORY.md":
        raise ValueError("A relative bucket file path is required")
    return "/".join(quote(p, safe="") for p in parts)


def _key(value):
    if not isinstance(value, str) or not _KEY.fullmatch(value):
        raise ValueError("idempotency_key must contain 1–200 letters, digits or . _ : -")
    return value


class MMATStudio:
    """The base URL includes /api/v1; use HTTPS outside localhost."""

    def __init__(self, base_url: str, token: str, *, timeout: float = 30):
        try:
            parsed = urlsplit(base_url)
            valid = (parsed.scheme in {"https", "http"} and parsed.hostname
                     and not parsed.username and not parsed.password
                     and not parsed.query and not parsed.fragment
                     and parsed.path.rstrip("/").endswith("/api/v1")
                     and (parsed.scheme == "https" or parsed.hostname in {"localhost", "127.0.0.1", "::1"}))
            parsed.port
        except (TypeError, ValueError):
            valid = False
        if not valid:
            raise ValueError("base_url must be an HTTPS API v1 URL without credentials or query parameters")
        if not isinstance(token, str) or not token or any(c.isspace() or not 33 <= ord(c) <= 126 for c in token):
            raise ValueError("A nonempty API token is required")
        self._base_url = base_url.rstrip("/")
        self._token = token
        self._timeout = _positive_number(timeout, "timeout")
        self._opener = urllib.request.build_opener(_NoRedirect())

    def __repr__(self):
        return "<MMATStudio client>"

    def _request(self, method, path, body=None, *, idempotency_key=None, query=None, text=False, binary=False, max_bytes=_MAX_RESPONSE):
        if not isinstance(path, str) or not path.startswith("/") or path.startswith("//") or "?" in path or "#" in path:
            raise ValueError("A relative API endpoint is required")
        url = self._base_url + path
        if query:
            url += "?" + urlencode({k: v for k, v in query.items() if v is not None})
        headers = {"Authorization": "Bearer " + self._token, "Accept": "application/json"}
        if idempotency_key is not None:
            headers["Idempotency-Key"] = _key(idempotency_key)
        data = None
        if body is not None:
            try:
                data = json.dumps(body, ensure_ascii=False, allow_nan=False).encode("utf-8")
            except (TypeError, ValueError):
                raise ValueError("The request body must contain valid JSON data") from None
            headers["Content-Type"] = "application/json"
        request = urllib.request.Request(url, data=data, headers=headers, method=method)
        try:
            with self._opener.open(request, timeout=self._timeout) as response:
                raw = response.read(max_bytes + 1)
                if len(raw) > max_bytes:
                    raise APIProtocolError()
        except urllib.error.HTTPError as exc:
            code = None
            try:
                error = json.loads(exc.read(16384))
                if isinstance(error, dict):
                    code = error.get("code")
            except (ValueError, OSError):
                pass
            finally:
                exc.close()
            raise APIError(exc.code, code=code) from None
        except (urllib.error.URLError, OSError, TimeoutError):
            raise APIError() from None
        if binary:
            return raw
        try:
            if text:
                return raw.decode("utf-8")
            if not raw:
                return None
            result = json.loads(raw)
        except (ValueError, UnicodeError):
            raise APIProtocolError() from None
        if not isinstance(result, dict):
            raise APIProtocolError()
        return result

    @staticmethod
    def _unwrap(response, key):
        if not isinstance(response, dict) or key not in response:
            raise APIProtocolError()
        return response[key]

    def health(self):
        return self._request("GET", "/health")

    def catalog(self):
        return self._request("GET", "/catalog")

    def me(self):
        return self._request("GET", "/me")

    def list_buckets(self):
        return self._unwrap(self._request("GET", "/buckets"), "buckets")

    def create_bucket(self, name, *, idempotency_key, description="", profile=None, files=None, memory=""):
        body = {"name": name, "description": description, "files": files if files is not None else [], "memory": memory}
        if profile is not None:
            body["profile"] = profile
        return self._unwrap(self._request("POST", "/buckets", body, idempotency_key=_key(idempotency_key)), "bucket")

    def get_bucket(self, bucket_id, *, version=None):
        return self._unwrap(self._request("GET", "/buckets/" + _part(bucket_id), query={"version": version}), "bucket")

    def update_bucket(self, bucket_id, *, expected_version, **changes):
        return self._unwrap(self._request("PATCH", "/buckets/" + _part(bucket_id), {**changes, "expected_version": _version(expected_version)}), "bucket")

    def delete_bucket(self, bucket_id):
        return self._unwrap(self._request("DELETE", "/buckets/" + _part(bucket_id)), "bucket")

    def bucket_action(self, bucket_id, action, *, idempotency_key):
        return self._unwrap(self._request("POST", "/buckets/" + _part(bucket_id) + "/actions", {"action": action}, idempotency_key=_key(idempotency_key)), "bucket")

    def bucket_versions(self, bucket_id):
        return self._unwrap(self._request("GET", "/buckets/" + _part(bucket_id) + "/versions"), "versions")

    def get_file(self, bucket_id, path):
        return self._request("GET", "/buckets/" + _part(bucket_id) + "/files/" + _file_path(path))

    def put_file(self, bucket_id, path, content, *, expected_version):
        return self._request("PUT", "/buckets/" + _part(bucket_id) + "/files/" + _file_path(path), {"content": content, "expected_version": _version(expected_version)})

    def delete_file(self, bucket_id, path, *, expected_version):
        return self._request("DELETE", "/buckets/" + _part(bucket_id) + "/files/" + _file_path(path), query={"expected_version": _version(expected_version)})

    def get_memory(self, bucket_id):
        return self._request("GET", "/buckets/" + _part(bucket_id) + "/memory")

    def put_memory(self, bucket_id, content, *, expected_version):
        return self._request("PUT", "/buckets/" + _part(bucket_id) + "/memory", {"content": content, "expected_version": _version(expected_version)})

    def run_bucket(self, bucket_id, prompt, *, idempotency_key, input_text="", mode="read_write", expected_version=None):
        body = {"prompt": prompt, "input_text": input_text, "mode": mode}
        if expected_version is not None:
            body["expected_version"] = _version(expected_version)
        return self._unwrap(self._request("POST", "/buckets/" + _part(bucket_id) + "/runs", body, idempotency_key=_key(idempotency_key)), "run")

    def bucket_runs(self, bucket_id):
        return self._unwrap(self._request("GET", "/buckets/" + _part(bucket_id) + "/runs"), "runs")

    def list_systems(self):
        return self._unwrap(self._request("GET", "/systems"), "systems")

    def create_system(self, name, *, idempotency_key, description="", color="#dc7359", template="empty"):
        body = {"name": name, "description": description, "color": color, "template": template}
        return self._unwrap(self._request("POST", "/systems", body, idempotency_key=_key(idempotency_key)), "system")

    def get_system(self, system_id):
        return self._unwrap(self._request("GET", "/systems/" + _part(system_id)), "system")

    def update_system(self, system_id, *, expected_version, **changes):
        return self._unwrap(self._request("PATCH", "/systems/" + _part(system_id), {**changes, "expected_version": _version(expected_version)}), "system")

    def delete_system(self, system_id):
        return self._unwrap(self._request("DELETE", "/systems/" + _part(system_id)), "system")

    def system_action(self, system_id, action, *, idempotency_key):
        return self._unwrap(self._request("POST", "/systems/" + _part(system_id) + "/actions", {"action": action}, idempotency_key=_key(idempotency_key)), "system")

    def get_graph(self, system_id):
        return self._request("GET", "/systems/" + _part(system_id) + "/graph")

    def put_graph(self, system_id, graph, *, expected_version):
        return self._request("PUT", "/systems/" + _part(system_id) + "/graph", {"graph": graph, "expected_version": _version(expected_version)})

    def get_nodes(self, system_id):
        return self._request("GET", "/systems/" + _part(system_id) + "/nodes")

    def put_nodes(self, system_id, nodes, edges, *, expected_version):
        return self._request("PUT", "/systems/" + _part(system_id) + "/nodes", {"nodes": nodes, "edges": edges, "expected_version": _version(expected_version)})

    def run_system(self, system_id, *, idempotency_key, input_text="", node_id=None):
        body = {"input_text": input_text}
        if node_id is not None:
            body["node_id"] = node_id
        return self._unwrap(self._request("POST", "/systems/" + _part(system_id) + "/runs", body, idempotency_key=_key(idempotency_key)), "run")

    def system_runs(self, system_id):
        return self._unwrap(self._request("GET", "/systems/" + _part(system_id) + "/runs"), "runs")

    def get_run(self, run_id):
        return self._unwrap(self._request("GET", "/runs/" + _part(run_id)), "run")

    def stop_run(self, run_id, *, wait=False, timeout=300, poll_interval=1, idempotency_key=None):
        run = self._unwrap(self._request("POST", "/runs/" + _part(run_id) + "/stop", {}, idempotency_key=idempotency_key), "run")
        return self.wait_run(run_id, timeout=timeout, poll_interval=poll_interval) if wait else run

    def get_run_node(self, run_id, node_id):
        return self._unwrap(self._request("GET", "/runs/" + _part(run_id) + "/nodes/" + _part(node_id)), "node")

    def run_events(self, run_id, *, after_seq=0, limit=100):
        return self._request("GET", "/runs/" + _part(run_id) + "/events", query={"after_seq": after_seq, "limit": limit})

    def wait_run(self, run_id, *, timeout=300, poll_interval=1):
        """Return a terminal run; a timeout leaves the server run untouched."""
        deadline = time.monotonic() + _positive_number(timeout, "timeout")
        interval = _positive_number(poll_interval, "poll_interval")
        while True:
            run = self.get_run(run_id)
            if not isinstance(run, dict):
                raise APIProtocolError()
            if run.get("status") in _TERMINAL:
                if run["status"] == "completed":
                    output = run.get("output")
                    if output is None and isinstance(run.get("result"), dict):
                        output = run["result"].get("output")
                    if run.get("ready") is not True or not isinstance(output, str) or not output.strip():
                        raise APIProtocolError()
                return run
            remaining = deadline - time.monotonic()
            if remaining <= 0:
                raise TimeoutError("Timed out waiting for the API run; the server run was not stopped")
            time.sleep(min(interval, remaining))

    def list_skills(self):
        return self._unwrap(self._request("GET", "/skills"), "skills")

    def create_skill(self, name, description, instructions, *, idempotency_key, **fields):
        body = {**fields, "name": name, "description": description, "instructions": instructions}
        return self._unwrap(self._request("POST", "/skills", body, idempotency_key=_key(idempotency_key)), "skill")

    def get_skill(self, skill_id, *, version=None):
        return self._unwrap(self._request("GET", "/skills/" + _part(skill_id), query={"version": version}), "skill")

    def update_skill(self, skill_id, *, expected_version, **changes):
        return self._unwrap(self._request("PATCH", "/skills/" + _part(skill_id), {**changes, "expected_version": _version(expected_version)}), "skill")

    def delete_skill(self, skill_id):
        return self._unwrap(self._request("DELETE", "/skills/" + _part(skill_id)), "skill")

    def skill_versions(self, skill_id):
        return self._unwrap(self._request("GET", "/skills/" + _part(skill_id) + "/versions"), "versions")

    def export_skill(self, skill_id, *, version=None):
        return self._request("GET", "/skills/" + _part(skill_id) + "/export", query={"version": version}, text=True)

    def skill_action(self, skill_id, action, *, idempotency_key):
        return self._unwrap(self._request("POST", "/skills/" + _part(skill_id) + "/actions", {"action": action}, idempotency_key=_key(idempotency_key)), "skill")

    def download_artifact(self, url, *, max_bytes=128 * 1024 * 1024):
        """Download an artifact URL returned by the API without forwarding auth."""
        if type(max_bytes) is not int or not 1 <= max_bytes <= 128 * 1024 * 1024:
            raise ValueError("max_bytes must be between 1 and 134217728")
        if not isinstance(url, str):
            raise ValueError("An API artifact URL is required")
        parsed, base = urlsplit(url), urlsplit(self._base_url)
        if parsed.query or parsed.fragment or parsed.username or parsed.password:
            raise ValueError("An API artifact URL is required")
        if parsed.scheme or parsed.netloc:
            if (parsed.scheme, parsed.netloc) != (base.scheme, base.netloc):
                raise ValueError("Artifact URL must use the same API origin")
        prefix = base.path.rstrip("/")
        relative = parsed.path[len(prefix):] if parsed.path.startswith(prefix + "/") else ""
        if not re.fullmatch(r"/runs/[A-Za-z0-9_-]+(?:/nodes/[A-Za-z0-9_-]+)?/artifacts/[0-9]+", relative):
            raise ValueError("Artifact URL must be a v1 run artifact endpoint")
        return self._request("GET", relative, binary=True, max_bytes=max_bytes)
