"""Database layer — Postgres when DATABASE_URL is set, otherwise local SQLite.

Supports feedback, audit logs, sessions, and SaaS tables via a thin
sqlite-compatible wrapper (`?` placeholders, dict rows, executescript).
"""

from __future__ import annotations

import re
import sqlite3
import uuid
from datetime import datetime
from pathlib import Path
from typing import Any, Iterable

from app.config import get_settings

_DB_PATH: Path | None = None


def is_postgres() -> bool:
    url = (get_settings().database_url or "").strip()
    return url.startswith("postgres://") or url.startswith("postgresql://")


def _db_path() -> Path:
    global _DB_PATH
    if _DB_PATH is None:
        settings = get_settings()
        _DB_PATH = settings.chroma_persist_dir.parent / "chatbot.db"
        _DB_PATH.parent.mkdir(parents=True, exist_ok=True)
    return _DB_PATH


def _adapt_sql(sql: str) -> str:
    """Convert SQLite-ish SQL to the active dialect."""
    if not is_postgres():
        return sql
    # INSERT OR IGNORE → ON CONFLICT DO NOTHING (simple single-table inserts)
    if re.match(r"(?is)^\s*INSERT\s+OR\s+IGNORE\s+INTO\s+", sql):
        sql = re.sub(r"(?is)^\s*INSERT\s+OR\s+IGNORE\s+INTO\s+", "INSERT INTO ", sql)
        if "ON CONFLICT" not in sql.upper():
            sql = sql.rstrip().rstrip(";") + " ON CONFLICT DO NOTHING"
    # Escape literal % (e.g. LIKE '...%') before ? → %s, or psycopg treats %' as a placeholder.
    sql = sql.replace("%", "%%")
    sql = re.sub(r"\?", "%s", sql)
    return sql


class _PgCursor:
    def __init__(self, cur: Any):
        self._cur = cur
        self.lastrowid: int | None = None
        self.rowcount = getattr(cur, "rowcount", -1)

    def fetchone(self):
        return self._cur.fetchone()

    def fetchall(self):
        return self._cur.fetchall()


class _PgConnection:
    """Minimal sqlite3.Connection-compatible wrapper over psycopg."""

    def __init__(self, conn: Any, *, pool: Any = None):
        self._conn = conn
        self._pool = pool

    def execute(self, sql: str, params: Iterable[Any] | None = None):
        adapted = _adapt_sql(sql)
        params = tuple(params) if params is not None else ()
        # Auto RETURNING id for plain INSERTs so callers can use lastrowid
        returning = False
        if (
            re.match(r"(?is)^\s*INSERT\s+INTO\s+\w+", adapted)
            and "RETURNING" not in adapted.upper()
            and "ON CONFLICT DO NOTHING" not in adapted.upper()
        ):
            # Only append when the table likely has a serial id (audit_log / feedback)
            table_m = re.match(r"(?is)^\s*INSERT\s+INTO\s+(\w+)", adapted)
            table = (table_m.group(1).lower() if table_m else "")
            if table in {"audit_log", "feedback"}:
                adapted = adapted.rstrip().rstrip(";") + " RETURNING id"
                returning = True
        cur = self._conn.execute(adapted, params)
        wrapped = _PgCursor(cur)
        if returning:
            row = cur.fetchone()
            if row:
                wrapped.lastrowid = int(row["id"] if isinstance(row, dict) else row[0])
        wrapped.rowcount = cur.rowcount
        return wrapped

    def executescript(self, script: str) -> None:
        # Strip SQL line comments and split on semicolons
        cleaned = re.sub(r"--.*?$", "", script, flags=re.M)
        for stmt in cleaned.split(";"):
            stmt = stmt.strip()
            if stmt:
                self.execute(stmt)

    def commit(self) -> None:
        self._conn.commit()

    def close(self) -> None:
        if self._pool is not None:
            self._pool.putconn(self._conn)
            return
        self._conn.close()

    def __enter__(self) -> _PgConnection:
        return self

    def __exit__(self, exc_type, exc, tb) -> None:
        try:
            if exc_type is None:
                self._conn.commit()
            else:
                self._conn.rollback()
        finally:
            self.close()


