from __future__ import annotations

import logging
import os
import secrets
import sqlite3
import time
from contextlib import contextmanager
from dataclasses import dataclass, field
from typing import Optional

from werkzeug.security import check_password_hash, generate_password_hash

logger = logging.getLogger("apex.rbac")

# ── DB path ──────────────────────────────────────────────────────────────────

_DB_PATH: str = ""

ADMIN_EMAIL = "admin@rptechindia.com"

# Permissions seeded at startup
DEFAULT_PERMISSIONS: list[tuple[str, str, str]] = [
    # (name, module, action)
    # Sales & Workspace
    ("sales.view",             "sales",     "view"),
    # Projects
    ("projects.view",          "projects",  "view"),
    ("projects.edit",          "projects",  "edit"),      # own department only
    ("projects.edit_any",      "projects",  "edit_any"),  # any department
    # SAP Reports
    ("reports.sap_hana",       "reports",   "sap_hana"),
    ("reports.sap_stock",      "reports",   "sap_stock"),
    ("reports.sap_po",         "reports",   "sap_po"),
    ("reports.sap_wos",        "reports",   "sap_wos"),
    ("reports.sap_aging",      "reports",   "sap_aging"),
    ("reports.sap_predictive", "reports",   "sap_predictive"),
    # ML Reports
    ("reports.ml_forecast",    "reports",   "ml_forecast"),
    ("reports.ml_demand",      "reports",   "ml_demand"),
    ("reports.ml_anomalies",   "reports",   "ml_anomalies"),
    ("reports.ml_segments",    "reports",   "ml_segments"),
    ("reports.ml_insights",    "reports",   "ml_insights"),
    # Inventory Hub
    ("inventory.view",         "inventory", "view"),
    # Service / AI Assistant
    ("service.view",           "service",   "view"),
]

# Permissions automatically granted to manager role (all of them)
MANAGER_PERMISSIONS = {p[0] for p in DEFAULT_PERMISSIONS}

# Set-password token TTL (24 hours)
TOKEN_TTL = 86_400


# ── Schema ───────────────────────────────────────────────────────────────────

_SCHEMA = """
CREATE TABLE IF NOT EXISTS users (
    id                         INTEGER PRIMARY KEY AUTOINCREMENT,
    email                      TEXT    UNIQUE NOT NULL,
    name                       TEXT    NOT NULL DEFAULT '',
    department                 TEXT    DEFAULT '',
    role                       TEXT    NOT NULL DEFAULT 'employee',
    status                     TEXT    NOT NULL DEFAULT 'pending',
    password_hash              TEXT    DEFAULT NULL,
    set_password_token         TEXT    DEFAULT NULL,
    set_password_token_expiry  REAL    DEFAULT NULL,
    approved_by                TEXT    DEFAULT NULL,
    approved_at                REAL    DEFAULT NULL,
    created_at                 REAL    NOT NULL
);

CREATE TABLE IF NOT EXISTS permissions (
    id          INTEGER PRIMARY KEY AUTOINCREMENT,
    name        TEXT    UNIQUE NOT NULL,
    module      TEXT    NOT NULL,
    action      TEXT    NOT NULL
);

CREATE TABLE IF NOT EXISTS user_permissions (
    user_id       INTEGER NOT NULL,
    permission_id INTEGER NOT NULL,
    granted_by    TEXT    NOT NULL,
    granted_at    REAL    NOT NULL,
    PRIMARY KEY (user_id, permission_id),
    FOREIGN KEY (user_id)       REFERENCES users(id)       ON DELETE CASCADE,
    FOREIGN KEY (permission_id) REFERENCES permissions(id) ON DELETE CASCADE
);

CREATE TABLE IF NOT EXISTS audit_log (
    id           INTEGER PRIMARY KEY AUTOINCREMENT,
    action       TEXT    NOT NULL,
    actor_email  TEXT    NOT NULL,
    target_email TEXT    DEFAULT NULL,
    detail       TEXT    DEFAULT '',
    created_at   REAL    NOT NULL
);

CREATE TABLE IF NOT EXISTS user_brands (
    user_id      INTEGER NOT NULL,
    brand_code   TEXT    NOT NULL,
    assigned_by  TEXT    NOT NULL DEFAULT '',
    assigned_at  REAL    NOT NULL,
    PRIMARY KEY (user_id, brand_code),
    FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);

CREATE TABLE IF NOT EXISTS user_requested_brands (
    user_id      INTEGER NOT NULL,
    brand_code   TEXT    NOT NULL,
    requested_at REAL    NOT NULL,
    PRIMARY KEY (user_id, brand_code),
    FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
"""


