import json
import os
import base64
import csv
import hashlib
import io
import re
import urllib.parse
import urllib.error
import urllib.parse
import urllib.request
from http.server import SimpleHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path

ROOT = Path(__file__).resolve().parent
CATALOG = ROOT / "data" / "catalog.json"
PROVIDERS = ROOT / "data" / "providers.json"
PROMPT_LIBRARY = ROOT / "data" / "prompt-library.json"
MAX_PROMPT_SOURCE_BYTES = 2 * 1024 * 1024
MAX_IMPORTED_PROMPTS = 300
ADMIN_USER = os.environ.get("ADMIN_USER", "admin")
ADMIN_PASSWORD = os.environ.get("ADMIN_PASSWORD", "")


def github_raw_url(value):
    parsed = urllib.parse.urlparse(str(value or "").strip())
    if parsed.scheme != "https" or parsed.username or parsed.password or parsed.port:
        raise ValueError("只支持 HTTPS GitHub 文件地址")
    parts = [urllib.parse.unquote(part) for part in parsed.path.split("/") if part]
    if parsed.hostname == "raw.githubusercontent.com" and len(parts) >= 4:
        owner, repo, branch = parts[:3]
        path = parts[3:]
    elif parsed.hostname == "github.com" and len(parts) >= 5 and parts[2] in ("blob", "raw"):
        owner, repo, _, branch = parts[:4]
        path = parts[4:]
    else:
        raise ValueError("请提供 raw.githubusercontent.com 文件地址或 github.com 仓库文件页面地址")
    safe_part = re.compile(r"^[A-Za-z0-9_.-]+$")
    if not all(safe_part.fullmatch(part) and part not in (".", "..") for part in [owner, repo, branch]) or not all(part not in (".", "..") and not any(ord(char) < 32 for char in part) for part in path):
        raise ValueError("GitHub 仓库路径包含不支持的字符")
    encoded_path = "/".join(urllib.parse.quote(part, safe="-._~") for part in path)
    return "https://raw.githubusercontent.com/{}/{}/{}/{}".format(owner, repo, branch, encoded_path)


class GitHubRedirectHandler(urllib.request.HTTPRedirectHandler):
    def redirect_request(self, request, response, code, message, headers, new_url):
        host = urllib.parse.urlparse(new_url).hostname or ""
        if host != "raw.githubusercontent.com" and not host.endswith(".githubusercontent.com"):
            raise ValueError("GitHub 文件重定向到非受信任域名")
        return super().redirect_request(request, response, code, message, headers, new_url)


def fetch_github_file(value):
    url = github_raw_url(value)
    request = urllib.request.Request(url, headers={"User-Agent": "GuangzhanPromptImporter/1.0", "Accept": "text/plain, application/json"})
    opener = urllib.request.build_opener(GitHubRedirectHandler)
    try:
        with opener.open(request, timeout=15) as response:
            payload = response.read(MAX_PROMPT_SOURCE_BYTES + 1)
    except (urllib.error.URLError, TimeoutError) as error:
        raise ValueError("GitHub 文件拉取失败，请确认公开访问和网络状态") from error
    if len(payload) > MAX_PROMPT_SOURCE_BYTES:
        raise ValueError("提示词文件超过 2 MB 导入限制")
    return url, payload.decode("utf-8-sig")


def prompt_rows_from_source(source_url, text):
    suffix = urllib.parse.urlparse(source_url).path.rsplit(".", 1)[-1].lower()
    if suffix == "json":
        value = json.loads(text)
        while isinstance(value, dict):
            value = next((value[key] for key in ("prompts", "items", "data", "records", "templates") if key in value), None)
            if value is None:
                raise ValueError("JSON 未找到 prompts/items/data/records/templates 列表")
        if not isinstance(value, list):
            raise ValueError("JSON 提示词数据必须是对象数组")
        return value
    if suffix == "csv":
        return list(csv.DictReader(io.StringIO(text)))
    if suffix in ("md", "markdown"):
        sections = re.split(r"(?m)^#{1,3}\s+", text)
        rows = []
        for section in sections:
            lines = section.strip().splitlines()
            if not lines:
                continue
            title = lines[0].strip().strip("# ")
            content = "\n".join(lines[1:]).strip()
            content = re.sub(r"(?s)^```[^\n]*\n|\n```\s*$", "", content).strip()
            if content:
                rows.append({"title": title, "content": content})
        return rows
    raise ValueError("仅支持 JSON、CSV、MD 或 Markdown 文件")


