import os
import uuid
import time
import logging
import hashlib
import json
import asyncio
import aiofiles
from pathlib import Path
from typing import Optional

from fastapi import FastAPI, File, UploadFile, Form, HTTPException, Header, Request, Query
from fastapi.responses import HTMLResponse, JSONResponse
from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware
import uvicorn

app = FastAPI()
app.add_middleware(ProxyHeadersMiddleware, trusted_hosts="*")

UPLOADS_ROOT  = Path(__file__).parent / "uploads"
SESSIONS_ROOT = Path(__file__).parent / ".sessions"   # temp chunk storage
UPLOADS_ROOT.mkdir(exist_ok=True)
SESSIONS_ROOT.mkdir(exist_ok=True)

HTML_FILE  = Path(__file__).parent / "upload.html"
CHUNK_SIZE = 5 * 1024 * 1024   # 5 MB chunks

# Максимальный размер одного чанка (чуть больше CHUNK_SIZE для допуска)
MAX_CHUNK_BYTES = CHUNK_SIZE + 1024

# Каждые сколько чанков пишем hashes.json на диск (диагностический файл).
# Раньше писали на КАЖДЫЙ чанк — это была основная причина задержки
# ~110ms/чанк (полная пересериализация + запись растущего JSON на каждый
# запрос, O(n^2) на весь аплоад). Теперь пишем раз в N чанков + всегда
# в start/finish, где полнота файла реально важна.
HASHES_FLUSH_EVERY = 20


# ── логирование / замер времени ───────────────────────────────────────────────

logging.basicConfig(
    level=os.environ.get("UPLOAD_LOG_LEVEL", "DEBUG"),
    format="%(asctime)s.%(msecs)03d [%(levelname)s] %(message)s",
    datefmt="%H:%M:%S",
)
log = logging.getLogger("upload")


@app.middleware("http")
async def full_request_timing(request: Request, call_next):
    """Меряет ПОЛНОЕ время запроса на уровне ASGI — от первого байта до
    отправки ответа, то есть включая приём тела с сети и разбор запроса
    (multipart/parsing), которые происходят ДО того, как выполнится код
    в самом хендлере, и до которых Stopwatch внутри хендлера не дотягивается.
    Если это число сильно больше, чем internal total из хендлера — значит,
    время уходит на приём/парсинг запроса, а не на нашу логику."""
    t0 = time.perf_counter()
    response = await call_next(request)
    if request.url.path == "/api/upload/chunk":
        full_ms = (time.perf_counter() - t0) * 1000
        log.debug(f"[asgi] {request.method} {request.url.path} full_request={full_ms:.1f}ms")
    return response


class Stopwatch:
    """Секундомер по этапам. sw.lap('read') фиксирует время с предыдущего lap()."""

    __slots__ = ("_t0", "_start", "marks")

    def __init__(self):
        self._start = time.perf_counter()
        self._t0 = self._start
        self.marks: list[tuple[str, float]] = []

    def lap(self, label: str) -> None:
        now = time.perf_counter()
        self.marks.append((label, (now - self._t0) * 1000))
        self._t0 = now

    @property
    def total_ms(self) -> float:
        return (time.perf_counter() - self._start) * 1000

    def summary(self) -> str:
        parts = " ".join(f"{label}={ms:.1f}ms" for label, ms in self.marks)
        return f"{parts} | total={self.total_ms:.1f}ms"


# ── security helpers ──────────────────────────────────────────────────────────
# Ниже — минимум проверок, реально защищающих от path traversal и путаницы
# сессий/индексов. Остальная "бумажная" валидация (зарезервированные Windows
# имена, посимвольная проверка hex и т.п.) убрана: сервер общается только с
# оригинальным upload.html, который никогда не пришлёт такое, а лишние
# проверки — это лишний код и лишние (пусть и небольшие) микросекунды на
# каждый запрос.

def resolve_subdir(subdir: str) -> Path:
    """Безопасное разрешение поддиректории внутри UPLOADS_ROOT."""
    if not subdir or subdir.strip() in ("", "/", "."):
        return UPLOADS_ROOT
    clean = Path(subdir.strip().replace("\x00", "")).as_posix().lstrip("/")
    if not clean or clean == ".":
        return UPLOADS_ROOT
    target = (UPLOADS_ROOT / clean).resolve()
    try:
        target.relative_to(UPLOADS_ROOT.resolve())
    except ValueError:
        raise HTTPException(400, "Invalid subdirectory")
    return target