# ── Connection helper ─────────────────────────────────────────────────────────

@contextmanager
def _conn():
    con = sqlite3.connect(_DB_PATH)
    con.row_factory = sqlite3.Row
    con.execute("PRAGMA foreign_keys = ON")
    try:
        yield con
        con.commit()
    except Exception:
        con.rollback()
        raise
    finally:
        con.close()


# ── Init ─────────────────────────────────────────────────────────────────────

def init_rbac(db_path: str) -> None:
    """Create tables, seed permissions, and ensure admin user exists."""
    global _DB_PATH
    _DB_PATH = db_path
    os.makedirs(os.path.dirname(db_path), exist_ok=True)

    with _conn() as con:
        con.executescript(_SCHEMA)

        # Migration: recreate user_brands if brand_code column is missing
        cols = {r[1] for r in con.execute("PRAGMA table_info(user_brands)").fetchall()}
        if "brand_code" not in cols:
            con.execute("DROP TABLE IF EXISTS user_brands")
            con.execute("""CREATE TABLE user_brands (
                user_id      INTEGER NOT NULL,
                brand_code   TEXT    NOT NULL,
                assigned_by  TEXT    NOT NULL DEFAULT '',
                assigned_at  REAL    NOT NULL,
                PRIMARY KEY (user_id, brand_code),
                FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
            )""")
            logger.info("[RBAC] Migrated user_brands: added brand_code column")

        # Seed permissions
        for name, module, action in DEFAULT_PERMISSIONS:
            con.execute(
                "INSERT OR IGNORE INTO permissions (name, module, action) VALUES (?,?,?)",
                (name, module, action),
            )
        # Prune stale permissions not in DEFAULT_PERMISSIONS (cascades via FK to user_permissions)
        _valid = [p[0] for p in DEFAULT_PERMISSIONS]
        _ph = ",".join("?" * len(_valid))
        con.execute(f"DELETE FROM permissions WHERE name NOT IN ({_ph})", _valid)
        # Ensure admin user exists
        exists = con.execute("SELECT id FROM users WHERE email=?", (ADMIN_EMAIL,)).fetchone()
        if not exists:
            con.execute(
                """INSERT INTO users (email, name, role, status, created_at)
                   VALUES (?, 'Admin', 'admin', 'active', ?)""",
                (ADMIN_EMAIL, time.time()),
            )
            logger.info("[RBAC] Admin user created: %s", ADMIN_EMAIL)

    logger.info("[RBAC] Initialized at %s", db_path)


# ── Dataclasses ───────────────────────────────────────────────────────────────

@dataclass
class UserRecord:
    id: int
    email: str
    name: str
    department: str
    role: str
    status: str
    approved_by: Optional[str]
    approved_at: Optional[float]
    created_at: float
    has_password: bool = False


@dataclass
class PermissionRecord:
    id: int
    name: str
    module: str
    action: str


# ── User operations ──────────────────────────────────────────────────────────

def create_user(email: str, name: str, department: str = "", password: str = "") -> tuple[bool, str]:
    email = email.strip().lower()
    try:
        pw_hash = generate_password_hash(password) if password else None
        with _conn() as con:
            con.execute(
                """INSERT INTO users (email, name, department, role, status, password_hash, created_at)
                   VALUES (?, ?, ?, 'employee', 'active', ?, ?)""",
                (email, name.strip(), department.strip(), pw_hash, time.time()),
            )
        _audit("user.signup", email, email, f"New signup: {name}")
        logger.info("[RBAC] New signup: %s", email)
        return True, "Account created successfully."
    except sqlite3.IntegrityError:
        return False, "An account with this email already exists."
    except Exception as exc:
        logger.error("[RBAC] create_user failed: %s", exc)
        return False, "Registration failed. Please try again."


