#!/usr/bin/env python3
"""
Chunked upload of a large file to MyAirBridge (WebDAV API).

Documentation: https://info.myairbridge.com/en/api#chunked-upload

How it works:
  1. POST  TARGET_DIRECTORY_URL/:chunked-upload  -> X-Upload-Chunk-Url header
  2. PUT   each chunk to X-Upload-Chunk-Url (X-File-Range header, half-open
           interval [from, to), e.g. 0-104857600 = first 100 MiB)
  3. Once the last chunk is uploaded, the server creates the target file.

On error, the script prints the request that failed (method, URL, headers,
body size; the password in Authorization is masked) and the server response.
With --verbose, every request and response is printed, even successful ones.

Query parameters of the --target URL (e.g. ?token=...) are added to every
request, including the chunk URL returned by the server.

Usage:
  # with authentication
  export MAB_PASSWORD='password'
  python3 mab_chunked_upload.py --user username \\
      --target 'https://.../my-directory' /path/to/large.file

  # password-protected directory (username 'guest' is used)
  python3 mab_chunked_upload.py --password 'password' \\
      --target 'https://.../my-directory' /path/to/large.file

  # without authentication (public directory)
  python3 mab_chunked_upload.py --target 'https://.../my-directory' file

Requirements: Python 3.9+, pip install requests
"""

import argparse
import base64
import getpass
import os
import re
import sys
import threading
import time
import uuid
from concurrent.futures import ThreadPoolExecutor, as_completed
from urllib.parse import parse_qsl, quote, urlsplit, urlunsplit

import requests

MIB = 1024 * 1024
DEFAULT_CHUNK_SIZE = 100 * MIB
MAX_NAME_LEN = 128  # per the docs: 1-128 UTF-8 characters
MAX_BODY_PRINT = 2000

VERBOSE = False     # set by --verbose
TARGET_QUERY = ""   # raw query string of --target, added to every request
_print_lock = threading.Lock()


class FatalError(Exception):
    """An error that is not worth retrying (e.g. HTTP 4xx)."""


def h(value) -> bytes:
    """Header value as UTF-8 bytes (requests would use latin-1 otherwise)."""
    return str(value).encode("utf-8")


def log(*args, **kwargs):
    """Thread-safe print to stderr."""
    with _print_lock:
        print(*args, file=sys.stderr, **kwargs)


# ---------------------------------------------------------------------------
# URLs
# ---------------------------------------------------------------------------

def join_url(base: str, segment: str) -> str:
    """Append a segment to the URL path, keeping the query string (?...).

    'https://x/dir?token=1' + ':chunked-upload'
        -> 'https://x/dir/:chunked-upload?token=1'
    """
    parts = urlsplit(base)
    path = parts.path.rstrip("/") + "/" + segment
    return urlunsplit(parts._replace(path=path))


def add_query(url: str, extra_query: str) -> str:
    """Add the parameters from extra_query to the URL.

    Parameters the URL already contains (by name) are left untouched, so the
    target parameters are never duplicated and values set by the server win.
    The original encoding of the parameters is preserved.
    """
    if not extra_query:
        return url
    parts = urlsplit(url)
    existing = {k for k, _ in parse_qsl(parts.query, keep_blank_values=True)}
    added = []
    for piece in extra_query.split("&"):
        if not piece:
            continue
        parsed = parse_qsl(piece, keep_blank_values=True)
        key = parsed[0][0] if parsed else piece.split("=", 1)[0]
        if key in existing:
            continue
        added.append(piece)
    if not added:
        return url
    query = "&".join(([parts.query] if parts.query else []) + added)
    return urlunsplit(parts._replace(query=query))


# ---------------------------------------------------------------------------
# Printing requests and responses
# ---------------------------------------------------------------------------

def _hv(v) -> str:
    return v.decode("utf-8", "replace") if isinstance(v, (bytes, bytearray)) else str(v)


