import os
import shutil
import tempfile
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path

import duckdb
from fastapi import FastAPI, UploadFile, File, HTTPException, Query, Header, Cookie, BackgroundTasks, Depends, Request, Response
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse
from pydantic import BaseModel
from slowapi import Limiter, _rate_limit_exceeded_handler
from slowapi.util import get_remote_address
from slowapi.errors import RateLimitExceeded

import jobs
import auth
from cleaning import clean_leads

# ---------------------------------------------------------------------------
# Config (env-driven, sensible local defaults)
# ---------------------------------------------------------------------------

MAX_UPLOAD_MB = int(os.environ.get("MAX_UPLOAD_MB", "1024"))
ALLOWED_ORIGINS = os.environ.get("ALLOWED_ORIGINS", "http://localhost:5173").split(",")
API_KEY = os.environ.get("API_KEY")  # optional: alternative to login, for scripts/automation
UPLOAD_RATE_LIMIT = os.environ.get("UPLOAD_RATE_LIMIT", "10/minute")
LOGIN_RATE_LIMIT = os.environ.get("LOGIN_RATE_LIMIT", "10/minute")
COOKIE_NAME = "session"
COOKIE_SECURE = os.environ.get("COOKIE_SECURE", "false").lower() == "true"

limiter = Limiter(key_func=get_remote_address)
app = FastAPI(title="Lead Cleaner API")
app.state.limiter = limiter
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)

app.add_middleware(
    CORSMiddleware,
    allow_origins=ALLOWED_ORIGINS,
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)


def require_auth(
    session: str | None = Cookie(default=None),
    x_api_key: str | None = Header(default=None, alias="X-API-Key"),
):
    """Accepts EITHER a valid logged-in session cookie OR the API key (for
    scripts/automation that shouldn't have to log in interactively)."""
    if API_KEY and x_api_key == API_KEY:
        return {"email": "api-key", "role": "admin", "id": None, "via_api_key": True}
    user = auth.get_session_user(session)
    if user:
        return {**user, "via_api_key": False}
    raise HTTPException(status_code=401, detail="Not authenticated")


def require_admin(user=Depends(require_auth)):
    if user["role"] != "admin":
        raise HTTPException(status_code=403, detail="Admin access required")
    return user


# ---------------------------------------------------------------------------
# Auth endpoints
# ---------------------------------------------------------------------------

class SetupBody(BaseModel):
    email: str
    password: str


class LoginBody(BaseModel):
    email: str
    password: str


class ChangePasswordBody(BaseModel):
    current_password: str
    new_password: str


class CreateUserBody(BaseModel):
    email: str
    password: str
    role: str  # "admin" | "member"


class UpdateRoleBody(BaseModel):
    role: str


class AdminResetPasswordBody(BaseModel):
    new_password: str


def _set_session_cookie(response: Response, token: str):
    response.set_cookie(
        key=COOKIE_NAME,
        value=token,
        httponly=True,
        secure=COOKIE_SECURE,
        samesite="lax",
        max_age=auth.SESSION_TTL_HOURS * 3600,
        path="/",
    )


@app.get("/api/auth/status")
def auth_status():
    return {"setup_complete": auth.is_setup_complete()}


@app.post("/api/auth/setup")
@limiter.limit(LOGIN_RATE_LIMIT)
def auth_setup(request: Request, response: Response, body: SetupBody):
    # This creates the FIRST account only, and it is always 'admin' —
    # ignoring any role input entirely avoids a race where someone posts
    # here with role=member and it's silently accepted before anyone exists.
    if auth.is_setup_complete():
        raise HTTPException(status_code=403, detail="Setup already complete. Please log in.")
    if not auth.validate_email(body.email):
        raise HTTPException(status_code=400, detail="Please enter a valid email address")
    pw_error = auth.validate_password_strength(body.password)
    if pw_error:
        raise HTTPException(status_code=400, detail=pw_error)

    user = auth.create_user(body.email, body.password, role="admin")
    token = auth.create_session(user["email"])
    _set_session_cookie(response, token)
    return {"id": user["id"], "email": user["email"], "role": user["role"]}


