"""
Standalone Captcha OCR API + Learn Mode.
Portal / Sarathi se alag. Test: http://127.0.0.1:8765/
"""
from __future__ import annotations

import base64
import re
from pathlib import Path
from typing import Optional

from fastapi import FastAPI, File, Form, HTTPException, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import HTMLResponse
from pydantic import BaseModel, Field

import ocr_engine as eng

ROOT = Path(__file__).resolve().parent

app = FastAPI(
    title="Captcha OCR + Learn Mode",
    description="Local OCR. Learn-mode: galat ho to sahi text bhejo — dataset me save, baad me retrain.",
    version="1.1.0",
)
app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],
    allow_methods=["*"],
    allow_headers=["*"],
)


class SolveJson(BaseModel):
    image_base64: str
    only_alnum: bool = True
    max_len: int = 8
    learn: bool = True  # auto-save raw for later labeling


class FeedbackJson(BaseModel):
    image_base64: str
    correct_text: str = Field(..., min_length=3, max_length=12)
    predicted: str = ""


class SolveResult(BaseModel):
    ok: bool
    text: str
    mode: str
    message: str = ""
    labeled_count: int = 0


@app.on_event("startup")
def _startup():
    try:
        eng.load_engines()
    except Exception as e:
        print("OCR load warning:", e)


@app.get("/", response_class=HTMLResponse)
def home():
    n = eng.labeled_count()
    mode = eng.engine_mode()
    return f"""
<!DOCTYPE html>
<html lang="hi">
<head>
  <meta charset="utf-8"/>
  <title>Captcha OCR + Learn</title>
  <style>
    body{{font-family:system-ui,sans-serif;max-width:680px;margin:36px auto;padding:0 16px;background:#0f172a;color:#e2e8f0}}
    h1{{font-size:1.3rem}} .card{{background:#1e293b;border-radius:12px;padding:18px;margin:14px 0}}
    input,button{{font-size:1rem;padding:10px 14px;border-radius:8px;border:0}}
    input[type=text]{{width:min(240px,75%);background:#0f172a;color:#e2e8f0;border:1px solid #334155}}
    button{{background:#38bdf8;color:#0f172a;font-weight:700;cursor:pointer;margin:6px 6px 0 0}}
    button.sec{{background:#22c55e;color:#052e16}} button.warn{{background:#64748b;color:#fff}}
    #out{{font-size:1.45rem;font-weight:800;letter-spacing:.06em;color:#86efac}}
    .hint{{color:#94a3b8;font-size:.9rem;line-height:1.45}} a{{color:#7dd3fc}}
    #msg{{color:#fde68a;min-height:1.2em}}
  </style>
</head>
<body>
  <h1>Captcha OCR + Learn Mode</h1>
  <p class="hint">Mode: <b id="mode">{mode or "…"}</b> · Labeled: <b id="cnt">{n}</b><br>
  <b>Capital/small ~90%+:</b> free Gemini key rakho file me
  <code>c:\\xampp\\htdocs\\captcha-ocr\\.env</code> → <code>GEMINI_API_KEY=...</code>
  (Google AI Studio se free key). Mode <b>gemini</b> dikhega.<br>
  Bina key ke local OCR chalta hai (case weak). Learn optional. <b>Ctrl+F5</b>.</p>
  <div class="card">
    <p>Captcha image:</p>
    <input type="file" id="f" accept="image/*"/>
    <button type="button" id="go">Read OCR</button>
    <p>OCR: <span id="out">—</span></p>
    <p>Sahi text (capital/small exact):</p>
    <input type="text" id="label" maxlength="12" placeholder="WEUVKt"/>
    <button type="button" id="learn" class="sec">Learn (save sahi)</button>
    <button type="button" id="ok" class="warn">OCR sahi tha (save)</button>
    <p id="msg"></p>
  </div>
  <p><a href="/docs">/docs</a> · <a href="/health">/health</a> · <a href="/stats">/stats</a></p>
  <script>
    let lastFile = null, lastPred = '';
    const setMsg = (t) => document.getElementById('msg').textContent = t;
    document.getElementById('f').onchange = e => {{ lastFile = e.target.files[0] || null; }};
    document.getElementById('go').onclick = async () => {{
      const f = document.getElementById('f').files[0];
      if (!f) return alert('Image select karo');
      lastFile = f;
      const fd = new FormData(); fd.append('file', f); fd.append('learn', 'true');
      document.getElementById('out').textContent = '…';
      const j = await (await fetch('/solve', {{ method:'POST', body: fd }})).json();
      lastPred = j.text || '';
      document.getElementById('out').textContent = j.ok ? j.text : ('ERR '+j.message);
      document.getElementById('label').value = j.text || '';
      document.getElementById('mode').textContent = j.mode;
      document.getElementById('cnt').textContent = j.labeled_count;
    }};
    async function saveLabel(label) {{
      const f = lastFile || document.getElementById('f').files[0];
      if (!f) return alert('Pehle image + OCR');
      if (!label || label.length < 3) return alert('Sahi text likho');
      const fd = new FormData();
      fd.append('file', f);
      fd.append('correct_text', label);
      fd.append('predicted', lastPred);
      const j = await (await fetch('/feedback', {{ method:'POST', body: fd }})).json();
      setMsg(j.ok ? ('Saved ✓ labeled=' + j.labeled_count) : ('Fail: ' + (j.message||'')));
      if (j.labeled_count) document.getElementById('cnt').textContent = j.labeled_count;
    }}
    document.getElementById('learn').onclick = () => saveLabel((document.getElementById('label').value||'').trim());
    document.getElementById('ok').onclick = () => saveLabel(lastPred || (document.getElementById('label').value||'').trim());
  </script>
</body>
</html>
"""