_pg_pool = None
_pg_pool_lock = __import__("threading").Lock()


def _normalize_pg_url(url: str) -> str:
    url = (url or "").strip()
    if url.startswith("postgres://"):
        return "postgresql://" + url[len("postgres://") :]
    return url


def _get_pg_pool():
    """Lazy singleton pool — avoids a new TCP handshake on every query."""
    global _pg_pool
    if _pg_pool is not None:
        return _pg_pool
    with _pg_pool_lock:
        if _pg_pool is not None:
            return _pg_pool
        from psycopg.rows import dict_row
        from psycopg_pool import ConnectionPool

        url = _normalize_pg_url(get_settings().database_url)
        _pg_pool = ConnectionPool(
            conninfo=url,
            min_size=1,
            max_size=8,
            timeout=30,
            kwargs={"row_factory": dict_row, "connect_timeout": 10},
            open=True,
        )
        return _pg_pool


def warm_pg_pool() -> None:
    """Open the pool early so the first UI request is not paying connect cost alone."""
    if not is_postgres():
        return
    pool = _get_pg_pool()
    conn = pool.getconn()
    try:
        conn.execute("SELECT 1")
        conn.commit()
    except Exception:
        conn.rollback()
        raise
    finally:
        pool.putconn(conn)


def get_connection():
    """Return a connection (Postgres or SQLite) usable as a context manager."""
    if is_postgres():
        pool = _get_pg_pool()
        raw = pool.getconn()
        return _PgConnection(raw, pool=pool)

    conn = sqlite3.connect(str(_db_path()))
    conn.row_factory = sqlite3.Row
    conn.execute("PRAGMA journal_mode=WAL")
    return conn


def table_columns(conn, table: str) -> set[str]:
    if is_postgres():
        rows = conn.execute(
            """
            SELECT column_name AS name
            FROM information_schema.columns
            WHERE table_schema = 'public' AND lower(table_name) = lower(?)
            """,
            (table,),
        ).fetchall()
        out: set[str] = set()
        for r in rows:
            name = r["name"] if isinstance(r, dict) else r[0]
            if name:
                out.add(str(name))
        return out
    rows = conn.execute(f"PRAGMA table_info({table})").fetchall()
    return {r[1] for r in rows}