def normalize_imported_prompts(rows, source_url, license_name, existing):
    existing_by_key = {str(item.get("key")): item for item in existing if item.get("key")}
    result = []
    seen = set()
    updated = 0
    for index, row in enumerate(rows[:MAX_IMPORTED_PROMPTS]):
        if not isinstance(row, dict):
            continue
        title = str(row.get("name") or row.get("title") or row.get("act") or row.get("key") or "").strip()
        content = str(row.get("value") or row.get("prompt") or row.get("content") or row.get("prompt_text") or "").strip()
        if not title or not content:
            continue
        basis = str(row.get("key") or title).lower()
        slug = re.sub(r"[^a-z0-9_]+", "_", basis).strip("_")[:40] or "prompt"
        digest = hashlib.sha1((source_url + "\n" + title).encode("utf-8")).hexdigest()[:8]
        key = "gh_{}_{}".format(slug, digest)
        if key in seen:
            continue
        seen.add(key)
        old = existing_by_key.get(key)
        if old:
            updated += 1
        result.append({
            **(old or {}),
            "id": (old or {}).get("id") or key,
            "key": key,
            "name": title[:120],
            "category": str(row.get("category") or row.get("tag") or "GitHub 导入")[:60],
            "description": str(row.get("description") or row.get("subtitle") or "")[:300],
            "value": content[:20000],
            "theme": (old or {}).get("theme") or "sky",
            "sort_order": (old or {}).get("sort_order") or 10000 + index,
            "is_featured": bool((old or {}).get("is_featured", False)),
            "is_active": bool((old or {}).get("is_active", False)),
            "is_builtin": False,
            "source_label": "GitHub 开源导入",
            "source_url": source_url,
            "source_license": license_name[:160],
        })
    if not result:
        raise ValueError("没有找到同时包含标题和提示词内容的条目")
    return result, updated


