"""Site model functions — CRUD, versioning, fields, actions, sessions, variants, waitlist."""

import json
import logging
import secrets
from datetime import datetime, timezone

logger = logging.getLogger(__name__)

from app.crypto import (
    decrypt_entity,
    decrypt_user_value,
    encrypt_entity,
    encrypt_user_submission,
    encrypt_user_value,
    encrypt_value,
    hash_value,
    is_encrypted,
    key_is_configured,
    try_decrypt,
    try_decrypt_entity,
    try_decrypt_user_submission,
    try_decrypt_user_value,
)
from app.db import get_db

# Import helpers
from app.helpers import (
    _check_spam_content,
    _current_month,
    _decrypt_submission_fields,
    _decrypt_user_row,
    _encrypt_submission_field,
    _get_site_owner_password_hash,
    _is_encrypted_value,
    parse_user_agent,
    resolve_geo,
)

# Import TIERS from config
from app.models_config import TIERS


def find_site_by_token(token):
    conn = None
    try:
        conn = get_db()
        site = conn.execute("SELECT * FROM sites WHERE token = ?", (token,)).fetchone()
    finally:
        if conn:
            conn.close()
    return dict(site) if site else None


def find_site_by_domain(domain):
    """Find a site by its custom_domain column."""
    conn = None
    try:
        conn = get_db()
        site = conn.execute("SELECT * FROM sites WHERE custom_domain = ?", (domain,)).fetchone()
    finally:
        if conn:
            conn.close()
    return dict(site) if site else None


def add_site(name, owner_email, smtp_from=None, user_id=None, field_config=None, webhook_url=None, team_id=None):
    import json

    token = f"fr-{secrets.token_hex(6)}"
    conn = None
    try:
        conn = get_db()
        fc_json = json.dumps(field_config) if field_config else None
        if team_id:
            conn.execute(
                "INSERT INTO sites (token, name, owner_email, smtp_from, user_id, field_config, webhook_url, team_id) VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
                (token, name, owner_email, smtp_from, user_id, fc_json, webhook_url, team_id),
            )
        else:
            conn.execute(
                "INSERT INTO sites (token, name, owner_email, smtp_from, user_id, field_config, webhook_url) VALUES (?, ?, ?, ?, ?, ?, ?)",
                (token, name, owner_email, smtp_from, user_id, fc_json, webhook_url),
            )
        conn.commit()
        site = conn.execute("SELECT * FROM sites WHERE token = ? ", (token,)).fetchone()
    finally:
        if conn:
            conn.close()
    site_dict = dict(site)
    # Encrypt webhook_url now that we have site_id
    if site_dict.get("webhook_url"):
        encrypted = encrypt_value(site_dict["id"], site_dict["webhook_url"])
        if encrypted:
            conn = None
            try:
                conn = get_db()
                conn.execute("UPDATE sites SET webhook_url = ? WHERE id = ?", (encrypted, site_dict["id"]))
                conn.commit()
            finally:
                if conn:
                    conn.close()
            site_dict["webhook_url"] = encrypted
    return site_dict


def list_sites():
    conn = None
    try:
        conn = get_db()
        sites = conn.execute("SELECT * FROM sites ORDER BY created_at DESC").fetchall()
    finally:
        if conn:
            conn.close()
    return [dict(s) for s in sites]


def delete_site(site_id):
    """Delete a site and all associated data (submissions, analytics, webhooks)."""
    conn = None
    try:
        conn = get_db()
        conn.execute("DELETE FROM webhook_logs WHERE site_id = ?", (site_id,))
        conn.execute(
            "DELETE FROM email_send_log WHERE submission_id IN (SELECT id FROM submissions WHERE site_id = ?)",
            (site_id,),
        )
        conn.execute("DELETE FROM documents WHERE site_id = ?", (site_id,))
        conn.execute("DELETE FROM site_analytics WHERE site_id = ?", (site_id,))
        conn.execute("DELETE FROM submissions WHERE site_id = ?", (site_id,))
        conn.execute("DELETE FROM sites WHERE id = ?", (site_id,))
        conn.commit()
    finally:
        if conn:
            conn.close()