def _mask_auth(value: str) -> str:
    """Basic auth: show the username, hide the password."""
    scheme, _, token = value.partition(" ")
    if scheme.lower() == "basic":
        try:
            user = base64.b64decode(token).decode("utf-8", "replace").split(":", 1)[0]
            return f"Basic <{user}:***>"
        except Exception:
            pass
    return "<hidden>"


def describe_request(req) -> str:
    if req is None:
        return "  (request not available)"
    lines = [f"  {req.method} {req.url}"]
    for k, v in req.headers.items():
        v = _hv(v)
        if k.lower() == "authorization":
            v = _mask_auth(v)
        lines.append(f"  {k}: {v}")
    body = req.body
    if body is None:
        desc = "(none)"
    elif isinstance(body, (bytes, bytearray, str)):
        desc = f"{len(body)} B (binary content not shown)"
    else:
        desc = "streamed from file"
    lines.append(f"  Body: {desc}")
    return "\n".join(lines)


def describe_response(r) -> str:
    lines = [f"  HTTP {r.status_code} {r.reason or ''}".rstrip()]
    lines += [f"  {k}: {v}" for k, v in r.headers.items()]
    text = r.text or ""
    if not text:
        lines.append("  Body: (empty)")
    else:
        cut = " ...(truncated)" if len(text) > MAX_BODY_PRINT else ""
        lines.append("  Body:\n" + text[:MAX_BODY_PRINT] + cut)
    return "\n".join(lines)


def error_report(msg, response=None, exc=None) -> str:
    """Error message including the request and the response / exception."""
    if response is not None:
        req = response.request
    else:
        req = getattr(exc, "request", None)
    parts = [msg, "--- Request ---", describe_request(req)]
    if response is not None:
        parts += ["--- Response ---", describe_response(response)]
    if exc is not None:
        parts += ["--- Exception ---", f"  {exc!r}"]
    return "\n".join(parts)


# ---------------------------------------------------------------------------
# HTTP
# ---------------------------------------------------------------------------

_local = threading.local()


def get_session(auth) -> requests.Session:
    """One Session per thread (requests.Session is not guaranteed thread-safe)."""
    s = getattr(_local, "session", None)
    if s is None:
        s = requests.Session()
        s.auth = auth
        _local.session = s
    return s


def send(auth, method, url, **kwargs) -> requests.Response:
    """Send a request; with --verbose, print the request and the response."""
    url = add_query(url, TARGET_QUERY)
    try:
        r = get_session(auth).request(method, url, **kwargs)
    except requests.RequestException as e:
        if VERBOSE:
            log(error_report(f"[verbose] {method} failed", exc=e) + "\n")
        raise
    if VERBOSE:
        log("\n".join([
            f"[verbose] {method} -> HTTP {r.status_code}",
            "--- Request ---", describe_request(r.request),
            "--- Response ---", describe_response(r),
        ]) + "\n")
    return r


def create_upload(auth, dir_url, upload_name, file_name, size, mtime):
    url = join_url(dir_url, ":chunked-upload")
    headers = {
        "X-Upload-Name": h(upload_name),
        "X-File-Name": h(file_name),
        "X-File-Size": h(size),
        "X-Last-Modified": h(mtime),
    }
    try:
        r = send(auth, "POST", url, headers=headers, timeout=60)
    except requests.RequestException as e:
        raise FatalError(error_report("Creating the upload failed (network error).", exc=e))
    if not r.ok:
        raise FatalError(error_report(
            f"Creating the upload failed: HTTP {r.status_code}", response=r))
    chunk_url = r.headers.get("X-Upload-Chunk-Url")
    if not chunk_url:
        raise FatalError(error_report(
            "The server did not return the X-Upload-Chunk-Url header.", response=r))
    return chunk_url