def get_user(email: str) -> Optional[UserRecord]:
    email = email.strip().lower()
    with _conn() as con:
        row = con.execute("SELECT * FROM users WHERE email=?", (email,)).fetchone()
    if not row:
        return None
    return _row_to_user(row)


def get_user_by_id(user_id: int) -> Optional[UserRecord]:
    with _conn() as con:
        row = con.execute("SELECT * FROM users WHERE id=?", (user_id,)).fetchone()
    return _row_to_user(row) if row else None


def get_all_users(status: str = "") -> list[UserRecord]:
    with _conn() as con:
        if status:
            rows = con.execute("SELECT * FROM users WHERE status=? ORDER BY created_at DESC", (status,)).fetchall()
        else:
            rows = con.execute("SELECT * FROM users ORDER BY created_at DESC").fetchall()
    return [_row_to_user(r) for r in rows]


def _row_to_user(row) -> UserRecord:
    return UserRecord(
        id=row["id"], email=row["email"], name=row["name"],
        department=row["department"] or "", role=row["role"],
        status=row["status"], approved_by=row["approved_by"],
        approved_at=row["approved_at"], created_at=row["created_at"],
        has_password=bool(row["password_hash"]),
    )


# ── Approval workflow ─────────────────────────────────────────────────────────

def approve_user(user_id: int, admin_email: str, permission_names: list[str]) -> tuple[bool, str, str]:
    """Approve user, assign permissions, generate set-password token.
    Returns (ok, message, set_password_token).
    """
    user = get_user_by_id(user_id)
    if not user:
        return False, "User not found.", ""
    if user.status == "active":
        return False, "User is already active.", ""

    token = secrets.token_urlsafe(32)
    expiry = time.time() + TOKEN_TTL

    with _conn() as con:
        con.execute(
            """UPDATE users SET status='active', approved_by=?, approved_at=?,
               set_password_token=?, set_password_token_expiry=? WHERE id=?""",
            (admin_email, time.time(), token, expiry, user_id),
        )
        # Assign permissions
        for pname in permission_names:
            perm = con.execute("SELECT id FROM permissions WHERE name=?", (pname,)).fetchone()
            if perm:
                con.execute(
                    """INSERT OR IGNORE INTO user_permissions (user_id, permission_id, granted_by, granted_at)
                       VALUES (?,?,?,?)""",
                    (user_id, perm["id"], admin_email, time.time()),
                )

    _audit("user.approved", admin_email, user.email,
           f"Approved with {len(permission_names)} permission(s): {', '.join(permission_names)}")
    return True, "User approved.", token



def resend_set_password_token(user_id: int, admin_email: str) -> tuple[bool, str, str]:
    """Generate a fresh set-password token for an active user who hasn't logged in yet."""
    user = get_user_by_id(user_id)
    if not user:
        return False, "User not found.", ""
    if user.status != "active":
        return False, "User is not active.", ""
    if user.has_password:
        return False, "User has already set their password.", ""
    token = secrets.token_urlsafe(32)
    expiry = time.time() + TOKEN_TTL
    with _conn() as con:
        con.execute(
            "UPDATE users SET set_password_token=?, set_password_token_expiry=? WHERE id=?",
            (token, expiry, user_id),
        )
    _audit("user.resend_token", admin_email, user.email, "Resent set-password email")
    return True, "Token generated.", token

def reject_user(user_id: int, admin_email: str, reason: str = "") -> tuple[bool, str]:
    user = get_user_by_id(user_id)
    if not user:
        return False, "User not found."
    with _conn() as con:
        con.execute("UPDATE users SET status='rejected' WHERE id=?", (user_id,))
    _audit("user.rejected", admin_email, user.email, reason or "No reason given")
    return True, "User rejected."