def get_site(site_id):
    """Get a site by ID. Returns dict or None."""
    conn = None
    try:
        conn = get_db()
        site = conn.execute("SELECT * FROM sites WHERE id = ?", (site_id,)).fetchone()
    finally:
        if conn:
            conn.close()
    if not site:
        return None
    site_dict = dict(site)
    site_dict["webhook_url"] = try_decrypt(site_id, site_dict.get("webhook_url"))
    return site_dict


def get_user_sites(user_id):
    """Get all sites owned by a user directly OR accessible through team membership."""
    conn = None
    try:
        conn = get_db()
        # Personal sites (user_id) + team sites (team_id where user is active member)
        rows = conn.execute(
            """
            SELECT s.*,
                   CASE WHEN s.team_id IS NULL THEN 'personal'
                        ELSE 'team' END AS site_type,
                   tm.role AS team_role,
                   t.name AS team_name
            FROM sites s
            LEFT JOIN team_members tm ON s.team_id = tm.team_id AND tm.user_id = ? AND tm.status = 'active'
            LEFT JOIN teams t ON s.team_id = t.id
            WHERE s.user_id = ? OR (s.team_id IS NOT NULL AND tm.status = 'active')
            ORDER BY s.created_at DESC
        """,
            (user_id, user_id),
        ).fetchall()
    finally:
        if conn:
            conn.close()
    sites = [dict(r) for r in rows]
    # Decrypt webhook_url for each site
    for s in sites:
        sid = s["id"]
        s["webhook_url"] = try_decrypt(sid, s.get("webhook_url"))
    return sites


# ─── Form versioning ────────────────────────────────────────────────────────────────


def get_version_limit(tier_key):
    """Return max version count for a tier."""
    limits = {
        "free": 10,
        "starter": 50,
        "pro": None,  # unlimited
        "teams": None,  # unlimited
        "teams_plus": None,
    }
    return limits.get(tier_key, 10)


def save_site_version(site_id, changed_by, change_reason=None):
    """Save current site field_config as a new version snapshot.

    Returns the new version number, or None if no field_config to save.
    """
    conn = None
    try:
        conn = get_db()
        site = conn.execute("SELECT field_config, metadata FROM sites WHERE id = ?", (site_id,)).fetchone()
        if not site or not site["field_config"]:
            return None

        # Get current max version for this site
        max_ver = conn.execute(
            "SELECT COALESCE(MAX(version), 0) as max_v FROM form_versions WHERE site_id = ?", (site_id,)
        ).fetchone()["max_v"]
        new_version = max_ver + 1

        conn.execute(
            "INSERT INTO form_versions (site_id, version, field_config, metadata, changed_by, change_reason) VALUES (?, ?, ?, ?, ?, ?)",
            (site_id, new_version, site["field_config"], site.get("metadata"), changed_by, change_reason),
        )
        conn.commit()
        return new_version
    finally:
        if conn:
            conn.close()


def get_site_versions(site_id, limit=None):
    """Get version history for a site, newest first.

    Args:
        site_id: Site ID
        limit: Max versions to return (None = all)

    Returns:
        List of version dicts
    """
    conn = None
    try:
        conn = get_db()
        if limit:
            rows = conn.execute(
                "SELECT * FROM form_versions WHERE site_id = ? ORDER BY version DESC LIMIT ?", (site_id, limit)
            ).fetchall()
        else:
            rows = conn.execute(
                "SELECT * FROM form_versions WHERE site_id = ? ORDER BY version DESC", (site_id,)
            ).fetchall()

        versions = []
        for row in rows:
            v = dict(row)
            try:
                v["field_config_parsed"] = json.loads(v.get("field_config", "[]"))
            except (json.JSONDecodeError, TypeError):
                v["field_config_parsed"] = []
            try:
                v["metadata_parsed"] = json.loads(v.get("metadata", "{}"))
            except (json.JSONDecodeError, TypeError):
                v["metadata_parsed"] = {}
            versions.append(v)
        return versions
    finally:
        if conn:
            conn.close()


