#!/usr/bin/env python3
"""dloffload receiver.

Accepts download tasks over the LAN from the Android side and fetches them
locally, so the phone never has to pull the bytes.

  POST /api/tasks   {url, headers?, filename?, dest_dir?, proxy?, engine?}  -> {task_id}
  GET  /api/tasks/<id>                                                     -> task state
  GET  /api/tasks?limit=N                                                  -> recent tasks
  GET  /api/health                                                         -> ok + disk + queue

Auth: Authorization: Bearer <token>   (token lives in the config file, 0600)

Only stdlib is used on purpose: the host is a small Linux box.
"""

from __future__ import annotations

import hmac
import json
import mimetypes
import os
import queue
import re
import secrets
import shutil
import sqlite3
import subprocess
import sys
import threading
import time
import uuid
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from urllib.parse import urlparse, unquote, quote, parse_qs

VERSION = "0.1.0"
DEFAULT_CONFIG = Path.home() / ".config" / "dloffload" / "recv.json"
SAFE_NAME = re.compile(r"[^\w.\-()\[\] ]+", re.UNICODE)


# --------------------------------------------------------------------------- config


def load_config(path: Path) -> dict:
    if not path.exists():
        cfg = {
            "token": secrets.token_urlsafe(24),
            "bind": "0.0.0.0",
            "port": 8760,
            "root": "/mnt/DataDisk/待处理",
            "default_dest_dir": "video",
            "default_proxy": "http://127.0.0.1:7897",
            "db": str(Path.home() / ".local" / "share" / "dloffload" / "tasks.db"),
            "workers": 1,
            "connect_timeout": 20,
            # Ranges fetched concurrently per task. The proxy in front of this
            # box throttles a single connection hard (230 KB/s measured vs
            # 1120 KB/s over 8), so >1 pays off on any large file.
            "connections": 8,
            "parallel_min_bytes": 8 << 20,
        }
        path.parent.mkdir(parents=True, exist_ok=True)
        path.write_text(json.dumps(cfg, indent=2, ensure_ascii=False))
        os.chmod(path, 0o600)
        print(f"[recv] created config {path}", flush=True)
    cfg = json.loads(path.read_text())
    had = set(cfg)
    cfg.setdefault("default_dest_dir", "video")
    cfg.setdefault("workers", 1)
    cfg.setdefault("connect_timeout", 20)
    cfg.setdefault("connections", 8)
    cfg.setdefault("parallel_min_bytes", 8 << 20)
    if set(cfg) != had:
        # Persist defaults added by a newer version: the file is supposed to be
        # the place a human looks to see (and tune) every knob.
        path.write_text(json.dumps(cfg, indent=2, ensure_ascii=False))
        os.chmod(path, 0o600)
    return cfg


# --------------------------------------------------------------------------- store

SCHEMA = """
CREATE TABLE IF NOT EXISTS tasks (
  id TEXT PRIMARY KEY,
  url TEXT NOT NULL,
  headers TEXT NOT NULL DEFAULT '{}',
  dest_dir TEXT,
  filename TEXT,
  proxy TEXT,
  engine TEXT,
  state TEXT NOT NULL,
  bytes INTEGER DEFAULT 0,
  total INTEGER DEFAULT 0,
  path TEXT,
  error TEXT,
  created_at REAL,
  updated_at REAL
);
"""