def init_db() -> None:
    pg = is_postgres()
    id_type = "BIGSERIAL PRIMARY KEY" if pg else "INTEGER PRIMARY KEY AUTOINCREMENT"
    with get_connection() as conn:
        conn.executescript(
            f"""
            CREATE TABLE IF NOT EXISTS audit_log (
                id          {id_type},
                session_id  TEXT,
                question    TEXT NOT NULL,
                answer      TEXT NOT NULL,
                citations   TEXT,
                duration_ms INTEGER,
                created_at  TEXT NOT NULL
            );

            CREATE TABLE IF NOT EXISTS feedback (
                id          {id_type},
                log_id      INTEGER,
                session_id  TEXT,
                question    TEXT NOT NULL,
                answer      TEXT NOT NULL,
                rating      TEXT NOT NULL CHECK(rating IN ('up','down')),
                correction  TEXT,
                reviewed    INTEGER DEFAULT 0,
                created_at  TEXT NOT NULL,
                org_id      TEXT,
                agent_id    TEXT,
                user_id     TEXT,
                message_id  TEXT,
                mode        TEXT,
                model_name  TEXT,
                updated_at  TEXT
            );

            CREATE TABLE IF NOT EXISTS sessions (
                id          TEXT PRIMARY KEY,
                title       TEXT,
                created_at  TEXT NOT NULL,
                last_seen   TEXT NOT NULL,
                message_count INTEGER DEFAULT 0
            );

            CREATE INDEX IF NOT EXISTS idx_audit_session ON audit_log(session_id);
            CREATE INDEX IF NOT EXISTS idx_feedback_rating ON feedback(rating);
            CREATE INDEX IF NOT EXISTS idx_feedback_reviewed ON feedback(reviewed);
            CREATE INDEX IF NOT EXISTS idx_sessions_last_seen ON sessions(last_seen);
            CREATE INDEX IF NOT EXISTS idx_feedback_org_agent ON feedback(org_id, agent_id);
            CREATE INDEX IF NOT EXISTS idx_feedback_message ON feedback(message_id);
            """
        )

        # Additive migrations for older SQLite DBs / partial Postgres creates
        fb_cols = table_columns(conn, "feedback")
        for col in (
            "org_id",
            "agent_id",
            "user_id",
            "message_id",
            "mode",
            "model_name",
            "updated_at",
        ):
            if col not in fb_cols:
                conn.execute(f"ALTER TABLE feedback ADD COLUMN {col} TEXT")

        sess_cols = table_columns(conn, "sessions")
        if "title" not in sess_cols:
            conn.execute("ALTER TABLE sessions ADD COLUMN title TEXT")

        # Unique message_id when present
        try:
            if pg:
                conn.execute(
                    """
                    CREATE UNIQUE INDEX IF NOT EXISTS idx_feedback_message_unique
                    ON feedback(message_id)
                    WHERE message_id IS NOT NULL AND message_id <> ''
                    """
                )
            else:
                conn.execute(
                    """
                    CREATE UNIQUE INDEX IF NOT EXISTS idx_feedback_message_unique
                    ON feedback(message_id)
                    WHERE message_id IS NOT NULL AND message_id != ''
                    """
                )
        except Exception:
            pass


# ── Audit log ────────────────────────────────────────────────────────────────

def log_query(
    question: str,
    answer: str,
    citations_json: str,
    duration_ms: int,
    session_id: str | None,
) -> int:
    if not get_settings().persist_user_data:
        return 0

    with get_connection() as conn:
        cur = conn.execute(
            """INSERT INTO audit_log(session_id, question, answer, citations, duration_ms, created_at)
               VALUES (?,?,?,?,?,?)""",
            (session_id, question, answer, citations_json, duration_ms, _now()),
        )
        log_id = cur.lastrowid or 0

    if session_id:
        touch_session(session_id, title_hint=question)

    return int(log_id or 0)


def get_audit_log(limit: int = 100, offset: int = 0) -> list[dict]:
    with get_connection() as conn:
        rows = conn.execute(
            "SELECT * FROM audit_log ORDER BY id DESC LIMIT ? OFFSET ?",
            (limit, offset),
        ).fetchall()
        return [dict(r) for r in rows]


def get_session_messages(session_id: str) -> list[dict]:
    if not get_settings().persist_user_data:
        return []
    with get_connection() as conn:
        rows = conn.execute(
            """SELECT id, question, answer, created_at, duration_ms
               FROM audit_log
               WHERE session_id = ?
               ORDER BY id ASC""",
            (session_id,),
        ).fetchall()
        return [dict(r) for r in rows]


# ── Feedback ─────────────────────────────────────────────────────────────────