def get_site_version(site_id, version_id):
    """Get a specific version by ID.

    Returns version dict or None.
    """
    conn = None
    try:
        conn = get_db()
        row = conn.execute("SELECT * FROM form_versions WHERE id = ? AND site_id = ?", (version_id, site_id)).fetchone()
        if not row:
            return None
        v = dict(row)
        try:
            v["field_config_parsed"] = json.loads(v.get("field_config", "[]"))
        except (json.JSONDecodeError, TypeError):
            v["field_config_parsed"] = []
        try:
            v["metadata_parsed"] = json.loads(v.get("metadata", "{}"))
        except (json.JSONDecodeError, TypeError):
            v["metadata_parsed"] = {}
        return v
    finally:
        if conn:
            conn.close()


def rollback_site_version(site_id, version_id, changed_by, change_reason=None):
    """Rollback a site to a previous version.

    Saves the current config as a new version first, then applies the old config.

    Returns the new version number after rollback, or None on failure.
    """
    conn = None
    try:
        conn = get_db()
        # Get target version
        target = conn.execute(
            "SELECT * FROM form_versions WHERE id = ? AND site_id = ?", (version_id, site_id)
        ).fetchone()
        if not target:
            return None

        # Get current max version
        max_ver = conn.execute(
            "SELECT COALESCE(MAX(version), 0) as max_v FROM form_versions WHERE site_id = ?", (site_id,)
        ).fetchone()["max_v"]

        # Save current state as new version before overwriting
        current_site = conn.execute("SELECT field_config, metadata FROM sites WHERE id = ?", (site_id,)).fetchone()
        if current_site and current_site["field_config"]:
            conn.execute(
                "INSERT INTO form_versions (site_id, version, field_config, metadata, changed_by, change_reason) VALUES (?, ?, ?, ?, ?, ?)",
                (
                    site_id,
                    max_ver + 1,
                    current_site["field_config"],
                    current_site.get("metadata"),
                    changed_by,
                    "Pre-rollback snapshot",
                ),
            )

        # Apply rollback
        conn.execute(
            "UPDATE sites SET field_config = ?, metadata = ? WHERE id = ?",
            (target["field_config"], target.get("metadata"), site_id),
        )
        conn.commit()
        return max_ver + 1
    finally:
        if conn:
            conn.close()


def prune_old_versions(site_id, max_versions):
    """Prune old versions beyond the tier limit.

    Returns count of pruned versions.
    """
    if max_versions is None:
        return 0
    conn = None
    try:
        conn = get_db()
        # Get versions to prune (oldest first, beyond limit)
        versions_to_keep = conn.execute(
            "SELECT id FROM form_versions WHERE site_id = ? ORDER BY version DESC LIMIT ?", (site_id, max_versions)
        ).fetchall()
        keep_ids = [v["id"] for v in versions_to_keep]

        if len(keep_ids) <= max_versions:
            return 0

        # Actually, we want to delete everything beyond the top N
        # The query above already limits to what we keep
        if len(keep_ids) > 0:
            placeholders = ",".join(["?"] * len(keep_ids))
            conn.execute(
                f"DELETE FROM form_versions WHERE site_id = ? AND id NOT IN ({placeholders})", [site_id] + keep_ids
            )
        conn.commit()

        # Count pruned
        remaining = conn.execute("SELECT COUNT(*) as cnt FROM form_versions WHERE site_id = ?", (site_id,)).fetchone()[
            "cnt"
        ]
        return max(0, len(keep_ids) - remaining)
    finally:
        if conn:
            conn.close()