def safe_filename(name: str) -> str:
    """basename без пути и нулевых байт — единственное, что реально защищает
    от записи файла вне UPLOADS_ROOT. Остальное (зарезервированные имена,
    скрытые файлы) не влияет на безопасность, поэтому не проверяем."""
    name = Path(name).name.replace("\x00", "").strip()
    if not name or name in (".", ".."):
        raise HTTPException(400, "Invalid filename")
    return name


def safe_session_id(sid: str) -> str:
    """UUID только — никаких path-трюков."""
    try:
        return str(uuid.UUID(sid))
    except (ValueError, AttributeError):
        raise HTTPException(400, "Invalid session id")


def safe_chunk_index(idx: int, total_chunks: int) -> int:
    if idx < 0 or idx >= total_chunks:
        raise HTTPException(400, f"chunk_index out of range [0, {total_chunks - 1}]")
    return idx


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


def list_subdirs() -> list[str]:
    result = [""]
    for root, dirs, _ in os.walk(UPLOADS_ROOT):
        dirs.sort()
        for d in dirs:
            full = Path(root) / d
            rel = str(full.relative_to(UPLOADS_ROOT)).replace("\\", "/")
            result.append(rel)
    return result


def sha256_bytes(data: bytes) -> str:
    return hashlib.sha256(data).hexdigest()


# ── session state ─────────────────────────────────────────────────────────────
sessions: dict[str, dict] = {}


def session_dir(sid: str) -> Path:
    return SESSIONS_ROOT / sid


def chunk_path(sid: str, idx: int) -> Path:
    return session_dir(sid) / f"chunk_{idx:08d}"


def hashes_path(sid: str) -> Path:
    return session_dir(sid) / "hashes.json"


async def save_hashes(sid: str) -> float:
    """Записываем текущее состояние хешей в JSON (диагностика). Возвращает
    время записи в мс — вызывающий код логирует его отдельным lap()."""
    t0 = time.perf_counter()
    sess = sessions[sid]
    data = {
        "session_id":   sid,
        "filename":     sess["filename"],
        "total_size":   sess["total_size"],
        "total_chunks": sess["total_chunks"],
        "chunks":       sess["hashes"],
    }
    async with aiofiles.open(hashes_path(sid), "w", encoding="utf-8") as f:
        await f.write(json.dumps(data, indent=2, ensure_ascii=False))
    return (time.perf_counter() - t0) * 1000


# ── routes ────────────────────────────────────────────────────────────────────

@app.get("/", response_class=HTMLResponse)
async def index():
    if not HTML_FILE.exists():
        raise HTTPException(500, "upload.html not found")
    return HTMLResponse(HTML_FILE.read_text(encoding="utf-8"))


@app.get("/api/subdirs")
async def api_subdirs():
    return JSONResponse({"subdirs": list_subdirs()})


@app.post("/api/mkdir")
async def api_mkdir(subdir: str = Form(...)):
    target = resolve_subdir(subdir)
    target.mkdir(parents=True, exist_ok=True)
    rel = str(target.relative_to(UPLOADS_ROOT)).replace("\\", "/")
    return JSONResponse({"ok": True, "path": rel})


# ── chunked upload ────────────────────────────────────────────────────────────