class Store:
    def __init__(self, db_path: str):
        Path(db_path).parent.mkdir(parents=True, exist_ok=True)
        self.db_path = db_path
        self.lock = threading.Lock()
        with self._conn() as c:
            c.executescript(SCHEMA)
        # anything left "running" from a previous boot is unknown state
        self.update_where("state IN ('queued','running')", state="failed", error="receiver restarted")

    def _conn(self) -> sqlite3.Connection:
        c = sqlite3.connect(self.db_path, timeout=30)
        c.row_factory = sqlite3.Row
        return c

    def insert(self, task: dict) -> None:
        now = time.time()
        with self.lock, self._conn() as c:
            c.execute(
                "INSERT INTO tasks (id,url,headers,dest_dir,filename,proxy,engine,state,created_at,updated_at)"
                " VALUES (?,?,?,?,?,?,?,'queued',?,?)",
                (
                    task["id"], task["url"], json.dumps(task.get("headers") or {}),
                    task.get("dest_dir"), task.get("filename"), task.get("proxy"),
                    task.get("engine"), now, now,
                ),
            )

    def update(self, tid: str, **fields) -> None:
        if not fields:
            return
        fields["updated_at"] = time.time()
        cols = ", ".join(f"{k}=?" for k in fields)
        with self.lock, self._conn() as c:
            c.execute(f"UPDATE tasks SET {cols} WHERE id=?", (*fields.values(), tid))

    def update_where(self, where: str, **fields) -> None:
        cols = ", ".join(f"{k}=?" for k in fields)
        with self.lock, self._conn() as c:
            c.execute(f"UPDATE tasks SET {cols} WHERE {where}", tuple(fields.values()))

    def get(self, tid: str) -> dict | None:
        with self._conn() as c:
            row = c.execute("SELECT * FROM tasks WHERE id=?", (tid,)).fetchone()
        return dict(row) if row else None

    def recent(self, limit: int = 20) -> list[dict]:
        with self._conn() as c:
            rows = c.execute("SELECT * FROM tasks ORDER BY created_at DESC LIMIT ?", (limit,)).fetchall()
        return [dict(r) for r in rows]


# --------------------------------------------------------------------------- paths


class PathError(Exception):
    pass


def resolve_dest(cfg: dict, dest_dir: str | None, filename: str) -> Path:
    root = Path(cfg["root"]).resolve()
    rel = dest_dir or cfg.get("default_dest_dir") or ""
    # accept both "video" and an absolute path that lives under root
    candidate = Path(rel)
    if candidate.is_absolute():
        target = candidate
    else:
        target = root / candidate
    target = Path(os.path.normpath(str(target)))
    # resolve parents that exist, then re-check containment (defeats ../ traversal)
    real_root = root
    probe = target
    while not probe.exists() and probe != probe.parent:
        probe = probe.parent
    if probe.exists() and not str(probe.resolve()).startswith(str(real_root)):
        raise PathError("dest escapes root")
    if not str(target).startswith(str(real_root)):
        raise PathError("dest escapes root")
    target.mkdir(parents=True, exist_ok=True)
    if not str(target.resolve()).startswith(str(real_root)):
        raise PathError("dest escapes root (symlink)")
    return target / filename


def safe_filename(url: str, explicit: str | None, ctype: str = "") -> str:
    if explicit:
        name = explicit
    else:
        name = unquote(os.path.basename(urlparse(url).path)) or ""
    name = name.strip().strip(".")
    if not name:
        ext = mimetypes.guess_extension((ctype or "").split(";")[0].strip()) or ""
        name = f"{uuid.uuid4().hex[:16]}{ext}"
    name = SAFE_NAME.sub("_", name)[:180]
    return name or uuid.uuid4().hex


def unique_path(path: Path) -> Path:
    if not path.exists():
        return path
    stem, suffix, i = path.stem, path.suffix, 1
    while True:
        cand = path.with_name(f"{stem}_{i}{suffix}")
        if not cand.exists():
            return cand
        i += 1


# --------------------------------------------------------------------------- download


def curl_head(url: str, proxy: str | None, headers: dict, timeout: int) -> tuple[int, str]:
    """Return (content_length, content_type) — best effort."""
    cmd = ["curl", "-sIL", "--fail", "--max-time", str(timeout)]
    if proxy:
        cmd += ["--proxy", proxy]
    for k, v in headers.items():
        cmd += ["-H", f"{k}: {v}"]
    cmd.append(url)
    try:
        out = subprocess.run(cmd, capture_output=True, text=True, timeout=timeout + 10).stdout
    except Exception:
        return 0, ""
    total, ctype = 0, ""
    for line in out.splitlines():
        low = line.lower()
        if low.startswith("content-length:"):
            try:
                total = int(line.split(":", 1)[1].strip())
            except ValueError:
                pass
        elif low.startswith("content-type:"):
            ctype = line.split(":", 1)[1].strip()
    return total, ctype