def save_feedback(
    log_id: int | None,
    session_id: str | None,
    question: str,
    answer: str,
    rating: str,
    correction: str | None,
    *,
    org_id: str | None = None,
    agent_id: str | None = None,
    user_id: str | None = None,
    message_id: str | None = None,
    mode: str | None = None,
    model_name: str | None = None,
) -> int:
    """Persist thumbs up/down for model training (always writes)."""
    if rating not in {"up", "down"}:
        raise ValueError("rating must be 'up' or 'down'")

    now = _now()
    model = model_name or get_settings().ollama_model

    with get_connection() as conn:
        existing = None
        if message_id:
            existing = conn.execute(
                "SELECT id FROM feedback WHERE message_id = ?",
                (message_id,),
            ).fetchone()

        if existing:
            eid = existing["id"] if isinstance(existing, dict) else existing[0]
            conn.execute(
                """UPDATE feedback SET
                       log_id=?, session_id=?, question=?, answer=?, rating=?,
                       correction=?, org_id=COALESCE(?, org_id),
                       agent_id=COALESCE(?, agent_id),
                       user_id=COALESCE(?, user_id),
                       mode=COALESCE(?, mode),
                       model_name=COALESCE(?, model_name),
                       updated_at=?,
                       reviewed=0
                   WHERE id=?""",
                (
                    log_id,
                    session_id,
                    question,
                    answer,
                    rating,
                    correction,
                    org_id,
                    agent_id,
                    user_id,
                    mode,
                    model,
                    now,
                    eid,
                ),
            )
            return int(eid)

        cur = conn.execute(
            """INSERT INTO feedback(
                   log_id, session_id, question, answer, rating, correction,
                   org_id, agent_id, user_id, message_id, mode, model_name,
                   created_at, updated_at
               ) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?)""",
            (
                log_id,
                session_id,
                question,
                answer,
                rating,
                correction,
                org_id,
                agent_id,
                user_id,
                message_id,
                mode,
                model,
                now,
                now,
            ),
        )
        return int(cur.lastrowid or 0)


def get_feedback(
    limit: int = 100,
    offset: int = 0,
    rating: str | None = None,
    reviewed: int | None = None,
    org_id: str | None = None,
    agent_id: str | None = None,
) -> list[dict]:
    clauses = []
    params: list = []
    if rating:
        clauses.append("rating = ?")
        params.append(rating)
    if reviewed is not None:
        clauses.append("reviewed = ?")
        params.append(reviewed)
    if org_id:
        clauses.append("org_id = ?")
        params.append(org_id)
    if agent_id:
        clauses.append("agent_id = ?")
        params.append(agent_id)
    where = ("WHERE " + " AND ".join(clauses)) if clauses else ""
    with get_connection() as conn:
        rows = conn.execute(
            f"SELECT * FROM feedback {where} ORDER BY id DESC LIMIT ? OFFSET ?",
            (*params, limit, offset),
        ).fetchall()
        return [dict(r) for r in rows]


def feedback_rating_counts(
    *,
    org_id: str | None = None,
    agent_id: str | None = None,
) -> dict[str, int]:
    """Return {up, down, total, with_correction} for scoped feedback."""
    clauses: list[str] = []
    params: list = []
    if org_id:
        clauses.append("org_id = ?")
        params.append(org_id)
    if agent_id:
        clauses.append("agent_id = ?")
        params.append(agent_id)
    where = ("WHERE " + " AND ".join(clauses)) if clauses else ""
    with get_connection() as conn:
        rows = conn.execute(
            f"""
            SELECT
              COALESCE(SUM(CASE WHEN rating = 'up' THEN 1 ELSE 0 END), 0) AS up_count,
              COALESCE(SUM(CASE WHEN rating = 'down' THEN 1 ELSE 0 END), 0) AS down_count,
              COUNT(*) AS total,
              COALESCE(SUM(CASE WHEN correction IS NOT NULL AND TRIM(correction) != '' THEN 1 ELSE 0 END), 0) AS with_correction
            FROM feedback
            {where}
            """,
            params,
        ).fetchone()
        if not rows:
            return {"up": 0, "down": 0, "total": 0, "with_correction": 0}
        d = dict(rows)
        return {
            "up": int(d.get("up_count") or 0),
            "down": int(d.get("down_count") or 0),
            "total": int(d.get("total") or 0),
            "with_correction": int(d.get("with_correction") or 0),
        }