@app.post("/api/upload/start")
async def upload_start(
    filename:   str = Form(...),
    subdir:     str = Form(""),
    total_size: int = Form(...),
):
    sw = Stopwatch()

    if total_size < 0 or total_size > 500 * 1024 * 1024 * 1024:  # 500 ГБ
        raise HTTPException(400, "Invalid total_size")

    fname  = safe_filename(filename)
    target = resolve_subdir(subdir)
    target.mkdir(parents=True, exist_ok=True)
    sw.lap("validate")

    sid  = str(uuid.uuid4())
    sdir = session_dir(sid)
    sdir.mkdir(parents=True, exist_ok=True)
    sw.lap("mkdir_session")

    total_chunks = max(1, -(-total_size // CHUNK_SIZE))   # ceil div

    sessions[sid] = {
        "filename":           fname,
        "subdir":             subdir,
        "total_size":         total_size,
        "total_chunks":       total_chunks,
        "session_dir":        sdir,
        "hashes":             {},
        "chunks_since_flush": 0,
    }

    flush_ms = await save_hashes(sid)
    sw.lap(f"save_hashes({flush_ms:.1f}ms)")

    log.info(f"[{sid[:8]}] START '{fname}' size={total_size} chunks={total_chunks} | {sw.summary()}")

    return JSONResponse({
        "session_id":      sid,
        "total_chunks":     total_chunks,
        "received_chunks":  [],
    })


@app.post("/api/upload/chunk")
async def upload_chunk(
    request:      Request,
    session_id:   str  = Query(...),
    chunk_index:  int  = Query(...),
    client_hash:  str  = Query(...),     # SHA-256 hex от клиента
):
    """
    Принять один чанк, проверить его хеш.
    При несовпадении возвращает ok=False — клиент должен переотправить.

    Метаданные идут в query-параметрах, а тело запроса — чистые байты
    чанка (Content-Type: application/octet-stream), БЕЗ multipart/form-data.
    Раньше файл принимался как UploadFile через multipart — Starlette в
    этом случае парсит тело через python-multipart и заворачивает файл в
    SpooledTemporaryFile, который на чанках 5MB почти всегда спилливается
    на диск ДО того, как выполнится код хендлера — отсюда была разница
    между ~5ms "internal" времени и ~90ms фактического ответа. Сырой body
    читается напрямую в bytes, без временного файла и без парсера форм.
    """
    sw = Stopwatch()

    sid = safe_session_id(session_id)
    if sid not in sessions:
        raise HTTPException(404, "Session not found — may have expired")
    sess = sessions[sid]
    safe_chunk_index(chunk_index, sess["total_chunks"])

    client_hash = client_hash.strip().lower()
    if len(client_hash) != 64:
        raise HTTPException(400, "Invalid client_hash format (expected SHA-256 hex)")

    # Ранний отсев слишком больших тел по Content-Length (best-effort —
    # заголовку не доверяем полностью, финальную проверку делаем по
    # факту прочитанных байт ниже).
    content_length = request.headers.get("content-length")
    if content_length is not None:
        try:
            if int(content_length) > MAX_CHUNK_BYTES:
                raise HTTPException(413, "Chunk too large")
        except ValueError:
            pass
    sw.lap("validate")

    data = await request.body()
    if len(data) > MAX_CHUNK_BYTES:
        raise HTTPException(413, "Chunk too large")
    sw.lap("read_upload")

    server_hash = sha256_bytes(data)
    sw.lap("sha256")

    hashes_entry = {
        "client": client_hash,
        "server": server_hash,
        "ok":     server_hash == client_hash,
        "size":   len(data),
    }
    sess["hashes"][str(chunk_index)] = hashes_entry

    if server_hash != client_hash:
        log.warning(f"[{sid[:8]}] chunk {chunk_index} HASH MISMATCH | {sw.summary()}")
        return JSONResponse({
            "ok":          False,
            "chunk_index": chunk_index,
            "server_hash": server_hash,
            "client_hash": client_hash,
            "error":       "Hash mismatch — please resend this chunk",
        }, status_code=200)

    cpath = chunk_path(sid, chunk_index)
    async with aiofiles.open(cpath, "wb") as f:
        await f.write(data)
    sw.lap("write_chunk")

    # Раз в HASHES_FLUSH_EVERY чанков подкидываем диагностику на диск.
    # Полная перезапись hashes.json на каждый чанк раньше и давала те самые
    # ~110ms — файл дописывался целиком и рос с каждым чанком.
    sess["chunks_since_flush"] += 1
    if sess["chunks_since_flush"] >= HASHES_FLUSH_EVERY:
        flush_ms = await save_hashes(sid)
        sess["chunks_since_flush"] = 0
        sw.lap(f"save_hashes({flush_ms:.1f}ms)")

    received_chunks = [int(k) for k, v in sess["hashes"].items() if v["ok"]]

    log.debug(f"[{sid[:8]}] chunk {chunk_index}/{sess['total_chunks']} ok | {sw.summary()}")

    return JSONResponse({
        "ok":              True,
        "chunk_index":     chunk_index,
        "server_hash":     server_hash,
        "received_chunks": received_chunks,
        "total_chunks":    sess["total_chunks"],
    })


@app.post("/api/upload/finish")
async def upload_finish(session_id: str = Form(...)):
    """Собрать все чанки в итоговый файл (по порядку индексов)."""
    sw = Stopwatch()

    sid = safe_session_id(session_id)
    if sid not in sessions:
        raise HTTPException(404, "Session not found")

    sess         = sessions[sid]
    total_chunks = sess["total_chunks"]
    hashes       = sess["hashes"]

    missing = []
    bad     = []
    for i in range(total_chunks):
        key = str(i)
        if key not in hashes:
            missing.append(i)
        elif not hashes[key]["ok"]:
            bad.append(i)
    sw.lap("verify")

    if missing or bad:
        return JSONResponse({
            "ok":      False,
            "missing": missing,
            "bad":     bad,
            "error":   "Not all chunks received or verified",
        }, status_code=409)

    target = resolve_subdir(sess["subdir"])
    dest   = unique_dest(target / sess["filename"])
    target.mkdir(parents=True, exist_ok=True)

    async with aiofiles.open(dest, "wb") as out:
        for i in range(total_chunks):
            cpath = chunk_path(sid, i)
            if not cpath.exists():
                raise HTTPException(500, f"Chunk file missing for index {i}")
            async with aiofiles.open(cpath, "rb") as inp:
                while True:
                    buf = await inp.read(1024 * 1024)
                    if not buf:
                        break
                    await out.write(buf)
    sw.lap("assemble")

    final_hashes_path = dest.with_suffix(dest.suffix + ".chunks.json")
    hashes_data = {
        "session_id":   sid,
        "filename":     sess["filename"],
        "total_size":   sess["total_size"],
        "total_chunks": total_chunks,
        "chunks": {k: v for k, v in sorted(hashes.items(), key=lambda x: int(x[0]))},
    }
    async with aiofiles.open(final_hashes_path, "w", encoding="utf-8") as f:
        await f.write(json.dumps(hashes_data, indent=2, ensure_ascii=False))
    sw.lap("write_final_hashes")

    import shutil
    try:
        shutil.rmtree(str(sess["session_dir"]))
    except Exception:
        pass
    del sessions[sid]
    sw.lap("cleanup")

    rel = str(dest.relative_to(UPLOADS_ROOT)).replace("\\", "/")
    log.info(f"[{sid[:8]}] FINISH '{rel}' size={dest.stat().st_size} | {sw.summary()}")

    return JSONResponse({
        "ok":          True,
        "path":        rel,
        "size":        dest.stat().st_size,
        "hashes_file": str(final_hashes_path.relative_to(UPLOADS_ROOT)).replace("\\", "/"),
    })


@app.get("/api/upload/status/{session_id}")
async def upload_status(session_id: str):
    sid = safe_session_id(session_id)
    if sid not in sessions:
        raise HTTPException(404, "Session not found")
    sess = sessions[sid]
    received_chunks = [int(k) for k, v in sess["hashes"].items() if v["ok"]]
    return JSONResponse({
        "total_chunks":    sess["total_chunks"],
        "received_chunks": received_chunks,
        "total_size":      sess["total_size"],
    })


@app.delete("/api/upload/{session_id}")
async def upload_cancel(session_id: str):
    sid = safe_session_id(session_id)
    if sid in sessions:
        import shutil
        try:
            shutil.rmtree(str(sessions[sid]["session_dir"]))
        except Exception:
            pass
        del sessions[sid]
    return JSONResponse({"ok": True})


# ── main ──────────────────────────────────────────────────────────────────────

if __name__ == "__main__":
    # ВАЖНО: sessions живёт в памяти одного процесса. При workers > 1 uvicorn
    # поднимает несколько НЕЗАВИСИМЫХ процессов с разными dict sessions —
    # запрос на чанк может прилететь не в тот воркер, где создавалась сессия
    # (upload/start), и тогда сервер честно ответит 404 "Session not found",
    # хотя с точки зрения клиента всё было отправлено правильно. Это не имеет
    # отношения к задержке 110ms, но может выглядеть как случайные обрывы
    # аплоада на многоядерной машине. Пока sessions не вынесены в общее
    # хранилище (redis/файл) — держим один воркер.
    workers = 1
    print(f"Upload server | uploads: {UPLOADS_ROOT} | http://localhost:8001 | workers={workers}")
    uvicorn.run(
        app, host="127.0.0.1", port=8001, workers=workers,
        timeout_keep_alive=3600, loop="asyncio",
    )
