"""SaaS user auth — password hashing + JWT with org context."""

from __future__ import annotations

import hashlib
import hmac
import secrets
from dataclasses import dataclass
from typing import Annotated, Any

from fastapi import Depends, Header, HTTPException, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer

from app.auth import create_access_token, decode_token
from app.saas import db as saas_db

bearer_scheme = HTTPBearer(auto_error=False)


def hash_password(password: str) -> str:
    # Prefer PBKDF2 to avoid passlib/bcrypt backend mismatches across versions
    salt = secrets.token_hex(16)
    digest = hashlib.pbkdf2_hmac("sha256", password.encode(), salt.encode(), 200_000).hex()
    return f"pbkdf2:{salt}:{digest}"


def verify_password(password: str, password_hash: str) -> bool:
    if password_hash.startswith("pbkdf2:"):
        _, salt, digest = password_hash.split(":", 2)
        check = hashlib.pbkdf2_hmac("sha256", password.encode(), salt.encode(), 200_000).hex()
        return hmac.compare_digest(check, digest)
    try:
        from passlib.context import CryptContext

        return CryptContext(schemes=["bcrypt"], deprecated="auto").verify(password, password_hash)
    except Exception:
        return False


def hash_api_key(raw_key: str) -> str:
    return hashlib.sha256(raw_key.encode()).hexdigest()


def generate_api_key() -> tuple[str, str, str]:
    """Return (full_key, prefix, hash)."""
    raw = f"sk_live_{secrets.token_urlsafe(32)}"
    prefix = raw[:16]
    return raw, prefix, hash_api_key(raw)


@dataclass
class AuthContext:
    user_id: str
    email: str
    full_name: str | None
    org_id: str
    org_name: str
    org_slug: str
    role: str
    kb_mode: str
    plan: str
    auth_via: str  # jwt | api_key


def _user_token_payload(user: dict[str, Any], org: dict[str, Any], role: str) -> dict[str, Any]:
    # Embed org fields so /saas/me auth can skip remote DB round-trips.
    return {
        "sub": user["id"],
        "email": user["email"],
        "full_name": user.get("full_name"),
        "org_id": org["id"],
        "org_name": org["name"],
        "org_slug": org.get("slug") or "",
        "role": role,
        "kb_mode": org.get("kb_mode") or "tenant",
        "plan": org.get("plan") or "free",
        "typ": "user",
    }


def issue_user_token(user: dict[str, Any], org: dict[str, Any], role: str) -> str:
    return create_access_token(_user_token_payload(user, org, role))


def get_auth_context(
    credentials: Annotated[HTTPAuthorizationCredentials | None, Depends(bearer_scheme)],
    x_api_key: Annotated[str | None, Header(alias="X-API-Key")] = None,
) -> AuthContext:
    """Require either user JWT or org API key."""
    if x_api_key:
        return _auth_from_api_key(x_api_key)
    if credentials:
        return _auth_from_jwt(credentials.credentials)
    raise HTTPException(
        status_code=status.HTTP_401_UNAUTHORIZED,
        detail="Authentication required",
        headers={"WWW-Authenticate": "Bearer"},
    )


def _auth_from_jwt(token: str) -> AuthContext:
    try:
        payload = decode_token(token)
    except Exception as exc:
        raise HTTPException(status_code=401, detail="Invalid or expired token") from exc

    if payload.get("typ") == "admin" or not payload.get("org_id"):
        # Reject bare admin tokens for SaaS routes
        if payload.get("typ") != "user":
            raise HTTPException(status_code=401, detail="User token required")

    user_id = str(payload.get("sub") or "")
    org_id = str(payload.get("org_id") or "")
    if not user_id or not org_id:
        raise HTTPException(status_code=401, detail="Invalid session")

    # Fast path: rich JWT from guest/login — no remote Postgres round-trips.
    if payload.get("email") and payload.get("org_name") is not None:
        return AuthContext(
            user_id=user_id,
            email=str(payload.get("email")),
            full_name=payload.get("full_name"),
            org_id=org_id,
            org_name=str(payload.get("org_name") or ""),
            org_slug=str(payload.get("org_slug") or ""),
            role=str(payload.get("role") or "member"),
            kb_mode=str(payload.get("kb_mode") or "tenant"),
            plan=str(payload.get("plan") or "free"),
            auth_via="jwt",
        )

    # Legacy tokens: one pooled connection for user + org + membership.
    from app.database import get_connection

    with get_connection() as conn:
        user = conn.execute("SELECT * FROM users WHERE id = ?", (user_id,)).fetchone()
        org = conn.execute("SELECT * FROM orgs WHERE id = ?", (org_id,)).fetchone()
        membership = conn.execute(
            "SELECT * FROM memberships WHERE user_id = ? AND org_id = ?",
            (user_id, org_id),
        ).fetchone()
    if not user or not org or not membership:
        raise HTTPException(status_code=401, detail="Invalid session")
    user_d, org_d, mem_d = dict(user), dict(org), dict(membership)
    return AuthContext(
        user_id=user_d["id"],
        email=user_d["email"],
        full_name=user_d.get("full_name"),
        org_id=org_d["id"],
        org_name=org_d["name"],
        org_slug=org_d.get("slug") or "",
        role=mem_d["role"],
        kb_mode=org_d.get("kb_mode") or "tenant",
        plan=org_d.get("plan") or "free",
        auth_via="jwt",
    )


def _auth_from_api_key(raw_key: str) -> AuthContext:
    prefix = raw_key[:16]
    row = saas_db.get_api_key_by_prefix(prefix)
    if not row or not hmac.compare_digest(row["key_hash"], hash_api_key(raw_key)):
        raise HTTPException(status_code=401, detail="Invalid API key")
    org = saas_db.get_org(row["org_id"])
    if not org:
        raise HTTPException(status_code=401, detail="Invalid API key org")
    saas_db.touch_api_key(row["id"])
    return AuthContext(
        user_id=row.get("created_by") or "api-key",
        email="api-key@" + org["slug"],
        full_name=row.get("name"),
        org_id=org["id"],
        org_name=org["name"],
        org_slug=org.get("slug") or "",
        role="api",
        kb_mode=org.get("kb_mode") or "tenant",
        plan=org.get("plan") or "free",
        auth_via="api_key",
    )