def feedback_counts_by_agent(org_id: str) -> dict[str, dict[str, int]]:
    """Map agent_id -> rating counts for an org."""
    with get_connection() as conn:
        rows = conn.execute(
            """
            SELECT
              agent_id,
              COALESCE(SUM(CASE WHEN rating = 'up' THEN 1 ELSE 0 END), 0) AS up_count,
              COALESCE(SUM(CASE WHEN rating = 'down' THEN 1 ELSE 0 END), 0) AS down_count,
              COUNT(*) AS total,
              COALESCE(SUM(CASE WHEN correction IS NOT NULL AND TRIM(correction) != '' THEN 1 ELSE 0 END), 0) AS with_correction
            FROM feedback
            WHERE org_id = ? AND agent_id IS NOT NULL AND agent_id != ''
            GROUP BY agent_id
            """,
            (org_id,),
        ).fetchall()
    out: dict[str, dict[str, int]] = {}
    for row in rows:
        d = dict(row)
        aid = str(d.get("agent_id") or "")
        if not aid:
            continue
        out[aid] = {
            "up": int(d.get("up_count") or 0),
            "down": int(d.get("down_count") or 0),
            "total": int(d.get("total") or 0),
            "with_correction": int(d.get("with_correction") or 0),
        }
    return out


def feedback_daily_counts(*, org_id: str, days: int = 7) -> list[dict[str, int | str]]:
    """Daily thumbs up/down for the last N days (inclusive)."""
    from datetime import date, timedelta

    days = max(1, min(int(days), 30))
    start = date.today() - timedelta(days=days - 1)
    cutoff = start.isoformat()
    with get_connection() as conn:
        rows = conn.execute(
            """
            SELECT
              SUBSTR(created_at, 1, 10) AS day,
              COALESCE(SUM(CASE WHEN rating = 'up' THEN 1 ELSE 0 END), 0) AS up_count,
              COALESCE(SUM(CASE WHEN rating = 'down' THEN 1 ELSE 0 END), 0) AS down_count
            FROM feedback
            WHERE org_id = ? AND SUBSTR(created_at, 1, 10) >= ?
            GROUP BY SUBSTR(created_at, 1, 10)
            ORDER BY day ASC
            """,
            (org_id, cutoff),
        ).fetchall()
    by_day = {str(dict(r)["day"]): dict(r) for r in rows if dict(r).get("day")}
    out: list[dict[str, int | str]] = []
    for i in range(days):
        d = start + timedelta(days=i)
        key = d.isoformat()
        row = by_day.get(key, {})
        out.append(
            {
                "date": key,
                "up": int(row.get("up_count") or 0),
                "down": int(row.get("down_count") or 0),
            }
        )
    return out


def get_training_pairs(
    *,
    org_id: str | None = None,
    agent_id: str | None = None,
    limit: int = 5000,
) -> list[dict]:
    return get_feedback(limit=limit, offset=0, org_id=org_id, agent_id=agent_id)


def update_feedback_review(feedback_id: int, reviewed: int) -> None:
    with get_connection() as conn:
        conn.execute(
            "UPDATE feedback SET reviewed=?, updated_at=? WHERE id=?",
            (reviewed, _now(), feedback_id),
        )


def update_feedback_correction(feedback_id: int, correction: str | None) -> None:
    with get_connection() as conn:
        conn.execute(
            "UPDATE feedback SET correction=?, updated_at=? WHERE id=?",
            (correction, _now(), feedback_id),
        )


# ── Sessions ─────────────────────────────────────────────────────────────────

def _title_from_question(question: str) -> str:
    text = (question or "").strip()
    if text.startswith("[uploaded PDF]"):
        text = text.replace("[uploaded PDF]", "").strip() or "PDF upload"
    text = " ".join(text.split())
    if len(text) > 48:
        return text[:45].rstrip() + "…"
    return text or "New chat"


def touch_session(session_id: str, title_hint: str | None = None) -> None:
    now = _now()
    title = _title_from_question(title_hint) if title_hint else None
    with get_connection() as conn:
        existing = conn.execute(
            "SELECT id, title FROM sessions WHERE id = ?",
            (session_id,),
        ).fetchone()
        if existing is None:
            conn.execute(
                """INSERT INTO sessions(id, title, created_at, last_seen, message_count)
                   VALUES (?,?,?,?,1)""",
                (session_id, title or "New chat", now, now),
            )
        else:
            existing_title = existing["title"] if isinstance(existing, dict) else existing[1]
            if title and not existing_title:
                conn.execute(
                    """UPDATE sessions
                       SET last_seen=?, message_count=message_count+1, title=?
                       WHERE id=?""",
                    (now, title, session_id),
                )
            else:
                conn.execute(
                    """UPDATE sessions
                       SET last_seen=?, message_count=message_count+1
                       WHERE id=?""",
                    (now, session_id),
                )