def curl_base(cfg: dict, proxy: str | None, headers: dict) -> list[str]:
    """Shared curl flags: retries, proxy and the headers we forward from the phone."""
    cmd = [
        "curl", "-L", "--fail", "--retry", "2", "--retry-delay", "2",
        "--connect-timeout", str(cfg["connect_timeout"]), "--no-progress-meter",
    ]
    if proxy:
        cmd += ["--proxy", proxy]
    for k, v in headers.items():
        if k.lower() in ("host", "content-length"):
            continue
        cmd += ["-H", f"{k}: {v}"]
    return cmd


def run_curl(url: str, part: Path, cfg: dict, proxy: str | None, headers: dict) -> None:
    cmd = curl_base(cfg, proxy, headers) + ["--continue-at", "-", "--output", str(part), url]
    proc = subprocess.run(cmd, capture_output=True, text=True)
    if proc.returncode != 0:
        raise RuntimeError(f"curl rc={proc.returncode} {proc.stderr.strip()[:300]}")


# ------------------------------------------------------------------ parallel download
#
# The proxy in front of the receiver caps a single connection hard: measured
# 2026-09-29 on a real CDN, one stream = 230 KB/s while 8 range requests over the
# same proxy = 1120 KB/s (4.9x). Most CDNs advertise accept-ranges, so split the
# file into contiguous ranges, fetch them concurrently and concatenate. If the
# probe says ranges are unsupported - or anything about the chunked path fails -
# fall back to the plain single-stream download.


class Progress:
    """Bytes on disk for the running task: the part file, or the chunk files."""

    def __init__(self, part: Path) -> None:
        self.paths = [part]

    def bytes(self) -> int:
        total = 0
        for p in self.paths:
            try:
                total += p.stat().st_size
            except OSError:
                pass
        return total


def parse_content_range(value: str) -> int | None:
    """'bytes 0-0/106906352' -> 106906352 ('*' or a malformed value -> None)."""
    m = re.match(r"bytes\s+\d+-\d+/(\d+)", value.strip(), re.I)
    return int(m.group(1)) if m else None


def plan_ranges(total: int, conns: int) -> list[tuple[int, int]]:
    """Split [0, total) into at most `conns` inclusive ranges, in order.

    Every byte is covered exactly once, so the concatenation is the file. Short
    files get fewer chunks than conns rather than useless 0-byte requests.
    """
    if total <= 0 or conns < 1:
        return []
    n = min(conns, total)
    size, rem = divmod(total, n)
    out, start = [], 0
    for i in range(n):
        length = size + (1 if i < rem else 0)
        out.append((start, start + length - 1))
        start += length
    return out


def probe_ranges(url: str, proxy: str | None, headers: dict, cfg: dict) -> tuple[int, bool]:
    """Ask for one byte. Returns (total_bytes, server_honours_ranges)."""
    cmd = curl_base(cfg, proxy, headers) + ["-D", "-", "-o", os.devnull, "-r", "0-0", url]
    try:
        proc = subprocess.run(cmd, capture_output=True, text=True, timeout=cfg["connect_timeout"] + 20)
    except subprocess.TimeoutExpired:
        return 0, False
    if proc.returncode != 0:
        return 0, False
    total = None
    for line in (proc.stdout or "").splitlines():
        if line.lower().startswith("content-range:"):
            total = parse_content_range(line.split(":", 1)[1])
    # A server that ignores the range answers 200 with the whole file (no
    # content-range); honouring it answers 206 with one.
    if total is None or "206" not in (proc.stdout or "").split("\r\n")[0]:
        return total or 0, False
    return total, True