def update_site_fields(
    site_id,
    field_config,
    webhook_url=None,
    webhook_enabled=None,
    webhook_events=None,
    honeypot_enabled=None,
    spam_filter_enabled=None,
    rate_limit_enabled=None,
    rate_limit_burst=None,
    rate_limit_refill=None,
    changed_by=None,
    change_reason=None,
):
    """Update field config, webhook settings, abuse protection, and rate limit config for a site.

    If field_config changes, saves the current state as a version snapshot first.
    """
    import json

    conn = None
    try:
        conn = get_db()
        updates = []
        params = []

        # Save current config as a version snapshot if field_config is changing
        if field_config is not None and changed_by:
            current = conn.execute("SELECT field_config, metadata FROM sites WHERE id = ?", (site_id,)).fetchone()
            if current and current["field_config"]:
                max_ver = conn.execute(
                    "SELECT COALESCE(MAX(version), 0) as max_v FROM form_versions WHERE site_id = ?", (site_id,)
                ).fetchone()["max_v"]
                conn.execute(
                    "INSERT INTO form_versions (site_id, version, field_config, metadata, changed_by, change_reason) VALUES (?, ?, ?, ?, ?, ?)",
                    (
                        site_id,
                        max_ver + 1,
                        current["field_config"],
                        current["metadata"],
                        changed_by,
                        change_reason or "Field config updated",
                    ),
                )
                conn.commit()

        if field_config is not None:
            updates.append("field_config = ?")
            params.append(json.dumps(field_config) if isinstance(field_config, list) else field_config)
        if webhook_url is not None:
            updates.append("webhook_url = ?")
            params.append(encrypt_value(site_id, webhook_url) if webhook_url else None)
        if webhook_enabled is not None:
            updates.append("webhook_enabled = ?")
            params.append(int(webhook_enabled))
        if webhook_events is not None:
            updates.append("webhook_events = ?")
            params.append(webhook_events)
        if honeypot_enabled is not None:
            updates.append("honeypot_enabled = ?")
            params.append(int(honeypot_enabled))
        if spam_filter_enabled is not None:
            updates.append("spam_filter_enabled = ?")
            params.append(int(spam_filter_enabled))
        if rate_limit_enabled is not None:
            updates.append("rate_limit_enabled = ?")
            params.append(1 if rate_limit_enabled else 0)
        if rate_limit_burst is not None:
            updates.append("rate_limit_burst = ?")
            params.append(max(1, int(rate_limit_burst)))
        if rate_limit_refill is not None:
            updates.append("rate_limit_refill = ?")
            params.append(max(0.01, float(rate_limit_refill)))

        if updates:
            params.append(site_id)
            conn.execute(f"UPDATE sites SET {', '.join(updates)} WHERE id = ?", params)
            conn.commit()

        site = conn.execute("SELECT * FROM sites WHERE id = ?", (site_id,)).fetchone()
    finally:
        if conn:
            conn.close()
    return dict(site) if site else None


def parse_site_fields(site):
    """Parse field_config JSON from site dict. Returns list of field definitions or None."""
    if not site or not site.get("field_config"):
        return None
    import json

    try:
        return json.loads(site["field_config"])
    except (json.JSONDecodeError, TypeError):
        return None


# ─── Column-name allowlist helpers ──────────────────────────────────────

_SITES_ALLOWED_COLUMNS = {
    "name",
    "description",
    "token",
    "active",
    "theme",
    "logo_url",
    "favicon_url",
    "domain",
    "redirect_url",
    "css",
    "js",
    "settings",
    "usage_limit",
    "usage_count",
    "created_at",
    "updated_at",
}

_FORM_ACTIONS_ALLOWED_COLUMNS = {
    "action_type",
    "config",
    "position",
    "active",
}

_FORM_VARIANTS_ALLOWED_COLUMNS = {
    "name",
    "fields",
    "config",
    "active",
    "position",
}

_FORM_SESSIONS_ALLOWED_COLUMNS = {
    "data",
    "submitted_at",
    "ip",
    "user_agent",
}