@app.post("/api/auth/login")
@limiter.limit(LOGIN_RATE_LIMIT)
def auth_login(request: Request, response: Response, body: LoginBody):
    user = auth.authenticate(body.email, body.password)
    if not user:
        raise HTTPException(status_code=401, detail="Incorrect email or password")
    token = auth.create_session(user["email"])
    _set_session_cookie(response, token)
    return {"id": user["id"], "email": user["email"], "role": user["role"]}


@app.post("/api/auth/logout")
def auth_logout(response: Response, session: str | None = Cookie(default=None)):
    auth.delete_session(session)
    response.delete_cookie(COOKIE_NAME, path="/")
    return {"ok": True}


@app.get("/api/auth/me")
def auth_me(user=Depends(require_auth)):
    return {"id": user.get("id"), "email": user["email"], "role": user["role"]}


@app.post("/api/auth/change-password")
def auth_change_password(body: ChangePasswordBody, user=Depends(require_auth)):
    if user.get("via_api_key"):
        raise HTTPException(status_code=403, detail="Log in with your account to change the password")
    pw_error = auth.validate_password_strength(body.new_password)
    if pw_error:
        raise HTTPException(status_code=400, detail=pw_error)
    if not auth.change_own_password(user["email"], body.current_password, body.new_password):
        raise HTTPException(status_code=401, detail="Current password is incorrect")
    return {"ok": True}


# ---------------------------------------------------------------------------
# User management (admin only)
# ---------------------------------------------------------------------------

@app.get("/api/users")
def list_users(admin=Depends(require_admin)):
    return {"users": auth.list_users()}


@app.post("/api/users")
def create_user(body: CreateUserBody, admin=Depends(require_admin)):
    if not auth.validate_email(body.email):
        raise HTTPException(status_code=400, detail="Please enter a valid email address")
    pw_error = auth.validate_password_strength(body.password)
    if pw_error:
        raise HTTPException(status_code=400, detail=pw_error)
    if body.role not in auth.VALID_ROLES:
        raise HTTPException(status_code=400, detail="role must be 'admin' or 'member'")
    try:
        user = auth.create_user(body.email, body.password, body.role)
    except ValueError as e:
        raise HTTPException(status_code=409, detail=str(e))
    return user


@app.delete("/api/users/{user_id}")
def delete_user(user_id: int, admin=Depends(require_admin)):
    target = auth.get_user_by_id(user_id)
    if not target:
        raise HTTPException(status_code=404, detail="User not found")
    if admin.get("id") == user_id:
        raise HTTPException(status_code=400, detail="You can't delete your own account while logged in as it")
    if target["role"] == "admin" and auth.count_admins() <= 1:
        raise HTTPException(status_code=400, detail="Can't delete the last remaining admin")
    auth.delete_user(user_id)
    return {"deleted": True}


@app.patch("/api/users/{user_id}/role")
def update_user_role(user_id: int, body: UpdateRoleBody, admin=Depends(require_admin)):
    target = auth.get_user_by_id(user_id)
    if not target:
        raise HTTPException(status_code=404, detail="User not found")
    if body.role not in auth.VALID_ROLES:
        raise HTTPException(status_code=400, detail="role must be 'admin' or 'member'")
    if target["role"] == "admin" and body.role == "member" and auth.count_admins() <= 1:
        raise HTTPException(status_code=400, detail="Can't demote the last remaining admin")
    auth.update_user_role(user_id, body.role)
    return {"ok": True}


@app.post("/api/users/{user_id}/reset-password")
def admin_reset_password(user_id: int, body: AdminResetPasswordBody, admin=Depends(require_admin)):
    target = auth.get_user_by_id(user_id)
    if not target:
        raise HTTPException(status_code=404, detail="User not found")
    pw_error = auth.validate_password_strength(body.new_password)
    if pw_error:
        raise HTTPException(status_code=400, detail=pw_error)
    auth.admin_reset_password(user_id, body.new_password)
    return {"ok": True}


# ---------------------------------------------------------------------------
# Background processing — bounded worker pool
#
# Uploads used to spawn a raw thread per job with no limit, so a burst of
# uploads (bulk-by-definition, now that multi-file upload exists) could spin
# up dozens of simultaneous Polars/DuckDB jobs fighting over CPU and RAM. A
# fixed-size pool means excess jobs simply wait (visible to the user as
# status="pending") instead of degrading the whole server.
# ---------------------------------------------------------------------------