def run_curl_parallel(
    url: str, part: Path, cfg: dict, proxy: str | None, headers: dict,
    total: int, conns: int, progress: Progress,
) -> None:
    """Fetch total bytes over `conns` ranges, then assemble them into `part`."""
    ranges = plan_ranges(total, conns)
    if not ranges:
        raise RuntimeError("nothing to download")
    # Chunk names keep the .part suffix so the media watcher on the receiving
    # host still treats them as in-progress files.
    chunks = [part.with_name(f"{part.name}.{i}.tmp.part") for i in range(len(ranges))]
    for c in chunks:
        try:
            c.unlink()
        except OSError:
            pass
    progress.paths = list(chunks)
    procs: list[tuple[subprocess.Popen, Path, tuple[int, int]]] = []
    try:
        for (start, end), out in zip(ranges, chunks):
            cmd = curl_base(cfg, proxy, headers) + ["-r", f"{start}-{end}", "--output", str(out), url]
            procs.append((subprocess.Popen(cmd, stdout=subprocess.DEVNULL,
                                           stderr=subprocess.PIPE, text=True), out, (start, end)))
        errors = []
        for proc, out, (start, end) in procs:
            _, err = proc.communicate()
            if proc.returncode != 0:
                errors.append(f"[{start}-{end}] rc={proc.returncode} {(err or '').strip()[:160]}")
        if errors:
            raise RuntimeError("chunked download failed: " + "; ".join(errors[:3]))
        for out, (start, end) in zip(chunks, ranges):
            want = end - start + 1
            got = out.stat().st_size
            if got != want:
                raise RuntimeError(f"chunk {start}-{end}: {got}/{want} bytes")
        with open(part, "wb") as fh:
            for c in chunks:
                with open(c, "rb") as src:
                    shutil.copyfileobj(src, fh, 1 << 20)
    finally:
        for c in chunks:
            try:
                c.unlink()
            except OSError:
                pass
        progress.paths = [part]


def run_ytdlp(url: str, dest: Path, cfg: dict, proxy: str | None, headers: dict) -> Path:
    out_tpl = str(dest.with_suffix("")) + ".%(ext)s"
    cmd = ["yt-dlp", "--no-playlist", "--no-progress", "-o", out_tpl, "--print", "after_move:filepath"]
    if proxy:
        cmd += ["--proxy", proxy]
    referer = headers.get("Referer") or headers.get("referer")
    if referer:
        cmd += ["--referer", referer]
    ua = headers.get("User-Agent") or headers.get("user-agent")
    if ua:
        cmd += ["--user-agent", ua]
    cookie = headers.get("Cookie") or headers.get("cookie")
    if cookie:
        cmd += ["--add-header", f"Cookie:{cookie}"]
    cmd.append(url)
    proc = subprocess.run(cmd, capture_output=True, text=True)
    if proc.returncode != 0:
        raise RuntimeError(f"yt-dlp rc={proc.returncode} {proc.stderr.strip()[:300]}")
    printed = (proc.stdout or "").strip().splitlines()
    return Path(printed[-1]) if printed else dest


def worker_loop(cfg: dict, store: Store, tasks: "queue.Queue[str]") -> None:
    while True:
        tid = tasks.get()
        task = store.get(tid)
        if not task:
            continue
        try:
            _run_task(cfg, store, task)
        except Exception as exc:  # never let the worker die
            store.update(tid, state="failed", error=str(exc)[:500])
            print(f"[recv] task {tid} failed: {exc}", flush=True)


def _run_task(cfg: dict, store: Store, task: dict) -> None:
    tid = task["id"]
    url = task["url"]
    headers = json.loads(task["headers"] or "{}")
    proxy = task.get("proxy") or cfg.get("default_proxy") or None
    store.update(tid, state="running", error=None)
    print(f"[recv] task {tid} start url={url} proxy={proxy}", flush=True)

    total, ctype = curl_head(url, proxy, headers, cfg["connect_timeout"])
    name = safe_filename(url, task.get("filename"), ctype)
    dest = unique_path(resolve_dest(cfg, task.get("dest_dir"), name))
    part = dest.with_name(dest.name + ".part")
    store.update(tid, path=str(dest), total=total)

    engine = (task.get("engine") or "curl").lower()
    stop = threading.Event()
    progress = Progress(part)

    def ticker() -> None:
        while not stop.is_set():
            try:
                store.update(tid, bytes=progress.bytes())
            except OSError:
                pass
            stop.wait(2)

    t = threading.Thread(target=ticker, daemon=True)
    t.start()
    try:
        if engine == "ytdlp":
            produced = run_ytdlp(url, dest, cfg, proxy, headers)
            store.update(tid, path=str(produced), bytes=produced.stat().st_size if produced.exists() else 0)
        else:
            conns = int(task.get("conn") or cfg.get("connections") or 1)
            min_parallel = int(cfg.get("parallel_min_bytes") or 0)
            parallel = False
            if conns > 1 and min_parallel > 0 and (total == 0 or total >= min_parallel):
                probe_total, ranges_ok = probe_ranges(url, proxy, headers, cfg)
                if probe_total:
                    total = probe_total
                parallel = ranges_ok and total >= min_parallel
            if parallel:
                print(f"[recv] task {tid} parallel: {conns} ranges, {total} bytes", flush=True)
                try:
                    run_curl_parallel(url, part, cfg, proxy, headers, total, conns, progress)
                except Exception as exc:
                    # Ranges are an optimisation, never a requirement: drop back
                    # to one stream rather than failing the download.
                    print(f"[recv] task {tid} parallel failed ({exc}); single stream", flush=True)
                    progress.paths = [part]
                    try:
                        part.unlink()
                    except OSError:
                        pass
                    run_curl(url, part, cfg, proxy, headers)
            else:
                run_curl(url, part, cfg, proxy, headers)
            size = part.stat().st_size
            if total and size < total:
                raise RuntimeError(f"incomplete: {size}/{total} bytes")
            if total and size > total:
                raise RuntimeError(f"oversized: {size}/{total} bytes")
            part.replace(dest)
            store.update(tid, bytes=size, path=str(dest))
    finally:
        stop.set()
    store.update(tid, state="done", error=None)
    print(f"[recv] task {tid} done -> {dest}", flush=True)