def _validate_columns(columns: list, allowed: set) -> list:
    """Validate column names against an allowlist. Returns list of 'col = ?' clauses."""
    valid = []
    for col in columns:
        if col in allowed:
            valid.append(f"{col} = ?")
        else:
            logger.warning(f"Blocked column in UPDATE: {col} (not in allowlist)")
    return valid


# ─── Tier enforcement & usage tracking ────────────────────────────────────────


def create_action(site_id, action_type, config, trigger_event="submission", enabled=True, execution_order=0):
    """Create a new form action for a site."""
    conn = None
    try:
        conn = get_db()
        cursor = conn.execute(
            """INSERT INTO form_actions (site_id, type, config, trigger_event, enabled, execution_order)
               VALUES (?, ?, ?, ?, ?, ?)""",
            (
                site_id,
                action_type,
                json.dumps(config) if isinstance(config, dict) else config,
                trigger_event,
                1 if enabled else 0,
                execution_order,
            ),
        )
        conn.commit()
        return cursor.lastrowid
    finally:
        if conn:
            conn.close()


def get_action(action_id):
    """Get a single action by ID."""
    conn = None
    try:
        conn = get_db()
        row = conn.execute("SELECT * FROM form_actions WHERE id = ?", (action_id,)).fetchone()
        if row:
            row = dict(row)
            if row.get("config"):
                try:
                    row["config"] = json.loads(row["config"])
                except (json.JSONDecodeError, TypeError):
                    pass
            return row
        return None
    finally:
        if conn:
            conn.close()


def get_actions_by_site(site_id):
    """Get all actions for a site, ordered by execution_order."""
    conn = None
    try:
        conn = get_db()
        rows = conn.execute(
            "SELECT * FROM form_actions WHERE site_id = ? ORDER BY execution_order ASC", (site_id,)
        ).fetchall()
        actions = []
        for row in rows:
            row = dict(row)
            if row.get("config"):
                try:
                    row["config"] = json.loads(row["config"])
                except (json.JSONDecodeError, TypeError):
                    pass
            actions.append(row)
        return actions
    finally:
        if conn:
            conn.close()


def update_action(action_id, **kwargs):
    """Update an action by ID. Only updates provided fields."""
    allowed = ("type", "config", "trigger_event", "enabled", "execution_order")
    updates = {k: v for k, v in kwargs.items() if k in allowed}
    if not updates:
        return False
    keys = list(updates.keys())
    values = list(updates.values())
    # Serialize config dict to JSON
    if "config" in updates and isinstance(updates["config"], dict):
        values[keys.index("config")] = json.dumps(updates["config"])
    values.append(action_id)
    conn = None
    try:
        conn = get_db()
        validated = _validate_columns(keys, _FORM_ACTIONS_ALLOWED_COLUMNS)
        if not validated:
            conn.close()
            return False
        conn.execute(f"UPDATE form_actions SET {', '.join(validated)} WHERE id = ?", values)
        conn.commit()
        return True
    finally:
        if conn:
            conn.close()


def delete_action(action_id):
    """Delete an action by ID."""
    conn = None
    try:
        conn = get_db()
        conn.execute("DELETE FROM form_actions WHERE id = ?", (action_id,))
        conn.commit()
        return True
    finally:
        if conn:
            conn.close()


def toggle_action(action_id):
    """Toggle enabled status of an action. Returns new enabled state."""
    conn = None
    try:
        conn = get_db()
        row = conn.execute("SELECT enabled FROM form_actions WHERE id = ?", (action_id,)).fetchone()
        if not row:
            return None
        new_state = 0 if row["enabled"] else 1
        conn.execute("UPDATE form_actions SET enabled = ? WHERE id = ?", (new_state, action_id))
        conn.commit()
        return bool(new_state)
    finally:
        if conn:
            conn.close()


# ─── Team model functions ─────────────────────────────────────────────────────