MAX_CONCURRENT_JOBS = int(os.environ.get("MAX_CONCURRENT_JOBS", "3"))
job_executor = ThreadPoolExecutor(max_workers=MAX_CONCURRENT_JOBS, thread_name_prefix="clean-job")


def _run_cleaning_job(job_id: str, params: dict):
    d = jobs.job_dir(job_id)
    inputs_dir = jobs.inputs_dir(job_id)
    try:
        jobs.set_status(job_id, status="processing")
        input_paths = sorted(inputs_dir.glob("*.csv"))
        if not input_paths:
            raise ValueError("No input files found for this job")
        result = clean_leads(input_paths, **params)
        result.clean_df.write_parquet(d / "clean.parquet")
        result.rejected_df.write_parquet(d / "rejected.parquet")
        jobs.set_status(
            job_id,
            status="done",
            stats=result.stats,
            category_column=result.category_column,
        )
    except Exception as e:
        jobs.set_status(job_id, status="error", error=str(e))
    finally:
        # Raw uploads no longer needed once Parquet outputs exist.
        if inputs_dir.exists():
            shutil.rmtree(inputs_dir, ignore_errors=True)


# ---------------------------------------------------------------------------
# Upload (supports multiple files at once, merged into one cleaned dataset)
# ---------------------------------------------------------------------------

@app.post("/api/upload", dependencies=[Depends(require_auth)])
@limiter.limit(UPLOAD_RATE_LIMIT)
async def upload_csv(
    request: Request,
    files: list[UploadFile] = File(...),
    dedupe_key: str = Query("email", pattern="^(email|domain|row)$"),
    explode_phones: bool = Query(True),
    exclude_placeholder_stores: bool = Query(False),
    user=Depends(require_auth),
):
    if not files:
        raise HTTPException(status_code=400, detail="No files provided")
    for f in files:
        if not f.filename.lower().endswith(".csv"):
            raise HTTPException(status_code=400, detail=f"'{f.filename}' is not a .csv file")

    job_id = jobs.create_job(
        [f.filename for f in files],
        created_by=user.get("email") if not user.get("via_api_key") else "api-key",
    )
    inputs_dir = jobs.inputs_dir(job_id)

    max_bytes = MAX_UPLOAD_MB * 1024 * 1024
    total_size = 0
    try:
        for i, f in enumerate(files):
            dest = inputs_dir / f"{i:04d}.csv"
            with open(dest, "wb") as out:
                while chunk := await f.read(1024 * 1024):
                    total_size += len(chunk)
                    if total_size > max_bytes:
                        raise HTTPException(
                            status_code=413,
                            detail=f"Combined upload exceeds {MAX_UPLOAD_MB}MB limit",
                        )
                    out.write(chunk)
            # quick content sniff — first bytes should look like text/CSV,
            # not e.g. a zip/exe magic number smuggled in with a .csv extension
            with open(dest, "rb") as check:
                head = check.read(4)
                if head[:2] in (b"PK", b"MZ") or head[:4] == b"\x7fELF":
                    raise HTTPException(status_code=400, detail=f"'{f.filename}' does not look like a valid CSV")
    except HTTPException:
        jobs.delete_job(job_id)
        raise

    params = dict(
        dedupe_key=dedupe_key,
        explode_phones=explode_phones,
        exclude_placeholder_stores=exclude_placeholder_stores,
    )
    job_executor.submit(_run_cleaning_job, job_id, params)

    return {"job_id": job_id}


# ---------------------------------------------------------------------------
# Job history
# ---------------------------------------------------------------------------

@app.get("/api/jobs", dependencies=[Depends(require_auth)])
def list_jobs(limit: int = Query(50, ge=1, le=200)):
    all_ids = jobs.list_jobs()
    entries = []
    for job_id in all_ids:
        st = jobs.get_status(job_id)
        if st:
            entries.append({"job_id": job_id, **st})
    entries.sort(key=lambda e: e.get("updated_at", ""), reverse=True)
    return {"jobs": entries[:limit]}