class Handler(SimpleHTTPRequestHandler):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, directory=str(ROOT), **kwargs)

    def do_GET(self):
        endpoint = self.path.split("?", 1)[0]
        if ADMIN_PASSWORD and (endpoint.startswith("/admin") or endpoint in ("/api/catalog", "/api/providers")) and not self.authorize_admin():
            return
        if endpoint == "/data/providers.json":
            self.send_error(404)
            return
        if endpoint == "/api/catalog":
            self.send_catalog()
            return
        if endpoint == "/api/providers":
            self.send_providers()
            return
        if endpoint == "/api/prompt-library":
            self.send_prompt_library()
            return
        super().do_GET()

    def do_POST(self):
        endpoint = self.path.split("?", 1)[0]
        if endpoint not in ("/api/catalog", "/api/providers", "/api/prompts/import-github"):
            self.send_error(404)
            return
        if ADMIN_PASSWORD and not self.authorize_admin():
            return
        try:
            length = int(self.headers.get("Content-Length", "0"))
            if length <= 0 or length > 16384:
                raise ValueError("请求内容无效或过大")
            payload = json.loads(self.rfile.read(length).decode("utf-8"))
            if endpoint == "/api/catalog" and (not isinstance(payload, dict) or not isinstance(payload.get("promptTemplates"), list) or not isinstance(payload.get("packages"), list)):
                raise ValueError("invalid catalog")
            if endpoint == "/api/catalog":
                self.save_catalog(payload)
                self.send_response(204)
                self.end_headers()
            elif endpoint == "/api/prompts/import-github":
                source_url = str(payload.get("url") or "").strip()
                license_name = str(payload.get("license") or "").strip()
                if not license_name:
                    raise ValueError("请先填写并核实该仓库的许可证")
                normalized_url, source_text = fetch_github_file(source_url)
                rows = prompt_rows_from_source(normalized_url, source_text)
                if len(rows) > MAX_IMPORTED_PROMPTS:
                    raise ValueError("单次最多导入 300 条，请拆分文件")
                catalog = json.loads(CATALOG.read_text(encoding="utf-8"))
                imported, updated = normalize_imported_prompts(rows, normalized_url, license_name, catalog["promptTemplates"])
                merged = {item.get("key"): item for item in catalog["promptTemplates"]}
                merged.update({item["key"]: item for item in imported})
                catalog["promptTemplates"] = list(merged.values())
                self.save_catalog(catalog)
                self.send_json({"imported": len(imported) - updated, "updated": updated, "pending_review": sum(not item.get("is_active") for item in imported)})
            else:
                self.save_providers(payload)
        except Exception as error:
            body = json.dumps({"message": str(error)}, ensure_ascii=False).encode("utf-8")
            self.send_response(400)
            self.send_header("Content-Type", "application/json; charset=utf-8")
            self.send_header("Content-Length", str(len(body)))
            self.end_headers()
            self.wfile.write(body)

    def send_catalog(self):
        body = CATALOG.read_bytes()
        self.send_response(200)
        self.send_header("Content-Type", "application/json; charset=utf-8")
        self.send_header("Cache-Control", "no-store")
        self.send_header("Content-Length", str(len(body)))
        self.end_headers()
        self.wfile.write(body)

    def save_catalog(self, payload):
        serialized = json.dumps(payload, ensure_ascii=False, indent=2)
        CATALOG.write_text(serialized + "\n", encoding="utf-8")
        module_data = json.dumps(payload, ensure_ascii=False, separators=(",", ":"))
        source = "(function (root) {\n  const catalog = " + module_data + ";\n  if (typeof module === 'object' && module.exports) module.exports = catalog;\n  if (root) root.CatalogSeed = catalog;\n})(typeof window !== 'undefined' ? window : null)\n"
        (ROOT / "data" / "catalog.js").write_text(source, encoding="utf-8")

    def send_json(self, payload, status=200):
        body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
        self.send_response(status)
        self.send_header("Content-Type", "application/json; charset=utf-8")
        self.send_header("Cache-Control", "no-store")
        self.send_header("Content-Length", str(len(body)))
        self.end_headers()
        self.wfile.write(body)

    def send_providers(self):
        if not PROVIDERS.exists():
            PROVIDERS.write_text("[]\n", encoding="utf-8")
        stored = json.loads(PROVIDERS.read_text(encoding="utf-8"))
        public = [{key: value for key, value in item.items() if key != "api_key"} | {"has_api_key": bool(item.get("api_key"))} for item in stored]
        self.send_json({"providers": public})

    def send_prompt_library(self):
        if not PROMPT_LIBRARY.exists():
            self.send_json({"updated_at": None, "total": 0, "items": [], "source_errors": []})
            return
        payload = json.loads(PROMPT_LIBRARY.read_text(encoding="utf-8"))
        query = urllib.parse.parse_qs(urllib.parse.urlsplit(self.path).query)
        page = max(1, int(query.get("page", ["1"])[0]))
        page_size = min(120, max(20, int(query.get("page_size", ["60"])[0])))
        start = (page - 1) * page_size
        end = start + page_size
        self.send_json({
            "updated_at": payload.get("updated_at"),
            "total": payload.get("total", len(payload.get("items", []))),
            "items": payload.get("items", [])[start:end],
            "has_more": end < len(payload.get("items", [])),
            "source_errors": payload.get("source_errors", []),
        })

    def authorize_admin(self):
        header = self.headers.get("Authorization", "")
        if header.startswith("Basic "):
            try:
                decoded = base64.b64decode(header[6:]).decode("utf-8")
                user, password = decoded.split(":", 1)
                if user == ADMIN_USER and password == ADMIN_PASSWORD:
                    return True
            except (ValueError, UnicodeDecodeError):
                pass
        self.send_response(401)
        self.send_header("WWW-Authenticate", 'Basic realm="Guangzhan Admin"')
        self.send_header("Content-Length", "0")
        self.end_headers()
        return False

    def save_providers(self, payload):
        if not isinstance(payload, dict) or not isinstance(payload.get("providers"), list):
            raise ValueError("invalid providers")
        old = json.loads(PROVIDERS.read_text(encoding="utf-8")) if PROVIDERS.exists() else []
        old_by_id = {item.get("id"): item for item in old}
        clean = []
        seen = set()
        default_count = 0
        for item in payload["providers"]:
            if not isinstance(item, dict):
                raise ValueError("invalid provider")
            provider_id = str(item.get("id", "")).strip()
            name = str(item.get("name", "")).strip()
            base_url = str(item.get("base_url", "")).strip()
            protocol = item.get("protocol")
            models = item.get("models")
            if not provider_id or provider_id in seen or not name or not base_url.startswith(("https://", "http://")) or protocol not in ("images", "gemini", "custom") or not isinstance(models, list) or not models:
                raise ValueError("invalid provider fields")
            seen.add(provider_id)
            saved = old_by_id.get(provider_id, {})
            api_key = str(item.get("api_key", "")) or saved.get("api_key", "")
            if not api_key:
                raise ValueError("API Key required")
            is_default = bool(item.get("is_default"))
            default_count += int(is_default)
            clean.append({"id": provider_id, "name": name, "base_url": base_url.rstrip("/"), "protocol": protocol, "models": [str(model) for model in models], "api_key": api_key, "is_default": is_default, "is_active": bool(item.get("is_active")), "updated_at": item.get("updated_at")})
        if default_count > 1:
            raise ValueError("only one default provider is allowed")
        PROVIDERS.write_text(json.dumps(clean, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
        public = [{key: value for key, value in item.items() if key != "api_key"} | {"has_api_key": bool(item.get("api_key"))} for item in clean]
        self.send_json({"providers": public})


if __name__ == "__main__":
    port = int(os.environ.get("PORT", "4173"))
    bind_host = os.environ.get("BIND_HOST", "127.0.0.1")
    ThreadingHTTPServer((bind_host, port), Handler).serve_forever()