def create_variant(site_id, variant_key, name, field_config, success_message=None):
    """Create a new A/B test variant. Returns variant dict or None."""
    conn = None
    try:
        conn = get_db()
        existing = conn.execute(
            "SELECT id FROM form_variants WHERE site_id = ? AND variant_key = ? AND active = 1", (site_id, variant_key)
        ).fetchone()
        if existing:
            return None

        cursor = conn.execute(
            """INSERT INTO form_variants (site_id, variant_key, name, field_config, success_message, weight)
               VALUES (?, ?, ?, ?, ?, ?)""",
            (site_id, variant_key, name, json.dumps(field_config), success_message, 50),
        )
        conn.commit()
        return get_variant(cursor.lastrowid)
    finally:
        if conn:
            conn.close()


def get_variants(site_id, active_only=True):
    """Get all variants for a site. Returns list of variant dicts."""
    conn = None
    try:
        conn = get_db()
        query = "SELECT * FROM form_variants WHERE site_id = ?"
        params = [site_id]
        if active_only:
            query += " AND active = 1"
        query += " ORDER BY variant_key"

        rows = conn.execute(query, params).fetchall()
        result = []
        for row in rows:
            result.append(
                {
                    "id": row["id"],
                    "site_id": row["site_id"],
                    "variant_key": row["variant_key"],
                    "name": row["name"],
                    "field_config": json.loads(row["field_config"]) if row["field_config"] else [],
                    "success_message": row["success_message"],
                    "weight": row["weight"],
                    "active": bool(row["active"]),
                    "created_at": row["created_at"],
                }
            )
        return result
    finally:
        if conn:
            conn.close()


def get_variant(variant_id):
    """Get a single variant by ID. Returns dict or None."""
    conn = None
    try:
        conn = get_db()
        row = conn.execute("SELECT * FROM form_variants WHERE id = ?", (variant_id,)).fetchone()
        if not row:
            return None
        return {
            "id": row["id"],
            "site_id": row["site_id"],
            "variant_key": row["variant_key"],
            "name": row["name"],
            "field_config": json.loads(row["field_config"]) if row["field_config"] else [],
            "success_message": row["success_message"],
            "weight": row["weight"],
            "active": bool(row["active"]),
            "created_at": row["created_at"],
        }
    finally:
        if conn:
            conn.close()


def assign_variant(site_id, visitor_id):
    """Determine which variant a visitor should see.

    Uses a deterministic hash based on visitor_id for consistent assignment.
    Returns variant dict or None if no active test exists.
    """
    conn = None
    try:
        conn = get_db()
        variants = conn.execute(
            "SELECT * FROM form_variants WHERE site_id = ? AND active = 1 ORDER BY variant_key", (site_id,)
        ).fetchall()

        if len(variants) < 2:
            return None

        # Deterministic assignment based on visitor_id hash
        import hashlib

        hash_val = int(hashlib.sha256(str(visitor_id).encode()).hexdigest(), 16)
        total_weight = sum(v["weight"] for v in variants)
        bucket = hash_val % total_weight

        cumulative = 0
        for variant in variants:
            cumulative += variant["weight"]
            if bucket < cumulative:
                return {
                    "variant_key": variant["variant_key"],
                    "variant_id": variant["id"],
                    "field_config": json.loads(variant["field_config"]) if variant["field_config"] else [],
                    "success_message": variant["success_message"],
                }

        # Fallback to last variant
        variant = variants[-1]
        return {
            "variant_key": variant["variant_key"],
            "variant_id": variant["id"],
            "field_config": json.loads(variant["field_config"]) if variant["field_config"] else [],
            "success_message": variant["success_message"],
        }
    finally:
        if conn:
            conn.close()