def upload_chunk(auth, chunk_url, path, upload_name, file_name,
                 start, length, retries):
    # X-File-Range is a half-open interval [from, to): 0-104857600 = first 100 MiB
    end = start + length
    with open(path, "rb") as f:
        f.seek(start)
        data = f.read(length)
    if len(data) != length:
        raise FatalError(f"Read {len(data)} B instead of {length} B (offset {start}).")

    headers = {
        "X-Upload-Name": h(upload_name),
        "X-File-Name": h(file_name),
        "X-File-Range": h(f"{start}-{end}"),
        "Content-Type": "application/octet-stream",
    }

    last_resp, last_exc = None, None
    for attempt in range(1, retries + 1):
        last_resp, last_exc = None, None
        try:
            r = send(auth, "PUT", chunk_url, data=data, headers=headers,
                     timeout=(30, 900))
            if r.ok:
                return
            last_resp = r
            if 400 <= r.status_code < 500 and r.status_code not in (408, 409, 429):
                raise FatalError(error_report(
                    f"Chunk {start}-{end} rejected: HTTP {r.status_code}",
                    response=r))
            short = f"HTTP {r.status_code}"
        except requests.RequestException as e:
            last_exc = e
            short = repr(e)
        if attempt < retries:
            wait = min(2 ** attempt, 60)
            log(f"  ! chunk {start}-{end}, attempt {attempt}/{retries} failed "
                f"({short}); retrying in {wait} s")
            time.sleep(wait)
    raise RuntimeError(error_report(
        f"Chunk {start}-{end} could not be uploaded after {retries} attempts.",
        response=last_resp, exc=last_exc))


def upload_whole(auth, dir_url, path, file_name):
    """Plain PUT (curl -T), used for an empty file."""
    url = join_url(dir_url, quote(file_name, safe=""))
    try:
        with open(path, "rb") as f:
            r = send(auth, "PUT", url, data=f, timeout=(30, 900))
    except requests.RequestException as e:
        raise FatalError(error_report("Upload failed (network error).", exc=e))
    if not r.ok:
        raise FatalError(error_report(f"Upload failed: HTTP {r.status_code}", response=r))


def verify(auth, dir_url, file_name, expected_size, attempts=5):
    """Best-effort check of the resulting file size via PROPFIND.

    Returns (ok, remote_size, error_report)."""
    url = join_url(dir_url, quote(file_name, safe=""))
    report = None
    for i in range(attempts):
        try:
            r = send(auth, "PROPFIND", url, headers={"Depth": "0"}, timeout=60)
            if r.ok:
                m = re.search(r"<d:getcontentlength>(\d+)</d:getcontentlength>", r.text)
                if m:
                    remote = int(m.group(1))
                    return remote == expected_size, remote, None
                report = error_report("PROPFIND did not return d:getcontentlength.",
                                      response=r)
            else:
                report = error_report(f"PROPFIND failed: HTTP {r.status_code}", response=r)
        except requests.RequestException as e:
            report = error_report("PROPFIND failed (network error).", exc=e)
        time.sleep(2 * (i + 1))
    return None, None, report


def fmt(n):
    return f"{n / MIB:,.1f} MiB"