def get_user_by_set_password_token(token: str) -> dict | None:
    """Look up name/expiry for a set-password token (used by GET page)."""
    if not token:
        return None
    with _conn() as con:
        row = con.execute(
            "SELECT name, set_password_token_expiry FROM users WHERE set_password_token=?",
            (token,),
        ).fetchone()
    return dict(row) if row else None


def set_password_via_token(token: str, new_password: str) -> tuple[bool, str]:
    """Set a user's password using the approval token."""
    with _conn() as con:
        row = con.execute(
            "SELECT * FROM users WHERE set_password_token=?", (token,)
        ).fetchone()
        if not row:
            return False, "Invalid or expired link."
        if time.time() > (row["set_password_token_expiry"] or 0):
            return False, "This link has expired. Please contact admin."
        if row["status"] != "active":
            return False, "Account is not active."
        pw_hash = generate_password_hash(new_password)
        con.execute(
            "UPDATE users SET password_hash=?, set_password_token=NULL, set_password_token_expiry=NULL WHERE id=?",
            (pw_hash, row["id"]),
        )
    _audit("user.set_password", row["email"], row["email"], "Password set after approval")
    return True, "Password set successfully. You can now log in."


def set_admin_password(admin_email: str, new_password: str) -> None:
    """Called on first boot to set admin password if not set."""
    with _conn() as con:
        row = con.execute("SELECT password_hash FROM users WHERE email=?", (admin_email,)).fetchone()
        if row and not row["password_hash"]:
            pw_hash = generate_password_hash(new_password)
            con.execute("UPDATE users SET password_hash=? WHERE email=?", (pw_hash, admin_email))
            logger.info("[RBAC] Admin password initialized")


# ── Authentication ────────────────────────────────────────────────────────────

def authenticate(email: str, password: str) -> tuple[bool, str, Optional[UserRecord]]:
    """Verify email/password. Returns (ok, message, user)."""
    email = email.strip().lower()
    with _conn() as con:
        row = con.execute("SELECT * FROM users WHERE email=?", (email,)).fetchone()
    if not row:
        return False, "Invalid email or password.", None
    if row["status"] == "pending":
        return False, "Your account is pending admin approval.", None
    if row["status"] == "rejected":
        return False, "Your access request was not approved. Contact admin.", None
    if row["status"] != "active":
        return False, "Account is inactive.", None
    if not row["password_hash"]:
        return False, "Password not set. Check your approval email for the setup link.", None
    if not check_password_hash(row["password_hash"], password):
        return False, "Invalid email or password.", None
    return True, "OK", _row_to_user(row)


# ── Permission operations ─────────────────────────────────────────────────────

def get_all_permissions() -> list[PermissionRecord]:
    with _conn() as con:
        rows = con.execute("SELECT * FROM permissions ORDER BY module, action").fetchall()
    return [PermissionRecord(id=r["id"], name=r["name"], module=r["module"], action=r["action"]) for r in rows]


def get_user_permissions(user_id: int) -> list[str]:
    with _conn() as con:
        rows = con.execute(
            """SELECT p.name FROM permissions p
               JOIN user_permissions up ON up.permission_id = p.id
               WHERE up.user_id = ?""",
            (user_id,),
        ).fetchall()
    return [r["name"] for r in rows]


def assign_permissions(user_id: int, permission_names: list[str], granted_by: str) -> int:
    """Assign permissions to a user. Returns count of new grants."""
    count = 0
    with _conn() as con:
        for pname in permission_names:
            perm = con.execute("SELECT id FROM permissions WHERE name=?", (pname,)).fetchone()
            if perm:
                try:
                    con.execute(
                        """INSERT OR IGNORE INTO user_permissions (user_id, permission_id, granted_by, granted_at)
                           VALUES (?,?,?,?)""",
                        (user_id, perm["id"], granted_by, time.time()),
                    )
                    count += 1
                except Exception:
                    pass
    if count:
        user = get_user_by_id(user_id)
        _audit("permission.assign", granted_by, user.email if user else str(user_id),
               f"Assigned: {', '.join(permission_names)}")
    return count