@app.get("/api/jobs/{job_id}/status", dependencies=[Depends(require_auth)])
def job_status(job_id: str):
    st = jobs.get_status(job_id)
    if st is None:
        raise HTTPException(status_code=404, detail="Job not found")
    return st


# ---------------------------------------------------------------------------
# Query helpers
# ---------------------------------------------------------------------------

def _require_done_job(job_id: str) -> dict:
    st = jobs.get_status(job_id)
    if st is None:
        raise HTTPException(status_code=404, detail="Job not found")
    if st.get("status") != "done":
        raise HTTPException(status_code=409, detail=f"Job status is '{st.get('status')}', not ready")
    return st


def _parquet_path(job_id: str, dataset: str) -> Path:
    if dataset not in ("clean", "rejected"):
        raise HTTPException(status_code=400, detail="dataset must be 'clean' or 'rejected'")
    p = jobs.job_dir(job_id) / f"{dataset}.parquet"
    if not p.exists():
        raise HTTPException(status_code=404, detail="Dataset not found for this job")
    return p


def _string_columns(con: duckdb.DuckDBPyConnection, parquet_path: Path) -> list[str]:
    desc = con.execute(f"DESCRIBE SELECT * FROM read_parquet('{parquet_path.as_posix()}')").fetchall()
    return [row[0] for row in desc if "VARCHAR" in row[1].upper()]


@app.get("/api/jobs/{job_id}/data", dependencies=[Depends(require_auth)])
def job_data(
    job_id: str,
    dataset: str = Query("clean"),
    category: str | None = Query(None),
    search: str | None = Query(None),
    email_filter: str = Query("all", pattern="^(all|has_email|no_email)$"),
    page: int = Query(1, ge=1),
    page_size: int = Query(50, ge=1, le=1000),
):
    st = _require_done_job(job_id)
    parquet_path = _parquet_path(job_id, dataset)
    cat_col = st.get("category_column")

    con = duckdb.connect()
    where_clauses = []
    params = []

    if category and cat_col:
        where_clauses.append(f'"{cat_col}" = ?')
        params.append(category)

    if email_filter != "all":
        email_col = _guess_email_col(con, parquet_path)
        if email_col:
            if email_filter == "has_email":
                where_clauses.append(f'"{email_col}" IS NOT NULL')
            else:  # no_email
                where_clauses.append(f'"{email_col}" IS NULL')

    if search:
        str_cols = _string_columns(con, parquet_path)
        if str_cols:
            like_clauses = " OR ".join(f'"{c}" ILIKE ?' for c in str_cols)
            where_clauses.append(f"({like_clauses})")
            params.extend([f"%{search}%"] * len(str_cols))

    where_sql = f"WHERE {' AND '.join(where_clauses)}" if where_clauses else ""

    total = con.execute(
        f"SELECT COUNT(*) FROM read_parquet('{parquet_path.as_posix()}') {where_sql}", params
    ).fetchone()[0]

    offset = (page - 1) * page_size
    rows_result = con.execute(
        f"SELECT * FROM read_parquet('{parquet_path.as_posix()}') {where_sql} "
        f"LIMIT {page_size} OFFSET {offset}",
        params,
    )
    columns = [d[0] for d in rows_result.description]
    rows = [dict(zip(columns, r)) for r in rows_result.fetchall()]

    return {
        "columns": columns,
        "rows": rows,
        "total": total,
        "page": page,
        "page_size": page_size,
    }


def _guess_email_col(con, parquet_path: Path) -> str | None:
    cols = [row[0] for row in con.execute(f"DESCRIBE SELECT * FROM read_parquet('{parquet_path.as_posix()}')").fetchall()]
    for c in cols:
        if "email" in c.lower():
            return c
    return None