# --------------------------------------------------------------------------- http


class Handler(BaseHTTPRequestHandler):
    server_version = f"dloffload-recv/{VERSION}"
    cfg: dict
    store: Store
    tasks: "queue.Queue[str]"

    def log_message(self, fmt: str, *args) -> None:  # quieter default logging
        if os.environ.get("DLOFFLOAD_VERBOSE"):
            super().log_message(fmt, *args)

    # -- helpers
    def _json(self, code: int, payload: dict) -> None:
        body = json.dumps(payload, ensure_ascii=False).encode()
        self.send_response(code)
        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 _authed(self) -> bool:
        auth = self.headers.get("Authorization", "")
        want = f"Bearer {self.cfg['token']}"
        return hmac.compare_digest(auth, want)

    def _body(self) -> dict:
        n = int(self.headers.get("Content-Length") or 0)
        if n <= 0:
            return {}
        return json.loads(self.rfile.read(n).decode("utf-8"))

    def _drain(self, limit: int = 8 << 20) -> None:
        """Consume the declared request body before answering an error.

        Answering a large upload without reading it leaves the client writing into
        a closed socket; behind a proxy that surfaces as a bogus 502 instead of
        the real status. Drain up to `limit`, then let the connection close.
        """
        n = int(self.headers.get("Content-Length") or 0)
        while n > 0 and limit > 0:
            chunk = self.rfile.read(min(1 << 20, n, limit))
            if not chunk:
                break
            n -= len(chunk)
            limit -= len(chunk)
        if n > 0:
            self.close_connection = True

    def _reject(self, code: int, msg: str) -> None:
        self._drain()
        self.close_connection = True
        self._json(code, {"error": msg})

    def _handle_upload(self) -> None:
        """Accept a file the phone already has on disk.

        This is the fallback channel: downloads that never went through the proxy
        (an app using its own downloader) are copied here instead of re-fetched.
        Headers: Authorization: Bearer <token>; query: ?name=<file>&dest_dir=<dir>
        Body: the raw bytes. Written to <name>.part first, then renamed - so the
        media watcher never sees a half file.
        """
        if not self._authed():
            return self._reject(401, "unauthorized")
        q = parse_qs(urlparse(self.path).query)
        name = unquote((q.get("name") or [""])[0]).strip()
        dest_dir = unquote((q.get("dest_dir") or [""])[0]).strip()
        if not name or "/" in name or "\\" in name or name.startswith(".") or name in (".", ".."):
            return self._reject(400, "bad name")
        try:
            dest = resolve_dest(self.cfg, dest_dir, name)
        except PathError as exc:
            return self._reject(403, str(exc))
        length = int(self.headers.get("Content-Length") or 0)
        if length <= 0:
            return self._reject(400, "empty body")

        # idempotent retry: a file with the same name and size is already here
        if dest.exists() and dest.stat().st_size == length:
            self.rfile.read(length)
            return self._json(200, {"ok": True, "path": str(dest), "bytes": length, "deduped": True})

        part = dest.with_name(dest.name + ".part")
        got = 0
        try:
            with open(part, "wb") as fh:
                remaining = length
                while remaining > 0:
                    chunk = self.rfile.read(min(1 << 20, remaining))
                    if not chunk:
                        break
                    fh.write(chunk)
                    got += len(chunk)
                    remaining -= len(chunk)
            if got != length:
                return self._json(400, {"error": f"short body {got}/{length}"})
            os.replace(part, dest)
        finally:
            if part.exists() and got != length:
                part.unlink()
        self._json(200, {"ok": True, "path": str(dest), "bytes": got})

    # -- routes
    def do_GET(self) -> None:
        path = urlparse(self.path).path.rstrip("/") or "/"
        if path == "/api/health":
            usage = shutil.disk_usage(self.cfg["root"])
            queued = self.store.recent(200)
            busy = sum(1 for t in queued if t["state"] in ("queued", "running"))
            return self._json(200, {
                "ok": True, "version": VERSION, "queue": busy,
                "free_bytes": usage.free, "root": self.cfg["root"],
            })
        if not self._authed():
            return self._json(401, {"error": "unauthorized"})
        if path == "/api/tasks":
            from urllib.parse import parse_qs
            qs = parse_qs(urlparse(self.path).query)
            limit = min(int(qs.get("limit", ["20"])[0]), 200)
            return self._json(200, {"tasks": self.store.recent(limit)})
        m = re.fullmatch(r"/api/tasks/([0-9a-f]{32})", path)
        if m:
            task = self.store.get(m.group(1))
            return self._json(200, task) if task else self._json(404, {"error": "not found"})
        return self._json(404, {"error": "no such route"})

    def do_POST(self) -> None:
        path = urlparse(self.path).path.rstrip("/")
        if path == "/api/upload":
            return self._handle_upload()
        if path != "/api/tasks":
            return self._json(404, {"error": "no such route"})
        if not self._authed():
            return self._json(401, {"error": "unauthorized"})
        try:
            body = self._body()
        except Exception as exc:
            return self._json(400, {"error": f"bad json: {exc}"})
        url = (body.get("url") or "").strip()
        scheme = urlparse(url).scheme.lower()
        if scheme not in ("http", "https"):
            return self._json(400, {"error": "url must be http(s)"})
        headers = body.get("headers") or {}
        if not isinstance(headers, dict):
            return self._json(400, {"error": "headers must be an object"})
        # validate the destination before queuing so the phone learns early
        try:
            resolve_dest(self.cfg, body.get("dest_dir"), "_probe_")
        except PathError as exc:
            return self._json(403, {"error": str(exc)})
        task = {
            "id": uuid.uuid4().hex,
            "url": url,
            "headers": headers,
            "dest_dir": body.get("dest_dir"),
            "filename": body.get("filename"),
            "proxy": body.get("proxy"),
            "engine": body.get("engine"),
        }
        self.store.insert(task)
        self.tasks.put(task["id"])
        return self._json(202, {"task_id": task["id"]})


def serve(cfg: dict) -> None:
    store = Store(cfg["db"])
    tasks: "queue.Queue[str]" = queue.Queue()
    for _ in range(max(1, int(cfg.get("workers", 1)))):
        threading.Thread(target=worker_loop, args=(cfg, store, tasks), daemon=True).start()
    Handler.cfg = cfg
    Handler.store = store
    Handler.tasks = tasks
    httpd = ThreadingHTTPServer((cfg["bind"], int(cfg["port"])), Handler)
    print(f"[recv] {VERSION} listening on {cfg['bind']}:{cfg['port']} root={cfg['root']}", flush=True)
    httpd.serve_forever()


def main() -> None:
    cfg_path = Path(sys.argv[1]) if len(sys.argv) > 1 else DEFAULT_CONFIG
    if cfg_path.is_dir():
        cfg_path = cfg_path / "recv.json"
    cfg = load_config(cfg_path)
    print(f"[recv] config {cfg_path} token={cfg['token'][:6]}...", flush=True)
    serve(cfg)


if __name__ == "__main__":
    main()