def revoke_permission(user_id: int, permission_name: str, revoked_by: str) -> bool:
    with _conn() as con:
        perm = con.execute("SELECT id FROM permissions WHERE name=?", (permission_name,)).fetchone()
        if not perm:
            return False
        con.execute(
            "DELETE FROM user_permissions WHERE user_id=? AND permission_id=?",
            (user_id, perm["id"]),
        )
    user = get_user_by_id(user_id)
    _audit("permission.revoke", revoked_by, user.email if user else str(user_id),
           f"Revoked: {permission_name}")
    return True


def revoke_all_permissions(user_id: int, revoked_by: str) -> None:
    with _conn() as con:
        con.execute("DELETE FROM user_permissions WHERE user_id=?", (user_id,))
    user = get_user_by_id(user_id)
    _audit("permission.revoke_all", revoked_by, user.email if user else str(user_id), "All permissions revoked")


# ── Permission check ──────────────────────────────────────────────────────────

def has_permission(user_email: str, permission_name: str) -> bool:
    """Check if a user has a specific permission.

    - admin: always True
    - manager: True for all non-admin_panel permissions
    - employee: only if explicitly assigned
    """
    email = (user_email or "").strip().lower()
    with _conn() as con:
        row = con.execute("SELECT id, role, status FROM users WHERE email=?", (email,)).fetchone()
    if not row or row["status"] != "active":
        return False
    role = row["role"]
    if role == "admin":
        return True
    if role == "manager":
        return permission_name in MANAGER_PERMISSIONS
    # employee — check explicit assignment
    with _conn() as con:
        found = con.execute(
            """SELECT 1 FROM user_permissions up
               JOIN permissions p ON p.id = up.permission_id
               WHERE up.user_id=? AND p.name=?""",
            (row["id"], permission_name),
        ).fetchone()
    return found is not None


def get_user_effective_permissions(user_email: str) -> list[str]:
    """Return the full effective permission list for a user."""
    email = (user_email or "").strip().lower()
    with _conn() as con:
        row = con.execute("SELECT id, role, status FROM users WHERE email=?", (email,)).fetchone()
    if not row or row["status"] != "active":
        return []
    role = row["role"]
    if role == "admin":
        return [p[0] for p in DEFAULT_PERMISSIONS]
    if role == "manager":
        return sorted(MANAGER_PERMISSIONS)
    return get_user_permissions(row["id"])


# ── Audit log ─────────────────────────────────────────────────────────────────

def _audit(action: str, actor: str, target: str = "", detail: str = "") -> None:
    try:
        with _conn() as con:
            con.execute(
                "INSERT INTO audit_log (action, actor_email, target_email, detail, created_at) VALUES (?,?,?,?,?)",
                (action, actor, target, detail, time.time()),
            )
    except Exception as exc:
        logger.warning("[RBAC] Audit log failed: %s", exc)


def get_audit_log(limit: int = 100) -> list[dict]:
    with _conn() as con:
        rows = con.execute(
            "SELECT * FROM audit_log ORDER BY created_at DESC LIMIT ?", (limit,)
        ).fetchall()
    return [dict(r) for r in rows]


# ── Permissions grouped by module ─────────────────────────────────────────────

def permissions_by_module() -> dict[str, list[str]]:
    result: dict[str, list[str]] = {}
    for name, module, _ in DEFAULT_PERMISSIONS:
        result.setdefault(module, []).append(name)
    return result


# ── Admin-initiated user creation ────────────────────────────────────────────