def update_variant(variant_id, **kwargs):
    """Update a variant. Allowed: name, field_config, success_message, weight, active."""
    allowed = {"name", "field_config", "success_message", "weight", "active"}
    updates = {k: v for k, v in kwargs.items() if k in allowed and v is not None}
    if not updates:
        return False

    conn = None
    try:
        conn = get_db()
        for key, val in updates.items():
            if key == "field_config" and isinstance(val, (list, dict)):
                val = json.dumps(val)
            updates[key] = val

        keys = list(updates.keys())
        validated = _validate_columns(keys, _FORM_VARIANTS_ALLOWED_COLUMNS)
        if not validated:
            conn.close()
            return False
        values = list(updates.values()) + [variant_id]
        conn.execute(f"UPDATE form_variants SET {', '.join(validated)} WHERE id = ?", values)
        conn.commit()
        return get_variant(variant_id)
    finally:
        if conn:
            conn.close()


def deactivate_variant(variant_id):
    """Deactivate a variant. Returns True if deactivated."""
    conn = None
    try:
        conn = get_db()
        cursor = conn.execute("UPDATE form_variants SET active = 0 WHERE id = ?", (variant_id,))
        conn.commit()
        return cursor.rowcount > 0
    finally:
        if conn:
            conn.close()


def create_or_update_session(
    site_id,
    session_id,
    visitor_id,
    event,
    client_ip=None,
    user_agent=None,
    referrer=None,
    fields_viewed=None,
    last_field_reached=None,
    step_reached=None,
):
    """Create or update a form session record.

    For 'session_start': inserts new record.
    For 'session_update': updates existing record by (site_id, session_id).
    For 'session_complete': updates existing record with completed_at.

    Returns dict with session record on success, None on failure.
    """
    import hashlib

    conn = None
    try:
        conn = get_db()
        ip_hash = None
        if client_ip:
            ip_hash = hashlib.sha256(client_ip.encode()).hexdigest()

        if event == "session_start":
            cursor = conn.execute(
                """INSERT INTO form_sessions
                   (site_id, session_id, visitor_id, ip_hash, user_agent, referrer,
                    event, fields_viewed, last_field_reached, step_reached)
                   VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
                (
                    site_id,
                    session_id,
                    visitor_id,
                    ip_hash,
                    user_agent,
                    referrer,
                    event,
                    json.dumps(fields_viewed) if fields_viewed else None,
                    last_field_reached,
                    step_reached,
                ),
            )
            conn.commit()
            return {"ok": True, "id": cursor.lastrowid}
        elif event in ("session_update", "session_complete"):
            updates = ["event = ?"]
            values = [event]

            if fields_viewed is not None:
                updates.append("fields_viewed = ?")
                values.append(json.dumps(fields_viewed))
            if last_field_reached is not None:
                updates.append("last_field_reached = ?")
                values.append(last_field_reached)
            if step_reached is not None:
                updates.append("step_reached = ?")
                values.append(step_reached)
            if event == "session_complete":
                updates.append("completed_at = CURRENT_TIMESTAMP")

            values.append(site_id)
            values.append(session_id)

            conn.execute(
                f"UPDATE form_sessions SET {', '.join(updates)} WHERE site_id = ? AND session_id = ?",
                values,
            )
            conn.commit()
            return {"ok": True}
        return {"ok": False, "error": f"Unknown event: {event}"}
    except Exception as e:
        print(f"[model] create_or_update_session error: {e}")
        return None
    finally:
        if conn:
            conn.close()


def add_waitlist(email, source="landing"):
    """Add email to waitlist. Returns True if added, False if duplicate."""
    conn = None
    try:
        conn = get_db()
        conn.execute(
            "INSERT INTO waitlist (email, source) VALUES (?, ?)",
            (email.lower().strip(), source),
        )
        conn.commit()
        added = True
    except sqlite3.IntegrityError:
        added = False
    finally:
        if conn:
            conn.close()
    return added


def waitlist_count():
    """Return total waitlist count."""
    conn = None
    try:
        conn = get_db()
        row = conn.execute("SELECT COUNT(*) as cnt FROM waitlist").fetchone()
    finally:
        if conn:
            conn.close()
    return row["cnt"] if row else 0


# ─── User Auth ───────────────────────────────────────────────────────────────