def list_conversations(limit: int = 40) -> list[dict]:
    if not get_settings().persist_user_data:
        return []
    with get_connection() as conn:
        orphan = conn.execute(
            """
            SELECT DISTINCT a.session_id
            FROM audit_log a
            WHERE a.session_id IS NOT NULL
              AND a.session_id != ''
              AND NOT EXISTS (SELECT 1 FROM sessions s WHERE s.id = a.session_id)
            """
        ).fetchall()
        now = _now()
        for row in orphan:
            sid = row["session_id"] if isinstance(row, dict) else row[0]
            first = conn.execute(
                "SELECT question, created_at FROM audit_log WHERE session_id=? ORDER BY id ASC LIMIT 1",
                (sid,),
            ).fetchone()
            last = conn.execute(
                "SELECT created_at FROM audit_log WHERE session_id=? ORDER BY id DESC LIMIT 1",
                (sid,),
            ).fetchone()
            count_row = conn.execute(
                "SELECT COUNT(*) AS c FROM audit_log WHERE session_id=?",
                (sid,),
            ).fetchone()
            count = count_row["c"] if isinstance(count_row, dict) else count_row[0]
            first_q = first["question"] if first and isinstance(first, dict) else (first[0] if first else "Chat")
            first_at = first["created_at"] if first and isinstance(first, dict) else (first[1] if first else now)
            last_at = last["created_at"] if last and isinstance(last, dict) else (last[0] if last else now)
            if is_postgres():
                conn.execute(
                    """INSERT INTO sessions(id, title, created_at, last_seen, message_count)
                       VALUES (?,?,?,?,?) ON CONFLICT (id) DO NOTHING""",
                    (sid, _title_from_question(first_q), first_at, last_at, count),
                )
            else:
                conn.execute(
                    """INSERT OR IGNORE INTO sessions(id, title, created_at, last_seen, message_count)
                       VALUES (?,?,?,?,?)""",
                    (sid, _title_from_question(first_q), first_at, last_at, count),
                )

        rows = conn.execute(
            """
            SELECT
                s.id,
                COALESCE(
                    NULLIF(s.title, ''),
                    (
                        SELECT CASE
                            WHEN a.question LIKE '[uploaded PDF]%' THEN trim(replace(a.question, '[uploaded PDF]', ''))
                            ELSE a.question
                        END
                        FROM audit_log a
                        WHERE a.session_id = s.id
                        ORDER BY a.id ASC
                        LIMIT 1
                    ),
                    'New chat'
                ) AS title,
                s.created_at,
                s.last_seen,
                s.message_count,
                (
                    SELECT COUNT(*) FROM audit_log a WHERE a.session_id = s.id
                ) AS turn_count
            FROM sessions s
            WHERE EXISTS (SELECT 1 FROM audit_log a WHERE a.session_id = s.id)
            ORDER BY s.last_seen DESC
            LIMIT ?
            """,
            (limit,),
        ).fetchall()
        return [dict(r) for r in rows]


def delete_conversation(session_id: str) -> None:
    with get_connection() as conn:
        conn.execute("DELETE FROM audit_log WHERE session_id = ?", (session_id,))
        conn.execute("DELETE FROM feedback WHERE session_id = ?", (session_id,))
        conn.execute("DELETE FROM sessions WHERE id = ?", (session_id,))


def new_session_id() -> str:
    return str(uuid.uuid4())


def db_backend() -> str:
    return "postgres" if is_postgres() else "sqlite"


def _now() -> str:
    return datetime.utcnow().isoformat(timespec="seconds") + "Z"