@app.get("/api/jobs/{job_id}/categories", dependencies=[Depends(require_auth)])
def job_categories(job_id: str, dataset: str = Query("clean")):
    st = _require_done_job(job_id)
    cat_col = st.get("category_column")
    if not cat_col:
        return {"category_column": None, "categories": []}
    parquet_path = _parquet_path(job_id, dataset)
    con = duckdb.connect()
    rows = con.execute(
        f'SELECT COALESCE("{cat_col}", \'(uncategorized)\') AS cat, COUNT(*) AS n '
        f"FROM read_parquet('{parquet_path.as_posix()}') GROUP BY 1 ORDER BY n DESC"
    ).fetchall()
    return {"category_column": cat_col, "categories": [{"name": r[0], "count": r[1]} for r in rows]}


# ---------------------------------------------------------------------------
# Download (streamed, CSV-injection-sanitized)
# ---------------------------------------------------------------------------

DANGEROUS_PREFIXES = ("=", "+", "-", "@")


def _sanitized_select(con: duckdb.DuckDBPyConnection, parquet_path: Path, columns: list[str] | None = None) -> str:
    desc = con.execute(f"DESCRIBE SELECT * FROM read_parquet('{parquet_path.as_posix()}')").fetchall()
    all_cols = [(r[0], r[1]) for r in desc]
    if columns:
        wanted = set(columns)
        filtered = [c for c in all_cols if c[0] in wanted]
        if not filtered:
            raise HTTPException(status_code=400, detail="None of the requested columns exist in this dataset")
        all_cols = filtered
    exprs = []
    for name, dtype in all_cols:
        if "VARCHAR" in dtype.upper():
            # Prefix a leading tab (Excel/Sheets treat this as forcing text
            # and it visually disappears) before any of = + - @ to defuse
            # formula injection on open, without altering normal values.
            exprs.append(
                f"""CASE WHEN "{name}" IS NOT NULL AND regexp_matches("{name}", '^[=+\\-@]')
                    THEN '\t' || "{name}" ELSE "{name}" END AS "{name}\""""
            )
        else:
            exprs.append(f'"{name}"')
    return ", ".join(exprs)


@app.get("/api/jobs/{job_id}/download", dependencies=[Depends(require_auth)])
def download(
    job_id: str,
    dataset: str = Query("clean"),
    category: str | None = Query(None),
    email_filter: str = Query("all", pattern="^(all|has_email|no_email)$"),
    columns: list[str] | None = Query(None),
    background_tasks: BackgroundTasks = None,
):
    st = _require_done_job(job_id)
    parquet_path = _parquet_path(job_id, dataset)
    cat_col = st.get("category_column")

    con = duckdb.connect()
    where_clauses = []
    params = []
    if category and cat_col:
        where_clauses.append(f'"{cat_col}" = ?')
        params.append(category)
    if email_filter != "all":
        email_col = _guess_email_col(con, parquet_path)
        if email_col:
            if email_filter == "has_email":
                where_clauses.append(f'"{email_col}" IS NOT NULL')
            else:  # no_email
                where_clauses.append(f'"{email_col}" IS NULL')
    where_sql = f"WHERE {' AND '.join(where_clauses)}" if where_clauses else ""

    select_expr = _sanitized_select(con, parquet_path, columns=columns)

    tmp = tempfile.NamedTemporaryFile(suffix=".csv", delete=False)
    tmp.close()
    tmp_path = tmp.name

    query = (
        f"COPY (SELECT {select_expr} FROM read_parquet('{parquet_path.as_posix()}') {where_sql}) "
        f"TO '{tmp_path}' (HEADER, DELIMITER ',')"
    )
    if params:
        con.execute(query, params)
    else:
        con.execute(query)

    filename = f"leads_{dataset}"
    if category:
        filename += f"_{category}"
    filename += ".csv"

    def cleanup():
        try:
            os.unlink(tmp_path)
        except OSError:
            pass

    if background_tasks is not None:
        background_tasks.add_task(cleanup)

    return FileResponse(tmp_path, media_type="text/csv", filename=filename, background=background_tasks)


# ---------------------------------------------------------------------------
# Cleanup
# ---------------------------------------------------------------------------

@app.delete("/api/jobs/{job_id}", dependencies=[Depends(require_auth)])
def delete_job(job_id: str):
    if not jobs.job_exists(job_id):
        raise HTTPException(status_code=404, detail="Job not found")
    jobs.delete_job(job_id)
    return {"deleted": True}


@app.get("/api/health")
def health():
    return {"ok": True}