def create_user_by_admin(
    email: str, full_name: str, department: str, role: str, admin_email: str
) -> tuple[bool, str, str]:
    """Admin creates a user directly — status is active, token generated for password setup."""
    email = email.strip().lower()
    token = secrets.token_urlsafe(32)
    expiry = time.time() + TOKEN_TTL
    try:
        with _conn() as con:
            con.execute(
                """INSERT INTO users (email, name, department, role, status,
                       set_password_token, set_password_token_expiry, approved_by, approved_at, created_at)
                   VALUES (?, ?, ?, ?, 'active', ?, ?, ?, ?, ?)""",
                (email, full_name.strip(), department.strip(), role,
                 token, expiry, admin_email, time.time(), time.time()),
            )
        _audit("user.create", admin_email, email, f"Created by admin, role={role}")
        logger.info("[RBAC] Admin created user: %s (role=%s)", email, role)
        return True, "User created.", token
    except sqlite3.IntegrityError:
        return False, "An account with this email already exists.", ""
    except Exception as exc:
        logger.error("[RBAC] create_user_by_admin failed: %s", exc)
        return False, "User creation failed.", ""


def set_user_status(user_id: int, new_status: str, admin_email: str) -> bool:
    """Set arbitrary status on a user (e.g. 'inactive')."""
    try:
        with _conn() as con:
            cur = con.execute(
                "UPDATE users SET status=? WHERE id=?", (new_status, user_id)
            )
        if cur.rowcount:
            _audit("user.status_change", admin_email, str(user_id), f"status→{new_status}")
            return True
        return False
    except Exception as exc:
        logger.error("[RBAC] set_user_status failed: %s", exc)
        return False



def set_user_role(user_id: int, new_role: str, admin_email: str) -> bool:
    """Change a user's role (employee / manager / admin)."""
    try:
        with _conn() as con:
            cur = con.execute("UPDATE users SET role=? WHERE id=?", (new_role, user_id))
        if cur.rowcount:
            _audit("user.role_change", admin_email, str(user_id), f"role→{new_role}")
            return True
        return False
    except Exception as exc:
        logger.error("[RBAC] set_user_role failed: %s", exc)
        return False


# ── Brand access ─────────────────────────────────────────────────────────────

def get_user_brands_by_email(email: str) -> list[str]:
    """Return brand codes assigned to a user. Empty list means no restriction."""
    email = (email or "").strip().lower()
    with _conn() as con:
        row = con.execute("SELECT id FROM users WHERE email=?", (email,)).fetchone()
        if not row:
            return []
        rows = con.execute(
            "SELECT brand_code FROM user_brands WHERE user_id=?", (row["id"],)
        ).fetchall()
    return [r["brand_code"] for r in rows]


def get_user_brands(user_id: int) -> list[str]:
    """Return brand codes assigned to a user by ID."""
    with _conn() as con:
        rows = con.execute(
            "SELECT brand_code FROM user_brands WHERE user_id=?", (user_id,)
        ).fetchall()
    return [r["brand_code"] for r in rows]


def assign_brands(user_id: int, brand_codes: list[str], assigned_by: str) -> None:
    """Replace a user's brand assignments entirely."""
    with _conn() as con:
        con.execute("DELETE FROM user_brands WHERE user_id=?", (user_id,))
        for code in brand_codes:
            con.execute(
                "INSERT OR IGNORE INTO user_brands (user_id, brand_code, assigned_by, assigned_at) VALUES (?,?,?,?)",
                (user_id, code.strip(), assigned_by, time.time()),
            )
    user = get_user_by_id(user_id)
    _audit("brands.assign", assigned_by, user.email if user else str(user_id),
           f"Assigned {len(brand_codes)} brand(s)")


def save_requested_brands(user_id: int, brand_codes: list[str]) -> None:
    """Save the brands a user requested at signup."""
    with _conn() as con:
        con.execute("DELETE FROM user_requested_brands WHERE user_id=?", (user_id,))
        for code in brand_codes:
            con.execute(
                "INSERT OR IGNORE INTO user_requested_brands (user_id, brand_code, requested_at) VALUES (?,?,?)",
                (user_id, code.strip(), time.time()),
            )


def get_requested_brands(user_id: int) -> list[str]:
    """Return brand codes the user requested at signup."""
    with _conn() as con:
        rows = con.execute(
            "SELECT brand_code FROM user_requested_brands WHERE user_id=?", (user_id,)
        ).fetchall()
    return [r["brand_code"] for r in rows]