def main():
    global VERBOSE, TARGET_QUERY

    p = argparse.ArgumentParser(description="Chunked upload to MyAirBridge.")
    p.add_argument("file", help="File to upload")
    p.add_argument("--target", required=True,
                   help="Target directory URL (the m:node-url value); its query "
                        "parameters are added to every request")
    p.add_argument("--user", default=None,
                   help="Username. Without --user and a password, no authentication "
                        "is used; with a password only, 'guest' is used.")
    p.add_argument("--password", default=None,
                   help="Password (preferably via the MAB_PASSWORD environment "
                        "variable; if --user is given without a password, you will "
                        "be prompted)")
    p.add_argument("--name", default=None,
                   help="Target file name (default: name of the local file)")
    p.add_argument("--upload-name", default=None,
                   help="Upload name, must be unique in the directory "
                        "(default: randomly generated, e.g. upload-<uuid>)")
    p.add_argument("--chunk-size", type=int, default=DEFAULT_CHUNK_SIZE // MIB,
                   help="Chunk size in MiB (default: 100)")
    p.add_argument("--parallel", type=int, default=1,
                   help="Number of chunks uploaded in parallel (default: 1; "
                        "each one holds a chunk in RAM)")
    p.add_argument("--retries", type=int, default=5,
                   help="Number of attempts per chunk (default: 5)")
    p.add_argument("--no-verify", action="store_true",
                   help="Do not verify the file size via PROPFIND at the end")
    p.add_argument("--verbose", action="store_true",
                   help="Print every request and response, even successful ones")
    args = p.parse_args()

    VERBOSE = args.verbose
    TARGET_QUERY = urlsplit(args.target).query

    path = args.file
    if not os.path.isfile(path):
        sys.exit(f"File not found: {path}")

    password = args.password or os.environ.get("MAB_PASSWORD")
    if args.user:
        if password is None:
            password = getpass.getpass(f"Password for {args.user}: ")
        auth = (args.user, password)
    elif password:
        # password-only protected directory: 'guest' always works per the docs
        auth = ("guest", password)
    else:
        auth = None  # no authentication

    file_name = args.name or os.path.basename(path)
    upload_name = args.upload_name or f"upload-{uuid.uuid4().hex}"
    for label, val in (("File name", file_name), ("Upload name", upload_name)):
        if not 1 <= len(val) <= MAX_NAME_LEN:
            sys.exit(f"{label} must be 1-{MAX_NAME_LEN} characters long: {val!r}")

    st = os.stat(path)
    size = st.st_size
    mtime = int(st.st_mtime)
    chunk_size = args.chunk_size * MIB

    try:
        if size == 0:
            print("The file is empty, using a plain PUT.")
            upload_whole(auth, args.target, path, file_name)
            print("Done.")
            return

        chunks = [(off, min(chunk_size, size - off))
                  for off in range(0, size, chunk_size)]
        print(f"File: {path} ({fmt(size)}), {len(chunks)} chunks of "
              f"{args.chunk_size} MiB, parallel: {args.parallel}")

        chunk_url = create_upload(auth, args.target, upload_name,
                                  file_name, size, mtime)
        print(f"Upload created (upload name: {upload_name}).")

        done_bytes = 0
        t0 = time.monotonic()

        def job(off, length):
            upload_chunk(auth, chunk_url, path, upload_name, file_name,
                         off, length, args.retries)
            return length

        ex = ThreadPoolExecutor(max_workers=max(1, args.parallel))
        try:
            futures = [ex.submit(job, off, ln) for off, ln in chunks]
            for fut in as_completed(futures):
                done_bytes += fut.result()  # raises if the chunk failed
                elapsed = max(time.monotonic() - t0, 1e-6)
                print(f"  {done_bytes / size * 100:5.1f} %  "
                      f"{fmt(done_bytes)} / {fmt(size)}  "
                      f"({done_bytes / elapsed / MIB:.1f} MiB/s)")
        finally:
            # on error, do not start any further chunks
            ex.shutdown(wait=True, cancel_futures=True)

        print("All chunks uploaded, the server is creating the file.")

        if not args.no_verify:
            ok, remote, report = verify(auth, args.target, file_name, size)
            if ok:
                print(f"Verified: the remote file has {remote} B.")
            elif ok is False:
                log(f"WARNING: remote size {remote} B != local size {size} B.")
                sys.exit(2)
            else:
                log("Could not verify the file size.\n" + (report or ""))

    except (FatalError, RuntimeError) as e:
        sys.exit(f"ERROR: {e}")
    except KeyboardInterrupt:
        sys.exit("\nInterrupted.")


if __name__ == "__main__":
    main()