@app.get("/health")
def health():
    try:
        if eng.engine_mode() == "none":
            eng.load_engines()
        return {"ok": True, "mode": eng.engine_mode(), "labeled": eng.labeled_count()}
    except Exception as e:
        return {"ok": False, "error": str(e)}


@app.get("/stats")
def stats():
    return {
        "ok": True,
        "mode": eng.engine_mode(),
        "labeled_count": eng.labeled_count(),
        "learn_log": eng.LEARN_LOG.is_file(),
        "has_custom_onnx": (eng.MODELS / "captcha.onnx").is_file(),
    }


@app.post("/solve", response_model=SolveResult)
async def solve_upload(
    file: UploadFile = File(...),
    only_alnum: bool = Form(True),
    max_len: int = Form(8),
    learn: bool = Form(True),
):
    data = await file.read()
    if not data:
        raise HTTPException(400, "Empty file")
    try:
        # Do NOT reload model every request (that made Read slow)
        if eng.engine_mode() == "none":
            eng.load_engines()
        text = eng.clean_text(eng.read_image_bytes(data), only_alnum=only_alnum, max_len=max_len)
        if learn:
            eng.DATASET_RAW.mkdir(parents=True, exist_ok=True)
            raw_path = eng.DATASET_RAW / f"raw_{int(__import__('time').time()*1000)}.png"
            raw_path.write_bytes(data)
            eng.learn_log({"event": "solve", "pred": text, "raw": str(raw_path.name)})
        return SolveResult(
            ok=bool(text),
            text=text,
            mode=eng.engine_mode(),
            message="" if text else "empty",
            labeled_count=eng.labeled_count(),
        )
    except Exception as e:
        return SolveResult(ok=False, text="", mode=eng.engine_mode(), message=str(e))


@app.post("/solve-json", response_model=SolveResult)
async def solve_json(body: SolveJson):
    b64 = body.image_base64.strip()
    if "," in b64 and b64.lower().startswith("data:"):
        b64 = b64.split(",", 1)[1]
    try:
        data = base64.b64decode(b64)
    except Exception:
        raise HTTPException(400, "Invalid base64")
    eng.load_engines() if eng.engine_mode() == "none" else None
    text = eng.clean_text(eng.read_image_bytes(data), only_alnum=body.only_alnum, max_len=body.max_len)
    if body.learn:
        eng.learn_log({"event": "solve_json", "pred": text})
    return SolveResult(ok=bool(text), text=text, mode=eng.engine_mode(), labeled_count=eng.labeled_count())


@app.post("/feedback")
async def feedback(
    file: UploadFile = File(...),
    correct_text: str = Form(...),
    predicted: str = Form(""),
):
    """Learn mode: sahi answer save → baad me retrain."""
    data = await file.read()
    if not data:
        raise HTTPException(400, "Empty file")
    try:
        path = eng.save_labeled_bytes(data, correct_text)
    except ValueError as e:
        raise HTTPException(400, str(e))
    eng.learn_log(
        {
            "event": "feedback",
            "pred": predicted,
            "correct": eng.clean_text(correct_text),
            "path": path.name,
        }
    )
    return {
        "ok": True,
        "path": str(path.relative_to(ROOT)),
        "labeled_count": eng.labeled_count(),
        "message": "Saved for learning. Retrain: train_all.bat",
    }


@app.post("/feedback-json")
async def feedback_json(body: FeedbackJson):
    b64 = body.image_base64.strip()
    if "," in b64 and b64.lower().startswith("data:"):
        b64 = b64.split(",", 1)[1]
    data = base64.b64decode(b64)
    path = eng.save_labeled_bytes(data, body.correct_text)
    eng.learn_log({"event": "feedback_json", "pred": body.predicted, "correct": body.correct_text, "path": path.name})
    return {"ok": True, "labeled_count": eng.labeled_count(), "path": path.name}


@app.post("/save-labeled")
async def save_labeled(file: UploadFile = File(...), label: str = Form(...)):
    data = await file.read()
    path = eng.save_labeled_bytes(data, label)
    return {"ok": True, "label": eng.clean_text(label), "path": str(path.relative_to(ROOT))}


if __name__ == "__main__":
    import os
    import uvicorn

    port = int(os.environ.get("PORT", "8765"))
    uvicorn.run("app:app", host="0.0.0.0", port=port)
